From af4ee19703cacfc938549695dd0ec64fc63a3772 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Petr=20Balv=C3=ADn?= Date: Thu, 3 Sep 2026 10:00:00 +0200 Subject: [PATCH] feat: initial release Assisted-by: GLM 5.3 Flash --- .gitea/workflows/race.yml | 58 + .gitea/workflows/release.yml | 183 + .gitea/workflows/test.yml | 109 + .gitignore | 5 + CHANGELOG.md | 348 ++ CONTRIBUTING.md | 125 + LICENSE | 21 + README.md | 727 ++++ SECURITY.md | 39 + doc.go | 58 + docs/API.md | 3583 +++++++++++++++++ docs/ARCHITECTURE.md | 207 + docs/BENCHMARKING.md | 82 + docs/DEVELOPMENT.md | 204 + dtypes_facade_test.go | 1354 +++++++ example_test.go | 308 ++ examples/deconv/main.go | 165 + examples/fft/main.go | 62 + examples/fits/main.go | 98 + examples/helmholtz/main.go | 242 ++ examples/hmc/main.go | 102 + examples/netcdf/main.go | 108 + examples/ode-fit/main.go | 135 + examples/pde/main.go | 103 + examples/pendulum/main.go | 72 + examples/qmc/main.go | 77 + examples/regression/main.go | 118 + examples/spectral/main.go | 136 + examples/wavelets/main.go | 139 + facade_generated.go | 785 ++++ facade_helpers_test.go | 24 + facade_test.go | 80 + go.mod | 3 + grad/adjoint.go | 313 ++ grad/adjoint_test.go | 349 ++ grad/backwardcoverage_test.go | 918 +++++ grad/batched.go | 165 + grad/batched_test.go | 188 + grad/bench_perf_test.go | 292 ++ grad/bench_test.go | 112 + grad/broadcast_backward_pins_test.go | 145 + grad/complex.go | 503 +++ grad/complex_test.go | 539 +++ grad/concat.go | 128 + grad/concat_test.go | 167 + grad/data_ops_test.go | 200 + grad/doc.go | 82 + grad/example_test.go | 218 + grad/fuzz_test.go | 187 + grad/grad_guard_pins_test.go | 875 ++++ grad/gradient_hygiene_pins_test.go | 230 ++ grad/helpers_test.go | 47 + grad/hessian.go | 178 + grad/hessian_nil_guard_test.go | 40 + grad/hessian_test.go | 193 + grad/hmc.go | 194 + grad/hmc_test.go | 217 + grad/hvp_shape_pin_test.go | 38 + grad/matmul_adjoint_pins_test.go | 158 + grad/newtoncg.go | 248 ++ grad/newtoncg_test.go | 197 + grad/permute.go | 42 + grad/permute_test.go | 100 + grad/pool.go | 352 ++ grad/pool_bench_test.go | 200 + grad/pooled_sweep_pin_test.go | 164 + grad/shape.go | 80 + grad/spectral.go | 246 ++ grad/spectral_test.go | 370 ++ grad/sweep_test.go | 286 ++ grad/tensor.go | 2303 +++++++++++ grad/tensor_test.go | 343 ++ integrate/bench_perf_test.go | 256 ++ integrate/bench_test.go | 305 ++ integrate/budget_pins_test.go | 75 + integrate/cubature.go | 363 ++ integrate/cubature_heap_test.go | 112 + integrate/cubature_test.go | 114 + integrate/doc.go | 109 + integrate/dtypes_census_test.go | 536 +++ integrate/event_boundary_pins_test.go | 125 + integrate/event_guard_pins_test.go | 236 ++ integrate/example_test.go | 225 ++ integrate/fem2d.go | 376 ++ integrate/fem2d_test.go | 477 +++ integrate/fem2dnilmesh_test.go | 21 + integrate/fem3d.go | 552 +++ integrate/fem3d_test.go | 565 +++ integrate/filon.go | 234 ++ integrate/filon_accuracy_test.go | 330 ++ integrate/filon_bounds_regression_test.go | 68 + integrate/filon_test.go | 189 + integrate/heat2d_pin_test.go | 165 + integrate/helpers_test.go | 24 + integrate/ode.go | 878 ++++ integrate/ode_test.go | 475 +++ integrate/odebdf2.go | 260 ++ integrate/odebdf2_test.go | 202 + integrate/odebdfvar.go | 414 ++ integrate/odebdfvar_test.go | 298 ++ integrate/odeboundary.go | 139 + integrate/odeboundary_test.go | 152 + integrate/odecolloc.go | 548 +++ integrate/odecolloc_test.go | 273 ++ integrate/odedae.go | 331 ++ integrate/odedae_test.go | 267 ++ integrate/odeevent.go | 208 + integrate/odegrid_pin_test.go | 41 + integrate/oderow.go | 264 ++ integrate/oderow_test.go | 266 ++ integrate/pde.go | 256 ++ integrate/pde2d.go | 423 ++ integrate/pde2d_test.go | 238 ++ integrate/pde_scratch_pin_test.go | 64 + integrate/pde_test.go | 145 + integrate/pdeadvect.go | 291 ++ integrate/pdeadvect_test.go | 387 ++ integrate/quad.go | 269 ++ integrate/quad_cub_bench_test.go | 60 + integrate/quad_test.go | 201 + integrate/symplectic.go | 129 + integrate/symplectic2.go | 284 ++ integrate/symplectic2_test.go | 370 ++ integrate/symplectic_test.go | 178 + integrate/tolerance_pins_test.go | 196 + integrate/tridiag.go | 39 + integrate/wave2d_pins_test.go | 391 ++ integrate/wave_bench_test.go | 152 + internal/base/base.go | 428 ++ internal/base/factor_bench_test.go | 431 ++ internal/base/factor_crossover_test.go | 189 + internal/base/scalar_abs_pin_test.go | 129 + internal/base/tridiag.go | 51 + internal/base/wrap_pin_test.go | 57 + internal/core/argmax_topk_test.go | 215 + internal/core/arrayutil.go | 2046 ++++++++++ internal/core/arrayutil_edge_pins_test.go | 174 + internal/core/arrayutil_new_test.go | 215 + internal/core/arrayutil_test.go | 397 ++ internal/core/axis.go | 1563 +++++++ internal/core/axis_accuracy_test.go | 345 ++ internal/core/axis_bench_test.go | 132 + internal/core/axis_test.go | 266 ++ internal/core/basebridge.go | 36 + internal/core/bench_einsum2_test.go | 639 +++ internal/core/bench_kernels_ab_test.go | 274 ++ internal/core/bench_mat2_test.go | 188 + internal/core/bench_narrow_test.go | 342 ++ internal/core/bench_ops2_test.go | 499 +++ internal/core/bench_probes_test.go | 744 ++++ internal/core/bench_sort2_test.go | 372 ++ internal/core/bench_test.go | 318 ++ internal/core/bessel_climb_pin_test.go | 81 + .../core/bessel_realorder_accuracy_test.go | 330 ++ internal/core/besseljreal_test.go | 180 + internal/core/besselmod.go | 300 ++ internal/core/besselmod_test.go | 152 + internal/core/broadcast.go | 244 ++ internal/core/broadcast_test.go | 234 ++ .../core/canonical_partition_pins_test.go | 132 + internal/core/centralsum_kernels_test.go | 364 ++ internal/core/chunk_precision_pins_test.go | 342 ++ internal/core/complex_test.go | 269 ++ internal/core/concat.go | 285 ++ internal/core/concat_test.go | 232 ++ internal/core/cosm1_test.go | 102 + internal/core/cumprod_half_test.go | 29 + internal/core/data_test.go | 42 + internal/core/defect_family_pins_test.go | 504 +++ internal/core/diff.go | 103 + internal/core/dtypes_promote_test.go | 345 ++ internal/core/einsum.go | 1653 ++++++++ internal/core/einsum_bench_test.go | 98 + internal/core/einsum_foldorder_test.go | 31 + .../core/einsum_norm_special_pins_test.go | 562 +++ internal/core/einsum_special_pins_test.go | 279 ++ internal/core/einsum_test.go | 256 ++ internal/core/elementwise_view_test.go | 219 + internal/core/elliptic.go | 510 +++ internal/core/elliptic2_test.go | 134 + internal/core/elliptic_test.go | 293 ++ internal/core/expint.go | 164 + internal/core/expint_test.go | 155 + internal/core/extrema.go | 213 + internal/core/extrema_test.go | 187 + internal/core/float16.go | 183 + internal/core/float16_test.go | 978 +++++ internal/core/float32_test.go | 334 ++ internal/core/fresnel.go | 199 + internal/core/grid.go | 118 + internal/core/guard_pins_test.go | 822 ++++ internal/core/index.go | 233 ++ internal/core/index_test.go | 317 ++ internal/core/indexing2.go | 395 ++ internal/core/init.go | 121 + internal/core/int_exactness_pins_test.go | 677 ++++ internal/core/interpolate2d.go | 199 + internal/core/interpolate2d_test.go | 233 ++ internal/core/interpolate_knot_test.go | 38 + internal/core/jacobian.go | 100 + internal/core/jacobian_test.go | 156 + internal/core/knot_ellipsis_pins_test.go | 292 ++ internal/core/mask.go | 1283 ++++++ internal/core/mask_test.go | 414 ++ internal/core/mat.go | 1345 +++++++ internal/core/mat_kernel_bench_test.go | 159 + internal/core/mat_kernel_test.go | 192 + internal/core/mat_tile_bench_test.go | 101 + internal/core/math_bench_extra_test.go | 98 + internal/core/mathfunc.go | 427 ++ internal/core/mathfunc_test.go | 259 ++ internal/core/matrix2.go | 171 + internal/core/misc_bench_test.go | 80 + internal/core/narrow_compare_alloc_test.go | 622 +++ internal/core/narrow_dispatch_pin_test.go | 432 ++ internal/core/narrowdt.go | 325 ++ internal/core/onehot.go | 57 + internal/core/onehot_test.go | 59 + internal/core/ops.go | 1094 +++++ internal/core/ops_test.go | 420 ++ internal/core/orthopoly.go | 312 ++ internal/core/orthopoly_test.go | 201 + internal/core/overflow_refusal_pin_test.go | 142 + internal/core/pchip.go | 117 + internal/core/quasirandom.go | 242 ++ internal/core/quasirandom_test.go | 238 ++ internal/core/random.go | 224 ++ internal/core/random_test.go | 149 + internal/core/randomsub_test.go | 130 + internal/core/reduce.go | 1494 +++++++ internal/core/reduce_accuracy_test.go | 357 ++ internal/core/reduce_test.go | 330 ++ internal/core/reduction2.go | 879 ++++ internal/core/reduction_route_test.go | 80 + .../core/reduction_shape_consistency_test.go | 197 + internal/core/reference_walks_test.go | 965 +++++ internal/core/reshape.go | 375 ++ internal/core/reshape_pad_test.go | 544 +++ internal/core/runtime.go | 59 + internal/core/runtime_test.go | 89 + internal/core/scalar_clamp_pow_bench_test.go | 71 + internal/core/scan_compensation_test.go | 249 ++ internal/core/shape_overflow_pin_test.go | 66 + internal/core/smallops_test.go | 198 + internal/core/sort.go | 642 +++ internal/core/sort_test.go | 387 ++ internal/core/sparse.go | 373 ++ internal/core/special.go | 410 ++ internal/core/special2.go | 449 +++ internal/core/special2_test.go | 144 + internal/core/special_edge_accuracy_test.go | 511 +++ internal/core/special_test.go | 383 ++ internal/core/tensor.go | 1246 ++++++ internal/core/tensor_bench_test.go | 32 + internal/core/tensor_test.go | 193 + internal/core/validation_pins_test.go | 492 +++ internal/core/view_payload_reads_test.go | 150 + internal/core/view_test.go | 347 ++ internal/core/walks_pin_test.go | 180 + internal/core/where_half_test.go | 34 + internal/engine/engine.go | 125 + internal/engine/engine_bench_test.go | 25 + internal/engine/engine_test.go | 171 + io/bench_test.go | 353 ++ io/csv.go | 471 +++ io/csv_bench_test.go | 65 + io/csv_parity_test.go | 242 ++ io/csv_test.go | 222 + io/doc.go | 61 + io/example_test.go | 237 ++ io/extent_wrap_pins_test.go | 573 +++ io/fits.go | 561 +++ io/fits_hostile_pins_test.go | 272 ++ io/fits_table_pins_test.go | 206 + io/fits_test.go | 338 ++ io/fitstable.go | 1082 +++++ io/fitstable_keyword_pin_test.go | 80 + io/fitstable_test.go | 441 ++ io/fuzz_test.go | 474 +++ io/hdf5.go | 2303 +++++++++++ io/hdf5_bool_fixture_test.go | 90 + io/hdf5_continuation_pins_test.go | 161 + io/hdf5_structure_pins_test.go | 527 +++ io/hdf5_test.go | 748 ++++ io/hdf5write.go | 786 ++++ io/hdf5write_estimate_test.go | 23 + io/hdf5write_group.go | 717 ++++ io/hdf5write_object.go | 449 +++ io/hdf5write_test.go | 1276 ++++++ io/helpers_test.go | 34 + io/hostile_test.go | 267 ++ io/mmap.go | 181 + io/mmap_test.go | 165 + io/netcdf.go | 1077 +++++ io/netcdf_payload_test.go | 102 + io/netcdf_test.go | 634 +++ io/netcdfrecord_test.go | 158 + .../fuzz/FuzzLoadHDF5/2fe5cd680e6d8b63 | 2 + .../fuzz/FuzzLoadHDF5/a7ab50e52191807a | 2 + .../fuzz/FuzzLoadHDF5/ae4657f5205db80f | 2 + .../fuzz/FuzzLoadNetCDF/14c890c96cda7d9c | 2 + .../fuzz/FuzzLoadNetCDF/c4c632324ec4a450 | 2 + io/testdata/h5/bool_enum.h5 | Bin 0 -> 712 bytes io/testdata/h5/fixture.h5 | Bin 0 -> 11432 bytes io/testdata/h5/fletcher.h5 | Bin 0 -> 3612 bytes io/testdata/h5/latest.h5 | Bin 0 -> 2080 bytes io/walk_cycle_pins_test.go | 204 + io/wave_bench_test.go | 61 + io/zz_guard_test.go | 61 + justfile | 295 ++ leak_test.go | 176 + linalg/bench_decomp_test.go | 512 +++ linalg/bench_factor_test.go | 421 ++ linalg/bench_iterative_test.go | 251 ++ linalg/bench_sparse_test.go | 318 ++ linalg/cholupdate.go | 120 + linalg/cholupdate_test.go | 158 + linalg/complex_refusal_pins_test.go | 170 + linalg/csvd.go | 315 ++ linalg/csvd_test.go | 191 + linalg/decomp.go | 596 +++ linalg/decomp2.go | 1334 ++++++ linalg/decomp3.go | 343 ++ linalg/decomp3_test.go | 293 ++ linalg/decomp_test.go | 479 +++ linalg/decompositions_test.go | 693 ++++ linalg/det_factor_dispatch_test.go | 34 + linalg/detfloat32_test.go | 17 + linalg/doc.go | 84 + linalg/dtype_shape_refusal_pins_test.go | 435 ++ linalg/dtypes_census_test.go | 1000 +++++ linalg/eigenreal.go | 446 ++ linalg/eigenreal_test.go | 307 ++ linalg/example_test.go | 203 + linalg/expm.go | 328 ++ linalg/expm_test.go | 323 ++ linalg/extreme_scale_pins_test.go | 1056 +++++ linalg/fitpoly_test.go | 33 + linalg/geneigen.go | 131 + linalg/geneigen_test.go | 131 + linalg/gmres.go | 310 ++ linalg/gmres_test.go | 167 + linalg/helpers.go | 59 + linalg/helpers_test.go | 64 + linalg/initialisers_sparse_test.go | 165 + linalg/least_squares_test.go | 190 + linalg/linalg.go | 343 ++ linalg/linalg_complex_test.go | 261 ++ linalg/linalg_test.go | 187 + linalg/mat_scratch_test.go | 40 + linalg/mat_test.go | 248 ++ linalg/matrixfunc.go | 403 ++ linalg/matrixfunc_test.go | 125 + linalg/nonfinite_refusal_pins_test.go | 197 + linalg/perf_bench_test.go | 132 + linalg/pipeline.go | 393 ++ linalg/pipeline_test.go | 112 + linalg/polyroots.go | 82 + linalg/polyroots_test.go | 135 + linalg/rrqr.go | 422 ++ linalg/rrqr_test.go | 304 ++ linalg/schur.go | 263 ++ linalg/schur_test.go | 197 + linalg/solver_dtype_pins_test.go | 169 + linalg/solver_guard_pins_test.go | 135 + linalg/sparse_guard_pins_test.go | 519 +++ linalg/sparsebicg.go | 142 + linalg/sparsebicg_test.go | 195 + linalg/sparsecholesky.go | 524 +++ linalg/sparsecholesky_test.go | 535 +++ linalg/sparsecholupdate.go | 160 + linalg/sparsecholupdate_test.go | 314 ++ linalg/sparsecomplex.go | 745 ++++ linalg/sparsecomplex_test.go | 207 + linalg/sparsecsc.go | 309 ++ linalg/sparsecsc_split_test.go | 123 + linalg/sparsecsr.go | 211 + linalg/sparsefill.go | 493 +++ linalg/sparsefill_bench_test.go | 161 + linalg/sparsefill_test.go | 337 ++ linalg/sparsegeneral.go | 327 ++ linalg/sparsegeneral_test.go | 233 ++ linalg/sparseigen.go | 432 ++ linalg/sparseigen_test.go | 503 +++ linalg/sparseilu.go | 197 + linalg/sparseilu_test.go | 208 + linalg/sparselsqr.go | 807 ++++ linalg/sparselsqr_test.go | 468 +++ linalg/sparselu.go | 352 ++ linalg/sparselu_test.go | 307 ++ linalg/sparseops.go | 150 + linalg/sparseops_test.go | 226 ++ linalg/sparsesolve.go | 207 + linalg/sparsesolve_test.go | 375 ++ linalg/sparsexp.go | 197 + linalg/sparsexp_test.go | 251 ++ linalg/spline.go | 134 + linalg/spline_test.go | 109 + linalg/svdsolve.go | 138 + linalg/svdsolve_test.go | 180 + linalg/tridiag.go | 141 + linalg/tridiag_test.go | 167 + optim/bench_leastsqfit_test.go | 86 + optim/bench_solvers_test.go | 370 ++ optim/bench_test.go | 187 + optim/bounded_test.go | 165 + optim/brent.go | 144 + optim/brent_test.go | 203 + optim/broyden.go | 126 + optim/broyden_test.go | 282 ++ optim/budget_exit_pin_test.go | 108 + optim/callback_hygiene_pins_test.go | 142 + optim/callback_pins_test.go | 73 + optim/devolution.go | 186 + optim/devolution_test.go | 114 + optim/doc.go | 84 + optim/dtypes_census_test.go | 348 ++ optim/example_test.go | 219 + optim/fallback_guard_pins_test.go | 194 + optim/global2.go | 600 +++ optim/global2_test.go | 378 ++ optim/helpers_test.go | 24 + optim/jacobianparallel_test.go | 341 ++ optim/lbfgs.go | 551 +++ optim/lbfgs_test.go | 109 + optim/leastsq.go | 694 ++++ optim/leastsq_mirror_pin_test.go | 51 + optim/leastsq_test.go | 295 ++ optim/leastsqfit_test.go | 609 +++ optim/linear.go | 308 ++ optim/linear_test.go | 150 + optim/linesearch_pins_test.go | 650 +++ optim/nonfinite_exit_pins_test.go | 157 + optim/nonlinearconstr.go | 313 ++ optim/nonlinearconstr_test.go | 306 ++ optim/optimise.go | 412 ++ optim/optimise_test.go | 175 + optim/qp.go | 447 ++ optim/qp_test.go | 381 ++ optim/rootsystem.go | 437 ++ optim/rootsystem_test.go | 188 + optim/rowblocks_pin_test.go | 48 + optim/simplex.go | 814 ++++ optim/simplex_test.go | 450 +++ optim/solvercoverage_test.go | 321 ++ oracle_norace_test.go | 10 + oracle_race_test.go | 13 + oracle_test.go | 2852 +++++++++++++ plot/dtypes_census_test.go | 179 + plot/plot.go | 240 ++ plot/plot_test.go | 184 + signal/arma.go | 476 +++ signal/arma_test.go | 427 ++ signal/bench_perf_test.go | 355 ++ signal/bench_test.go | 138 + signal/chirp.go | 60 + signal/chirp_test.go | 158 + signal/complex_refusal_pins_test.go | 119 + signal/conv.go | 1011 +++++ signal/conv_bench_extra_test.go | 74 + signal/conv_bench_test.go | 48 + signal/conv_fused_pin_test.go | 302 ++ signal/conv_pooling_test.go | 481 +++ signal/conv_transpose_test.go | 120 + signal/conv_window_pins_test.go | 801 ++++ signal/correlate.go | 250 ++ signal/correlate_test.go | 152 + signal/dct.go | 209 + signal/dct_test.go | 225 ++ signal/doc.go | 55 + signal/dtypes_census_test.go | 723 ++++ signal/entry_gate_pins_test.go | 227 ++ signal/example_test.go | 195 + signal/examples_test.go | 36 + signal/fft.go | 874 ++++ signal/fft_nd_test.go | 247 ++ signal/fft_precision_test.go | 233 ++ signal/fft_test.go | 124 + signal/filter.go | 515 +++ signal/filter_dtype_pin_test.go | 92 + signal/filter_test.go | 349 ++ signal/filterdesign.go | 635 +++ signal/filterdesign_test.go | 217 + signal/filtfilt.go | 87 + signal/filtfilt_test.go | 273 ++ signal/float32_test.go | 26 + signal/helpers_test.go | 50 + signal/hilbert.go | 72 + signal/hilbert_test.go | 91 + signal/kalman.go | 941 +++++ signal/kalman_test.go | 739 ++++ signal/misc_bench_test.go | 159 + signal/nan_propagation_pins_test.go | 177 + signal/nufft.go | 124 + signal/nufft_test.go | 119 + signal/periodogram.go | 202 + signal/periodogram_test.go | 312 ++ signal/poisson.go | 112 + signal/poisson_complex_pins_test.go | 198 + signal/poisson_test.go | 143 + signal/poisson_view_pin_test.go | 99 + signal/poissondirichlet.go | 468 +++ signal/poissondirichlet_test.go | 130 + signal/pool.go | 815 ++++ signal/pooling2d_test.go | 194 + signal/rankfilter.go | 258 ++ signal/rankfilter_test.go | 353 ++ signal/resample.go | 249 ++ signal/resample_test.go | 184 + signal/stencil.go | 218 + signal/stencil_test.go | 256 ++ signal/stft.go | 158 + signal/stft_test.go | 128 + signal/transform_cache_pins_test.go | 112 + signal/wave_bench_test.go | 446 ++ signal/wavelet.go | 310 ++ signal/wavelet_test.go | 267 ++ signal/waveletdwt.go | 323 ++ signal/waveletdwt_test.go | 370 ++ signal/welch.go | 169 + signal/windows.go | 212 + signal/windows_test.go | 298 ++ spmd/arg.go | 248 ++ spmd/arg_test.go | 158 + spmd/bench_tcp_test.go | 177 + spmd/bench_test.go | 118 + spmd/collective.go | 285 ++ spmd/collective_test.go | 211 + spmd/doc.go | 42 + spmd/framepool.go | 110 + spmd/framepool_test.go | 212 + spmd/halo.go | 209 + spmd/halo_grid_test.go | 372 ++ spmd/halo_test.go | 180 + spmd/partition.go | 79 + spmd/partition_test.go | 44 + spmd/proc_test.go | 201 + spmd/reduce.go | 1291 ++++++ spmd/reduce_test.go | 886 ++++ spmd/transport.go | 237 ++ spmd/wire.go | 315 ++ spmd/wire_test.go | 175 + spmd/world.go | 576 +++ spmd/world_test.go | 416 ++ stats/anova.go | 228 ++ stats/anova_test.go | 224 ++ stats/bench_kernels_test.go | 216 + stats/bench_quantilereg_test.go | 38 + stats/bench_regression_test.go | 638 +++ stats/bench_test.go | 265 ++ stats/cdf.go | 735 ++++ stats/cdf_test.go | 278 ++ stats/cluster.go | 771 ++++ stats/cluster_test.go | 436 ++ stats/contingency.go | 355 ++ stats/contingency_test.go | 513 +++ stats/deep_tail_pin_test.go | 154 + stats/degenerate_input_pin_test.go | 74 + stats/distrib.go | 182 + stats/distrib2.go | 433 ++ stats/distrib2_test.go | 436 ++ stats/distrib_test.go | 94 + stats/doc.go | 90 + stats/dtypes_census_test.go | 710 ++++ stats/estimation_guards_pin_test.go | 152 + stats/example_test.go | 156 + stats/glm.go | 699 ++++ stats/glmkdemvn_test.go | 283 ++ stats/gmm_posterior_reference_test.go | 176 + stats/gp.go | 398 ++ stats/gp_test.go | 424 ++ stats/helpers_test.go | 64 + stats/hierarchy.go | 349 ++ stats/hierarchy_test.go | 242 ++ stats/histogram2d.go | 214 + stats/hmm.go | 565 +++ stats/hmm_test.go | 345 ++ stats/inference.go | 432 ++ stats/inference_test.go | 198 + stats/kde.go | 125 + stats/lasso.go | 509 +++ stats/lasso_test.go | 581 +++ stats/mixed.go | 826 ++++ stats/mixed_test.go | 484 +++ stats/multipletest.go | 112 + stats/multipletest_test.go | 175 + stats/mvn.go | 210 + stats/negbinom_far_pin_test.go | 60 + stats/noncentral.go | 441 ++ stats/noncentral_test.go | 497 +++ stats/pca.go | 413 ++ stats/pca_test.go | 441 ++ stats/pins_test.go | 1006 +++++ stats/poissonregression_test.go | 293 ++ stats/pool_reuse_pin_test.go | 56 + stats/quantile_newton_test.go | 243 ++ stats/quantile_test.go | 67 + stats/quantilereg.go | 458 +++ stats/quantilereg_test.go | 296 ++ stats/rankcorr.go | 278 ++ stats/rankcorr_test.go | 127 + stats/regression.go | 579 +++ stats/regression_edge_pin_test.go | 241 ++ stats/regression_test.go | 187 + stats/robust.go | 102 + stats/robustregression.go | 652 +++ stats/robustregression_test.go | 345 ++ stats/rolling.go | 253 ++ stats/rolling_prec_test.go | 230 ++ stats/smallops_test.go | 474 +++ stats/stats.go | 543 +++ stats/stats_random_test.go | 141 + stats/stats_test.go | 232 ++ stats/tail_and_extreme_pin_test.go | 394 ++ stats/variance_prec_test.go | 182 + stats/view_tail_pin_test.go | 191 + stats/zz_guard_test.go | 40 + 617 files changed, 191195 insertions(+) create mode 100644 .gitea/workflows/race.yml create mode 100644 .gitea/workflows/release.yml create mode 100644 .gitea/workflows/test.yml create mode 100644 .gitignore create mode 100644 CHANGELOG.md create mode 100644 CONTRIBUTING.md create mode 100644 LICENSE create mode 100644 README.md create mode 100644 SECURITY.md create mode 100644 doc.go create mode 100644 docs/API.md create mode 100644 docs/ARCHITECTURE.md create mode 100644 docs/BENCHMARKING.md create mode 100644 docs/DEVELOPMENT.md create mode 100644 dtypes_facade_test.go create mode 100644 example_test.go create mode 100644 examples/deconv/main.go create mode 100644 examples/fft/main.go create mode 100644 examples/fits/main.go create mode 100644 examples/helmholtz/main.go create mode 100644 examples/hmc/main.go create mode 100644 examples/netcdf/main.go create mode 100644 examples/ode-fit/main.go create mode 100644 examples/pde/main.go create mode 100644 examples/pendulum/main.go create mode 100644 examples/qmc/main.go create mode 100644 examples/regression/main.go create mode 100644 examples/spectral/main.go create mode 100644 examples/wavelets/main.go create mode 100644 facade_generated.go create mode 100644 facade_helpers_test.go create mode 100644 facade_test.go create mode 100644 go.mod create mode 100644 grad/adjoint.go create mode 100644 grad/adjoint_test.go create mode 100644 grad/backwardcoverage_test.go create mode 100644 grad/batched.go create mode 100644 grad/batched_test.go create mode 100644 grad/bench_perf_test.go create mode 100644 grad/bench_test.go create mode 100644 grad/broadcast_backward_pins_test.go create mode 100644 grad/complex.go create mode 100644 grad/complex_test.go create mode 100644 grad/concat.go create mode 100644 grad/concat_test.go create mode 100644 grad/data_ops_test.go create mode 100644 grad/doc.go create mode 100644 grad/example_test.go create mode 100644 grad/fuzz_test.go create mode 100644 grad/grad_guard_pins_test.go create mode 100644 grad/gradient_hygiene_pins_test.go create mode 100644 grad/helpers_test.go create mode 100644 grad/hessian.go create mode 100644 grad/hessian_nil_guard_test.go create mode 100644 grad/hessian_test.go create mode 100644 grad/hmc.go create mode 100644 grad/hmc_test.go create mode 100644 grad/hvp_shape_pin_test.go create mode 100644 grad/matmul_adjoint_pins_test.go create mode 100644 grad/newtoncg.go create mode 100644 grad/newtoncg_test.go create mode 100644 grad/permute.go create mode 100644 grad/permute_test.go create mode 100644 grad/pool.go create mode 100644 grad/pool_bench_test.go create mode 100644 grad/pooled_sweep_pin_test.go create mode 100644 grad/shape.go create mode 100644 grad/spectral.go create mode 100644 grad/spectral_test.go create mode 100644 grad/sweep_test.go create mode 100644 grad/tensor.go create mode 100644 grad/tensor_test.go create mode 100644 integrate/bench_perf_test.go create mode 100644 integrate/bench_test.go create mode 100644 integrate/budget_pins_test.go create mode 100644 integrate/cubature.go create mode 100644 integrate/cubature_heap_test.go create mode 100644 integrate/cubature_test.go create mode 100644 integrate/doc.go create mode 100644 integrate/dtypes_census_test.go create mode 100644 integrate/event_boundary_pins_test.go create mode 100644 integrate/event_guard_pins_test.go create mode 100644 integrate/example_test.go create mode 100644 integrate/fem2d.go create mode 100644 integrate/fem2d_test.go create mode 100644 integrate/fem2dnilmesh_test.go create mode 100644 integrate/fem3d.go create mode 100644 integrate/fem3d_test.go create mode 100644 integrate/filon.go create mode 100644 integrate/filon_accuracy_test.go create mode 100644 integrate/filon_bounds_regression_test.go create mode 100644 integrate/filon_test.go create mode 100644 integrate/heat2d_pin_test.go create mode 100644 integrate/helpers_test.go create mode 100644 integrate/ode.go create mode 100644 integrate/ode_test.go create mode 100644 integrate/odebdf2.go create mode 100644 integrate/odebdf2_test.go create mode 100644 integrate/odebdfvar.go create mode 100644 integrate/odebdfvar_test.go create mode 100644 integrate/odeboundary.go create mode 100644 integrate/odeboundary_test.go create mode 100644 integrate/odecolloc.go create mode 100644 integrate/odecolloc_test.go create mode 100644 integrate/odedae.go create mode 100644 integrate/odedae_test.go create mode 100644 integrate/odeevent.go create mode 100644 integrate/odegrid_pin_test.go create mode 100644 integrate/oderow.go create mode 100644 integrate/oderow_test.go create mode 100644 integrate/pde.go create mode 100644 integrate/pde2d.go create mode 100644 integrate/pde2d_test.go create mode 100644 integrate/pde_scratch_pin_test.go create mode 100644 integrate/pde_test.go create mode 100644 integrate/pdeadvect.go create mode 100644 integrate/pdeadvect_test.go create mode 100644 integrate/quad.go create mode 100644 integrate/quad_cub_bench_test.go create mode 100644 integrate/quad_test.go create mode 100644 integrate/symplectic.go create mode 100644 integrate/symplectic2.go create mode 100644 integrate/symplectic2_test.go create mode 100644 integrate/symplectic_test.go create mode 100644 integrate/tolerance_pins_test.go create mode 100644 integrate/tridiag.go create mode 100644 integrate/wave2d_pins_test.go create mode 100644 integrate/wave_bench_test.go create mode 100644 internal/base/base.go create mode 100644 internal/base/factor_bench_test.go create mode 100644 internal/base/factor_crossover_test.go create mode 100644 internal/base/scalar_abs_pin_test.go create mode 100644 internal/base/tridiag.go create mode 100644 internal/base/wrap_pin_test.go create mode 100644 internal/core/argmax_topk_test.go create mode 100644 internal/core/arrayutil.go create mode 100644 internal/core/arrayutil_edge_pins_test.go create mode 100644 internal/core/arrayutil_new_test.go create mode 100644 internal/core/arrayutil_test.go create mode 100644 internal/core/axis.go create mode 100644 internal/core/axis_accuracy_test.go create mode 100644 internal/core/axis_bench_test.go create mode 100644 internal/core/axis_test.go create mode 100644 internal/core/basebridge.go create mode 100644 internal/core/bench_einsum2_test.go create mode 100644 internal/core/bench_kernels_ab_test.go create mode 100644 internal/core/bench_mat2_test.go create mode 100644 internal/core/bench_narrow_test.go create mode 100644 internal/core/bench_ops2_test.go create mode 100644 internal/core/bench_probes_test.go create mode 100644 internal/core/bench_sort2_test.go create mode 100644 internal/core/bench_test.go create mode 100644 internal/core/bessel_climb_pin_test.go create mode 100644 internal/core/bessel_realorder_accuracy_test.go create mode 100644 internal/core/besseljreal_test.go create mode 100644 internal/core/besselmod.go create mode 100644 internal/core/besselmod_test.go create mode 100644 internal/core/broadcast.go create mode 100644 internal/core/broadcast_test.go create mode 100644 internal/core/canonical_partition_pins_test.go create mode 100644 internal/core/centralsum_kernels_test.go create mode 100644 internal/core/chunk_precision_pins_test.go create mode 100644 internal/core/complex_test.go create mode 100644 internal/core/concat.go create mode 100644 internal/core/concat_test.go create mode 100644 internal/core/cosm1_test.go create mode 100644 internal/core/cumprod_half_test.go create mode 100644 internal/core/data_test.go create mode 100644 internal/core/defect_family_pins_test.go create mode 100644 internal/core/diff.go create mode 100644 internal/core/dtypes_promote_test.go create mode 100644 internal/core/einsum.go create mode 100644 internal/core/einsum_bench_test.go create mode 100644 internal/core/einsum_foldorder_test.go create mode 100644 internal/core/einsum_norm_special_pins_test.go create mode 100644 internal/core/einsum_special_pins_test.go create mode 100644 internal/core/einsum_test.go create mode 100644 internal/core/elementwise_view_test.go create mode 100644 internal/core/elliptic.go create mode 100644 internal/core/elliptic2_test.go create mode 100644 internal/core/elliptic_test.go create mode 100644 internal/core/expint.go create mode 100644 internal/core/expint_test.go create mode 100644 internal/core/extrema.go create mode 100644 internal/core/extrema_test.go create mode 100644 internal/core/float16.go create mode 100644 internal/core/float16_test.go create mode 100644 internal/core/float32_test.go create mode 100644 internal/core/fresnel.go create mode 100644 internal/core/grid.go create mode 100644 internal/core/guard_pins_test.go create mode 100644 internal/core/index.go create mode 100644 internal/core/index_test.go create mode 100644 internal/core/indexing2.go create mode 100644 internal/core/init.go create mode 100644 internal/core/int_exactness_pins_test.go create mode 100644 internal/core/interpolate2d.go create mode 100644 internal/core/interpolate2d_test.go create mode 100644 internal/core/interpolate_knot_test.go create mode 100644 internal/core/jacobian.go create mode 100644 internal/core/jacobian_test.go create mode 100644 internal/core/knot_ellipsis_pins_test.go create mode 100644 internal/core/mask.go create mode 100644 internal/core/mask_test.go create mode 100644 internal/core/mat.go create mode 100644 internal/core/mat_kernel_bench_test.go create mode 100644 internal/core/mat_kernel_test.go create mode 100644 internal/core/mat_tile_bench_test.go create mode 100644 internal/core/math_bench_extra_test.go create mode 100644 internal/core/mathfunc.go create mode 100644 internal/core/mathfunc_test.go create mode 100644 internal/core/matrix2.go create mode 100644 internal/core/misc_bench_test.go create mode 100644 internal/core/narrow_compare_alloc_test.go create mode 100644 internal/core/narrow_dispatch_pin_test.go create mode 100644 internal/core/narrowdt.go create mode 100644 internal/core/onehot.go create mode 100644 internal/core/onehot_test.go create mode 100644 internal/core/ops.go create mode 100644 internal/core/ops_test.go create mode 100644 internal/core/orthopoly.go create mode 100644 internal/core/orthopoly_test.go create mode 100644 internal/core/overflow_refusal_pin_test.go create mode 100644 internal/core/pchip.go create mode 100644 internal/core/quasirandom.go create mode 100644 internal/core/quasirandom_test.go create mode 100644 internal/core/random.go create mode 100644 internal/core/random_test.go create mode 100644 internal/core/randomsub_test.go create mode 100644 internal/core/reduce.go create mode 100644 internal/core/reduce_accuracy_test.go create mode 100644 internal/core/reduce_test.go create mode 100644 internal/core/reduction2.go create mode 100644 internal/core/reduction_route_test.go create mode 100644 internal/core/reduction_shape_consistency_test.go create mode 100644 internal/core/reference_walks_test.go create mode 100644 internal/core/reshape.go create mode 100644 internal/core/reshape_pad_test.go create mode 100644 internal/core/runtime.go create mode 100644 internal/core/runtime_test.go create mode 100644 internal/core/scalar_clamp_pow_bench_test.go create mode 100644 internal/core/scan_compensation_test.go create mode 100644 internal/core/shape_overflow_pin_test.go create mode 100644 internal/core/smallops_test.go create mode 100644 internal/core/sort.go create mode 100644 internal/core/sort_test.go create mode 100644 internal/core/sparse.go create mode 100644 internal/core/special.go create mode 100644 internal/core/special2.go create mode 100644 internal/core/special2_test.go create mode 100644 internal/core/special_edge_accuracy_test.go create mode 100644 internal/core/special_test.go create mode 100644 internal/core/tensor.go create mode 100644 internal/core/tensor_bench_test.go create mode 100644 internal/core/tensor_test.go create mode 100644 internal/core/validation_pins_test.go create mode 100644 internal/core/view_payload_reads_test.go create mode 100644 internal/core/view_test.go create mode 100644 internal/core/walks_pin_test.go create mode 100644 internal/core/where_half_test.go create mode 100644 internal/engine/engine.go create mode 100644 internal/engine/engine_bench_test.go create mode 100644 internal/engine/engine_test.go create mode 100644 io/bench_test.go create mode 100644 io/csv.go create mode 100644 io/csv_bench_test.go create mode 100644 io/csv_parity_test.go create mode 100644 io/csv_test.go create mode 100644 io/doc.go create mode 100644 io/example_test.go create mode 100644 io/extent_wrap_pins_test.go create mode 100644 io/fits.go create mode 100644 io/fits_hostile_pins_test.go create mode 100644 io/fits_table_pins_test.go create mode 100644 io/fits_test.go create mode 100644 io/fitstable.go create mode 100644 io/fitstable_keyword_pin_test.go create mode 100644 io/fitstable_test.go create mode 100644 io/fuzz_test.go create mode 100644 io/hdf5.go create mode 100644 io/hdf5_bool_fixture_test.go create mode 100644 io/hdf5_continuation_pins_test.go create mode 100644 io/hdf5_structure_pins_test.go create mode 100644 io/hdf5_test.go create mode 100644 io/hdf5write.go create mode 100644 io/hdf5write_estimate_test.go create mode 100644 io/hdf5write_group.go create mode 100644 io/hdf5write_object.go create mode 100644 io/hdf5write_test.go create mode 100644 io/helpers_test.go create mode 100644 io/hostile_test.go create mode 100644 io/mmap.go create mode 100644 io/mmap_test.go create mode 100644 io/netcdf.go create mode 100644 io/netcdf_payload_test.go create mode 100644 io/netcdf_test.go create mode 100644 io/netcdfrecord_test.go create mode 100644 io/testdata/fuzz/FuzzLoadHDF5/2fe5cd680e6d8b63 create mode 100644 io/testdata/fuzz/FuzzLoadHDF5/a7ab50e52191807a create mode 100644 io/testdata/fuzz/FuzzLoadHDF5/ae4657f5205db80f create mode 100644 io/testdata/fuzz/FuzzLoadNetCDF/14c890c96cda7d9c create mode 100644 io/testdata/fuzz/FuzzLoadNetCDF/c4c632324ec4a450 create mode 100644 io/testdata/h5/bool_enum.h5 create mode 100644 io/testdata/h5/fixture.h5 create mode 100644 io/testdata/h5/fletcher.h5 create mode 100644 io/testdata/h5/latest.h5 create mode 100644 io/walk_cycle_pins_test.go create mode 100644 io/wave_bench_test.go create mode 100644 io/zz_guard_test.go create mode 100644 justfile create mode 100644 leak_test.go create mode 100644 linalg/bench_decomp_test.go create mode 100644 linalg/bench_factor_test.go create mode 100644 linalg/bench_iterative_test.go create mode 100644 linalg/bench_sparse_test.go create mode 100644 linalg/cholupdate.go create mode 100644 linalg/cholupdate_test.go create mode 100644 linalg/complex_refusal_pins_test.go create mode 100644 linalg/csvd.go create mode 100644 linalg/csvd_test.go create mode 100644 linalg/decomp.go create mode 100644 linalg/decomp2.go create mode 100644 linalg/decomp3.go create mode 100644 linalg/decomp3_test.go create mode 100644 linalg/decomp_test.go create mode 100644 linalg/decompositions_test.go create mode 100644 linalg/det_factor_dispatch_test.go create mode 100644 linalg/detfloat32_test.go create mode 100644 linalg/doc.go create mode 100644 linalg/dtype_shape_refusal_pins_test.go create mode 100644 linalg/dtypes_census_test.go create mode 100644 linalg/eigenreal.go create mode 100644 linalg/eigenreal_test.go create mode 100644 linalg/example_test.go create mode 100644 linalg/expm.go create mode 100644 linalg/expm_test.go create mode 100644 linalg/extreme_scale_pins_test.go create mode 100644 linalg/fitpoly_test.go create mode 100644 linalg/geneigen.go create mode 100644 linalg/geneigen_test.go create mode 100644 linalg/gmres.go create mode 100644 linalg/gmres_test.go create mode 100644 linalg/helpers.go create mode 100644 linalg/helpers_test.go create mode 100644 linalg/initialisers_sparse_test.go create mode 100644 linalg/least_squares_test.go create mode 100644 linalg/linalg.go create mode 100644 linalg/linalg_complex_test.go create mode 100644 linalg/linalg_test.go create mode 100644 linalg/mat_scratch_test.go create mode 100644 linalg/mat_test.go create mode 100644 linalg/matrixfunc.go create mode 100644 linalg/matrixfunc_test.go create mode 100644 linalg/nonfinite_refusal_pins_test.go create mode 100644 linalg/perf_bench_test.go create mode 100644 linalg/pipeline.go create mode 100644 linalg/pipeline_test.go create mode 100644 linalg/polyroots.go create mode 100644 linalg/polyroots_test.go create mode 100644 linalg/rrqr.go create mode 100644 linalg/rrqr_test.go create mode 100644 linalg/schur.go create mode 100644 linalg/schur_test.go create mode 100644 linalg/solver_dtype_pins_test.go create mode 100644 linalg/solver_guard_pins_test.go create mode 100644 linalg/sparse_guard_pins_test.go create mode 100644 linalg/sparsebicg.go create mode 100644 linalg/sparsebicg_test.go create mode 100644 linalg/sparsecholesky.go create mode 100644 linalg/sparsecholesky_test.go create mode 100644 linalg/sparsecholupdate.go create mode 100644 linalg/sparsecholupdate_test.go create mode 100644 linalg/sparsecomplex.go create mode 100644 linalg/sparsecomplex_test.go create mode 100644 linalg/sparsecsc.go create mode 100644 linalg/sparsecsc_split_test.go create mode 100644 linalg/sparsecsr.go create mode 100644 linalg/sparsefill.go create mode 100644 linalg/sparsefill_bench_test.go create mode 100644 linalg/sparsefill_test.go create mode 100644 linalg/sparsegeneral.go create mode 100644 linalg/sparsegeneral_test.go create mode 100644 linalg/sparseigen.go create mode 100644 linalg/sparseigen_test.go create mode 100644 linalg/sparseilu.go create mode 100644 linalg/sparseilu_test.go create mode 100644 linalg/sparselsqr.go create mode 100644 linalg/sparselsqr_test.go create mode 100644 linalg/sparselu.go create mode 100644 linalg/sparselu_test.go create mode 100644 linalg/sparseops.go create mode 100644 linalg/sparseops_test.go create mode 100644 linalg/sparsesolve.go create mode 100644 linalg/sparsesolve_test.go create mode 100644 linalg/sparsexp.go create mode 100644 linalg/sparsexp_test.go create mode 100644 linalg/spline.go create mode 100644 linalg/spline_test.go create mode 100644 linalg/svdsolve.go create mode 100644 linalg/svdsolve_test.go create mode 100644 linalg/tridiag.go create mode 100644 linalg/tridiag_test.go create mode 100644 optim/bench_leastsqfit_test.go create mode 100644 optim/bench_solvers_test.go create mode 100644 optim/bench_test.go create mode 100644 optim/bounded_test.go create mode 100644 optim/brent.go create mode 100644 optim/brent_test.go create mode 100644 optim/broyden.go create mode 100644 optim/broyden_test.go create mode 100644 optim/budget_exit_pin_test.go create mode 100644 optim/callback_hygiene_pins_test.go create mode 100644 optim/callback_pins_test.go create mode 100644 optim/devolution.go create mode 100644 optim/devolution_test.go create mode 100644 optim/doc.go create mode 100644 optim/dtypes_census_test.go create mode 100644 optim/example_test.go create mode 100644 optim/fallback_guard_pins_test.go create mode 100644 optim/global2.go create mode 100644 optim/global2_test.go create mode 100644 optim/helpers_test.go create mode 100644 optim/jacobianparallel_test.go create mode 100644 optim/lbfgs.go create mode 100644 optim/lbfgs_test.go create mode 100644 optim/leastsq.go create mode 100644 optim/leastsq_mirror_pin_test.go create mode 100644 optim/leastsq_test.go create mode 100644 optim/leastsqfit_test.go create mode 100644 optim/linear.go create mode 100644 optim/linear_test.go create mode 100644 optim/linesearch_pins_test.go create mode 100644 optim/nonfinite_exit_pins_test.go create mode 100644 optim/nonlinearconstr.go create mode 100644 optim/nonlinearconstr_test.go create mode 100644 optim/optimise.go create mode 100644 optim/optimise_test.go create mode 100644 optim/qp.go create mode 100644 optim/qp_test.go create mode 100644 optim/rootsystem.go create mode 100644 optim/rootsystem_test.go create mode 100644 optim/rowblocks_pin_test.go create mode 100644 optim/simplex.go create mode 100644 optim/simplex_test.go create mode 100644 optim/solvercoverage_test.go create mode 100644 oracle_norace_test.go create mode 100644 oracle_race_test.go create mode 100644 oracle_test.go create mode 100644 plot/dtypes_census_test.go create mode 100644 plot/plot.go create mode 100644 plot/plot_test.go create mode 100644 signal/arma.go create mode 100644 signal/arma_test.go create mode 100644 signal/bench_perf_test.go create mode 100644 signal/bench_test.go create mode 100644 signal/chirp.go create mode 100644 signal/chirp_test.go create mode 100644 signal/complex_refusal_pins_test.go create mode 100644 signal/conv.go create mode 100644 signal/conv_bench_extra_test.go create mode 100644 signal/conv_bench_test.go create mode 100644 signal/conv_fused_pin_test.go create mode 100644 signal/conv_pooling_test.go create mode 100644 signal/conv_transpose_test.go create mode 100644 signal/conv_window_pins_test.go create mode 100644 signal/correlate.go create mode 100644 signal/correlate_test.go create mode 100644 signal/dct.go create mode 100644 signal/dct_test.go create mode 100644 signal/doc.go create mode 100644 signal/dtypes_census_test.go create mode 100644 signal/entry_gate_pins_test.go create mode 100644 signal/example_test.go create mode 100644 signal/examples_test.go create mode 100644 signal/fft.go create mode 100644 signal/fft_nd_test.go create mode 100644 signal/fft_precision_test.go create mode 100644 signal/fft_test.go create mode 100644 signal/filter.go create mode 100644 signal/filter_dtype_pin_test.go create mode 100644 signal/filter_test.go create mode 100644 signal/filterdesign.go create mode 100644 signal/filterdesign_test.go create mode 100644 signal/filtfilt.go create mode 100644 signal/filtfilt_test.go create mode 100644 signal/float32_test.go create mode 100644 signal/helpers_test.go create mode 100644 signal/hilbert.go create mode 100644 signal/hilbert_test.go create mode 100644 signal/kalman.go create mode 100644 signal/kalman_test.go create mode 100644 signal/misc_bench_test.go create mode 100644 signal/nan_propagation_pins_test.go create mode 100644 signal/nufft.go create mode 100644 signal/nufft_test.go create mode 100644 signal/periodogram.go create mode 100644 signal/periodogram_test.go create mode 100644 signal/poisson.go create mode 100644 signal/poisson_complex_pins_test.go create mode 100644 signal/poisson_test.go create mode 100644 signal/poisson_view_pin_test.go create mode 100644 signal/poissondirichlet.go create mode 100644 signal/poissondirichlet_test.go create mode 100644 signal/pool.go create mode 100644 signal/pooling2d_test.go create mode 100644 signal/rankfilter.go create mode 100644 signal/rankfilter_test.go create mode 100644 signal/resample.go create mode 100644 signal/resample_test.go create mode 100644 signal/stencil.go create mode 100644 signal/stencil_test.go create mode 100644 signal/stft.go create mode 100644 signal/stft_test.go create mode 100644 signal/transform_cache_pins_test.go create mode 100644 signal/wave_bench_test.go create mode 100644 signal/wavelet.go create mode 100644 signal/wavelet_test.go create mode 100644 signal/waveletdwt.go create mode 100644 signal/waveletdwt_test.go create mode 100644 signal/welch.go create mode 100644 signal/windows.go create mode 100644 signal/windows_test.go create mode 100644 spmd/arg.go create mode 100644 spmd/arg_test.go create mode 100644 spmd/bench_tcp_test.go create mode 100644 spmd/bench_test.go create mode 100644 spmd/collective.go create mode 100644 spmd/collective_test.go create mode 100644 spmd/doc.go create mode 100644 spmd/framepool.go create mode 100644 spmd/framepool_test.go create mode 100644 spmd/halo.go create mode 100644 spmd/halo_grid_test.go create mode 100644 spmd/halo_test.go create mode 100644 spmd/partition.go create mode 100644 spmd/partition_test.go create mode 100644 spmd/proc_test.go create mode 100644 spmd/reduce.go create mode 100644 spmd/reduce_test.go create mode 100644 spmd/transport.go create mode 100644 spmd/wire.go create mode 100644 spmd/wire_test.go create mode 100644 spmd/world.go create mode 100644 spmd/world_test.go create mode 100644 stats/anova.go create mode 100644 stats/anova_test.go create mode 100644 stats/bench_kernels_test.go create mode 100644 stats/bench_quantilereg_test.go create mode 100644 stats/bench_regression_test.go create mode 100644 stats/bench_test.go create mode 100644 stats/cdf.go create mode 100644 stats/cdf_test.go create mode 100644 stats/cluster.go create mode 100644 stats/cluster_test.go create mode 100644 stats/contingency.go create mode 100644 stats/contingency_test.go create mode 100644 stats/deep_tail_pin_test.go create mode 100644 stats/degenerate_input_pin_test.go create mode 100644 stats/distrib.go create mode 100644 stats/distrib2.go create mode 100644 stats/distrib2_test.go create mode 100644 stats/distrib_test.go create mode 100644 stats/doc.go create mode 100644 stats/dtypes_census_test.go create mode 100644 stats/estimation_guards_pin_test.go create mode 100644 stats/example_test.go create mode 100644 stats/glm.go create mode 100644 stats/glmkdemvn_test.go create mode 100644 stats/gmm_posterior_reference_test.go create mode 100644 stats/gp.go create mode 100644 stats/gp_test.go create mode 100644 stats/helpers_test.go create mode 100644 stats/hierarchy.go create mode 100644 stats/hierarchy_test.go create mode 100644 stats/histogram2d.go create mode 100644 stats/hmm.go create mode 100644 stats/hmm_test.go create mode 100644 stats/inference.go create mode 100644 stats/inference_test.go create mode 100644 stats/kde.go create mode 100644 stats/lasso.go create mode 100644 stats/lasso_test.go create mode 100644 stats/mixed.go create mode 100644 stats/mixed_test.go create mode 100644 stats/multipletest.go create mode 100644 stats/multipletest_test.go create mode 100644 stats/mvn.go create mode 100644 stats/negbinom_far_pin_test.go create mode 100644 stats/noncentral.go create mode 100644 stats/noncentral_test.go create mode 100644 stats/pca.go create mode 100644 stats/pca_test.go create mode 100644 stats/pins_test.go create mode 100644 stats/poissonregression_test.go create mode 100644 stats/pool_reuse_pin_test.go create mode 100644 stats/quantile_newton_test.go create mode 100644 stats/quantile_test.go create mode 100644 stats/quantilereg.go create mode 100644 stats/quantilereg_test.go create mode 100644 stats/rankcorr.go create mode 100644 stats/rankcorr_test.go create mode 100644 stats/regression.go create mode 100644 stats/regression_edge_pin_test.go create mode 100644 stats/regression_test.go create mode 100644 stats/robust.go create mode 100644 stats/robustregression.go create mode 100644 stats/robustregression_test.go create mode 100644 stats/rolling.go create mode 100644 stats/rolling_prec_test.go create mode 100644 stats/smallops_test.go create mode 100644 stats/stats.go create mode 100644 stats/stats_random_test.go create mode 100644 stats/stats_test.go create mode 100644 stats/tail_and_extreme_pin_test.go create mode 100644 stats/variance_prec_test.go create mode 100644 stats/view_tail_pin_test.go create mode 100644 stats/zz_guard_test.go diff --git a/.gitea/workflows/race.yml b/.gitea/workflows/race.yml new file mode 100644 index 0000000..0e2ac0e --- /dev/null +++ b/.gitea/workflows/race.yml @@ -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/... diff --git a/.gitea/workflows/release.yml b/.gitea/workflows/release.yml new file mode 100644 index 0000000..99c23e6 --- /dev/null +++ b/.gitea/workflows/release.yml @@ -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}; + ' diff --git a/.gitea/workflows/test.yml b/.gitea/workflows/test.yml new file mode 100644 index 0000000..7034e44 --- /dev/null +++ b/.gitea/workflows/test.yml @@ -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 . diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..8f720fc --- /dev/null +++ b/.gitignore @@ -0,0 +1,5 @@ +.idea/ +.zcode/ +bin/ +coverage.out +*.test diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..a110b54 --- /dev/null +++ b/CHANGELOG.md @@ -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. diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md new file mode 100644 index 0000000..822cad0 --- /dev/null +++ b/CONTRIBUTING.md @@ -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 + 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. diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..f83dd2a --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 Petr Balvín (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. diff --git a/README.md b/README.md new file mode 100644 index 0000000..c6abd81 --- /dev/null +++ b/README.md @@ -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) diff --git a/SECURITY.md b/SECURITY.md new file mode 100644 index 0000000..ea19fdc --- /dev/null +++ b/SECURITY.md @@ -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. diff --git a/doc.go b/doc.go new file mode 100644 index 0000000..87c097e --- /dev/null +++ b/doc.go @@ -0,0 +1,58 @@ +// Copyright (c) 2026 Petr Balvín (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 diff --git a/docs/API.md b/docs/API.md new file mode 100644 index 0000000..2259fb0 --- /dev/null +++ b/docs/API.md @@ -0,0 +1,3583 @@ +# API + +The reference for the exported surface of Tensor, one section per +package. Each section lists every exported symbol in the groups the +package is read by, documents the options and result structs field by +field with their defaults, names the conditions that come back as +errors, and ends with the package's flagship call sequence. + +The signatures live in the source, where `go doc` is the authority on +them and the godoc comments explain each one. This document is the +map: what each part of the surface is for, how the parts fit together, +and which options a call takes. Every package carries runnable +`Example` functions beside its tests, compiled and executed by the +suite, so the documented call sequences cannot rot; the longer +workflows are the programs in [`examples/`](../examples). + +Conventions every package shares: functions return `(value, error)` +and wrap errors with context, every error message carries the +`tensor: ` prefix, operations never modify their inputs, and scalars +come back as `Scalar` when the dtype follows the input. + +## Library + +### Package `tensor` (root) + +The facade and the array core. `import "sourcedock.dev/petrbalvin/tensor"` is the +one import the library needs: the domain packages are re-exported here under +their own names, so `tensor.Solve` and `linalg.Solve` are one function, and the +shared array core, which has no domain package of its own, is exported here +outright. What the documented surface of a call looks like in `go doc +sourcedock.dev/petrbalvin/tensor` is therefore the whole library in one place. + +The value everything is built on is `Array`: an immutable, shape-checked +n-dimensional array in which every operation returns a new array and writes +neither its receiver nor its arguments. Elements are `int64`, IEEE 754 binary16 +bits, `float32`, `float64` or `complex128`, any rank is allowed, and storage is +row-major. Mixing dtypes promotes along one ladder and shapes never broadcast +silently, so a mismatch is an error naming both shapes rather than a wrong +number. + +#### Importing + +```go +import ( + tensor "sourcedock.dev/petrbalvin/tensor" +) +``` + +`tensor.Array` is a type alias for the core array, not a wrapper: `type Array = +core.Array` in `facade_generated.go`, and the same holds for `Dtype`, `Scalar`, +`Generator`, `SparseCOO` and `JacobianOptions`. Everything the array core +offers is exported here directly, so the root package is the only import a core +user needs. The domain packages are importable on their own +(`sourcedock.dev/petrbalvin/tensor/linalg` and its siblings) for a caller that +wants one of them without the rest, but their entry points take the core array +as `*core.Array`, so a caller who imports one of them alone needs this package +beside it to build the argument: the general constructors are rooted here. + +`linalg.ArrayFromFloatsSafe` is the one constructor that exists outside this +package, a float vector of a given length built from a slice the caller keeps, +and it is there for exactly that caller. + +#### Dtypes + +| Name | What it is | +|---|---| +| `Bool` | the boolean element type: the comparisons' answer, and logical vectors for masking, `Where` and masked reads. It carries no arithmetic of its own; the logic entries `And`, `Or`, `Xor` and `Not` are its surface. | +| `Int8` | the `int8` element type. | +| `Uint8` | the `uint8` element type: byte payloads for images, masks and the byte classes of the file formats. | +| `Int16` | the `int16` element type. | +| `Uint16` | the `uint16` element type. | +| `Int32` | the `int32` element type. | +| `Uint32` | the `uint32` element type. | +| `Int` | the `int64` element type. | +| `Float16` | the IEEE 754 binary16 element type: the payload holds `uint16` bit patterns, widened exactly on every read. | +| `Float32` | the `float32` element type. | +| `Float` | the `float64` element type, the default real dtype. | +| `Complex` | the `complex128` element type. | +| `Dtype` | the element type of an `Array`; `Dtype.String()` renders it as it appears in diagnostics ("bool", "int8", "uint8", "int16", "uint16", "int32", "uint32", "int", "float16", "float32", "float", "complex"). | +| `BoolAt(a, index...)` | reads any element as a bool: the boolean payload's own value, and every other dtype read against zero. | +| `HalfFromFloat64(f)` | narrows a float64 to the half bit pattern with round-to-nearest-even, gradual underflow through the subnormals, overflow to the signed infinity at `\|f\| >= 65520`, NaN canonicalised, signed zero preserved. | +| `HalfToFloat64(h)` | widens a half bit pattern back to float64 exactly, with no rounding and no range loss. | + +The ladder runs the integer class first, then the floats, then complex: +`Bool` below the small integers below `Int` below `Float16` below +`Float32` below `Float` below `Complex`. Mixed signedness integer pairs +promote to the smallest dtype whose value range contains both operands, +so `Int8` with `Uint8` answers `Int16`, `Int16` with `Uint16` answers +`Int32` and `Int32` with `Uint32` answers `Int`; across classes the +higher ladder position decides, the rule the int-to-float-to-complex +promotions have always walked. The constants' numeric order is +deliberately not the ladder's order: the seven new dtypes append after +`Float16`, because the recorded oracle digests pin the ordinals, and +the promotion tables are explicit rather than those numbers. The +element-wise kernels compute the narrow integers at their own width +with the machine's wrap-around semantics; reductions widen exactly, so +`Sum` of a `Bool` array counts its trues into an int scalar; `Astype` +range-checks every conversion into a narrow target with a loud error +naming the value and its index, while the conversions among the +original five dtypes keep their historical cast semantics. Surfaces a +later round will carry refuse the new dtypes with an error naming +`Astype`: the matrix products and `Einsum`, the sparse constructors +and products, `Kron`, `Sort` and `ArgSort`, `Prod`, `CumSum`, +`CumProd`, the `Norm` family, the clipping entries, `Sqrt` and the +transcendental families (`Abs` answers nil, its signature having no +error channel). `Interpolate2D` refuses a narrow grid the same way. +`Bytes` and `FromBytes` keep their `Int`-only contract for this round. + +#### Constructors + +| Call | What it does | +|---|---| +| `FromInts(values, shape...)` | builds an int array of the given shape from values, copying them. | +| `FromFloats(values, shape...)` | builds a float array of the given shape from values, copying them. | +| `FromFloat16s(values, shape...)` | builds a float16 array from float64 values, copying them and narrowing each to the nearest half. | +| `FromFloat32s(values, shape...)` | builds a float32 array of the given shape from values, copying them. | +| `FromComplexes(values, shape...)` | builds a complex array of the given shape from values, copying them. | +| `FromBools(values, shape...)` | builds a bool array from values, copying them. | +| `FromInt8s(values, shape...)` | builds an int8 array from values, copying them. | +| `FromUint8s(values, shape...)` | builds a uint8 array from values, copying them. | +| `FromInt16s(values, shape...)` | builds an int16 array from values, copying them. | +| `FromUint16s(values, shape...)` | builds a uint16 array from values, copying them. | +| `FromInt32s(values, shape...)` | builds an int32 array from values, copying them. | +| `FromUint32s(values, shape...)` | builds a uint32 array from values, copying them. | +| `FloatsFromArray(values, shape...)` | builds a float array that takes ownership of values: no copy, and the caller must not touch the slice afterwards. | +| `IntsFromArray(values, shape...)` | the same ownership contract for an int array. | +| `BoolsFromArray(values, shape...)` | the same ownership contract for a bool array. | +| `Int8sFromArray(values, shape...)` | the same ownership contract for an int8 array. | +| `Uint8sFromArray(values, shape...)` | the same for a uint8 array. | +| `Int16sFromArray(values, shape...)` | the same for an int16 array. | +| `Uint16sFromArray(values, shape...)` | the same for a uint16 array. | +| `Int32sFromArray(values, shape...)` | the same for an int32 array. | +| `Uint32sFromArray(values, shape...)` | the same for a uint32 array. | +| `HalvesFromArray(values, shape...)` | the same for a float16 array, taking raw half bit patterns with no conversion. | +| `ComplexFromArray(values, shape...)` | the same for a complex array. | +| `FromFloatSlice(values, shape...)` | aliases a `[]float64` into a float array with no copy at all: reads reflect later writes, so it fits build-then-consume windows and nothing else. | +| `FromFloat32Slice(values, shape...)` | the float32 twin, the same contract. | +| `Zeros(dt, shape...)` | builds an array of the given dtype and shape filled with zeros. | +| `Ones(dt, shape...)` | builds an array of the given dtype and shape filled with ones. | +| `FullI(v, shape...)` | fills an int array with an int. | +| `FullF(v, shape...)` | fills a float array with a float64. | +| `FullF16(v, shape...)` | fills a float16 array with a float64, narrowed to the nearest half. | +| `FullF32s(v, shape...)` | fills a float32 array with a float32. | +| `FullC(v, shape...)` | fills a complex array with a complex128. | +| `Range(start, stop)` | builds the int array start, start+1, …, stop-1; start at or above stop yields an empty array. | +| `RangeBy(start, stop, step)` | builds the int array start, start+step, …, staying below stop for a positive step and above it for a negative one. | +| `Linspace(start, stop, n)` | returns n evenly spaced float values from start to stop, both inclusive. | +| `Grid(x, y)` | builds the two coordinate matrices of the meshgrid pattern from two 1-D vectors, x varying along the columns and y along the rows. | +| `New(dt, shape...)` | allocates a zeroed array with no error return, for internally derived shapes: an invalid argument yields a nil array, not a panic. | +| `FromBytes(b)` | wraps raw bytes as an int array, the inverse of `Array.Bytes`. | +| `Identity(dt, n)` | builds the n×n identity matrix of the given dtype. | +| `Copy(a)` | returns a fresh array with the same data, shape and dtype. | +| `ZerosLike(a)` | returns a zero-filled array of the same shape and dtype, as a copy. | +| `OnesLike(a)` | the same filled with ones. | + +#### The array value + +| Call | What it does | +|---|---| +| `Array.Shape()` | returns a copy of the dimensions. | +| `Array.NDim()` | returns the number of dimensions. | +| `Array.Len()` | returns the number of elements, the product of the shape, which for a view is shorter than the storage it aliases. | +| `Array.Dtype()` | returns the element type. | +| `Array.String()` | renders the dtype, the shape and the values, as in `int (2, 2) [1, 2, 3, 4]`; a debugging aid, not a format. | +| `Array.Strided()` | reports whether the array carries a non-trivial stride layout; a dense array answers false. | +| `Array.CopyRows(indices)` | returns a new array holding the rows of a 2+-D array at the given leading-dimension indices, preserving every other dimension and the dtype. | +| `Equal(a, b)` | reports whether two arrays have the same dtype, shape and values; the dtype is part of the identity, so an int 1 does not equal a float 1.0, and arrays holding NaN are never equal. | +| `Astype(a, dt)` | returns a copy converted to dt: int and float convert like Go's casts, complex to float keeps the real part, complex to int or float16 is an error, real to complex adds a zero imaginary part, and a conversion to the array's own dtype still copies. | +| `Array.Bytes()` | returns the payload interpreted as raw bytes; only int arrays carry one, other dtypes return nil. | + +#### Element-wise operations + +| Call | What it does | +|---|---| +| `Add(a, b)` | the element-wise sum of two arrays of the same shape. | +| `Sub(a, b)` | the element-wise difference. | +| `Mul(a, b)` | the element-wise product. | +| `Div(a, b)` | the element-wise true division; the result is float for real operands and complex when a complex operand takes part, and division by zero follows IEEE 754 rather than raising an error. | +| `Quo(a, b)` | the element-wise integer division of two int arrays; a float or complex operand, or a zero divisor element, is an error. | +| `QuoI(a, v)` | the same against an int scalar, with a zero scalar refused. | +| `AddI(a, v)` / `SubI(a, v)` / `MulI(a, v)` / `DivI(a, v)` | the int scalar variants. Each dtype keeps its kind, except `DivI`, which is true division and therefore lifts an int array to float. | +| `AddF(a, v)` / `SubF(a, v)` / `MulF(a, v)` / `DivF(a, v)` | the float scalar variants: an int array becomes float64, a floating array keeps its dtype, a complex array stays complex. | +| `AddC(a, v)` / `SubC(a, v)` / `MulC(a, v)` / `DivC(a, v)` | the complex scalar variants; the result is always complex. All of these return the array alone, with no error, because a scalar cannot make a shape disagree. | +| `Pow(a, b)` | the element-wise power: int with int stays int with wrapping on overflow, a negative exponent on int operands is an error, and any float or complex operand promotes per the ladder. | +| `PowI(a, n)` | raises each element to an int exponent; complex arrays use exact repeated squaring and a negative exponent takes the reciprocal. | +| `Maximum(a, b)` | the element-wise larger of two same-shape arrays; NaN propagates. | +| `Minimum(a, b)` | the element-wise smaller; NaN propagates. | +| `Where(cond, x, y)` | selects element-wise between x and y: a set element of cond picks x, an unset one picks y, all three shapes must agree, and the result promotes like the arithmetic. cond is a bool mask or an int array, whose non-zero elements stand for true. | +| `Select(a, mask)` | a 1-D copy of the elements where the mask is set: a bool mask selects where it reads true, an int mask where it reads non-zero. | +| `ClipF(a, lo, hi)` | clamps a real array into `[lo, hi]`; an int array becomes float, a float16 or float32 array keeps its dtype, and `lo > hi` is an error. | +| `ClipI(a, lo, hi)` | the same clamp with int bounds, the dtype keeping its kind. | +| `Abs(a)` | the element-wise absolute value; a complex array yields its float magnitudes. | +| `Sign(a)` | the sign of every element as a float array; a NaN element reports 0. | +| `Sqrt(a)` | the square root of each element; negatives yield NaN. | +| `Exp(a)` | e raised to each element. | +| `Log(a)` / `Log2(a)` / `Log10(a)` | the natural, base-2 and base-10 logarithms; negatives yield NaN. | +| `Sin(a)` / `Cos(a)` / `Tan(a)` / `Tanh(a)` | the trigonometric and hyperbolic elements of the library, in radians. | +| `Sigmoid(a)` | the logistic function `1/(1+e^-x)` of each element. | +| `Sinc(a)` | the normalised `sin(πx)/(πx)` with `Sinc(0) = 1`. | +| `Cosm1(a)` | `cos(x) − 1` kept accurate for small arguments, where the direct float64 subtraction has no correct significant bit under about `\|x\| ≈ 1e-8`; the factored series holds full precision down to the smallest representable argument, where the answer is exactly `−x²/2` as far as the format can see it. | +| `Ceil(a)` / `Floor(a)` / `Round(a)` / `Trunc(a)` | the rounding family: up, down, half away from zero, and toward zero. An int array is an identity copy in all four. | + +#### Comparisons and masks + +| Call | What it does | +|---|---| +| `Eq(a, b)` / `Ne(a, b)` | a bool mask, true where the operands are equal, respectively different, element-wise. | +| `Lt(a, b)` / `Le(a, b)` / `Gt(a, b)` / `Ge(a, b)` | a bool mask, true where a is below, at most, above, at least b. | +| `EqI` / `NeI` / `LtI` / `LeI` / `GtI` / `GeI` | each comparison against an int scalar, answered as a bool mask; the names follow the same order as the pairs above. | +| `EqF` / `NeF` / `LtF` / `LeF` / `GtF` / `GeF` | each comparison against a float scalar, answered as a bool mask. | +| `IsNaN(a)` | an int mask marking NaN elements; complex arrays are not supported. | +| `IsInf(a)` | an int mask marking positive and negative infinity. | +| `IsFinite(a)` | an int mask marking values that are neither NaN nor infinite. | + +The comparisons answer bool masks, one byte per element, the predicate +itself. `Select` and `Where` take a bool mask, and they still take the +int mask of 0 and 1 `IsNaN`, `IsInf` and `IsFinite` answer, so a mask of +either dtype feeds both directly and `And`, `Or`, `Xor` and `Not` +compose the bool ones. + +`Sum`, `Dot`, `Min` and `Max` answer a `Scalar`, which boxes one numeric result +whose dtype is known only at runtime. + +| Call | What it does | +|---|---| +| `Scalar.Int()` | the value as an int64, converting a float or complex scalar by truncation. | +| `Scalar.Float()` | the value as a float64; a complex scalar contributes its real part. | +| `Scalar.Complex()` | the value as a complex128. | +| `Scalar.IsFloat()` | whether the scalar came from a float computation. | +| `Scalar.IsComplex()` | whether the scalar came from a complex computation. | +| `Scalar.String()` | renders the scalar with its dtype, as in `int 7` or `complex (4-2i)`. | + +Use `IsFloat` and `IsComplex` where the distinction between an integer result +and a converted one matters, because `Int` alone truncates silently. + +#### Reductions + +| Call | What it does | +|---|---| +| `Sum(a)` | the sum of all elements as a `Scalar`, which is int for an int array, float when a float16 or float32 array accumulated in float64, and complex for a complex array; an empty array sums to zero. | +| `SumAxis(a, dim)` | the sums along the given dimension. | +| `Mean(a)` | the arithmetic mean as a float64, computed in float64 even for an int array; an empty or complex array is an error. | +| `MeanAxis(a, dim)` | the float means along the given dimension. | +| `Prod(a, dim, keepDim)` | the product along dim; a complex array is refused. | +| `Min(a)` | the smallest element as a `Scalar`; an empty or complex array is an error. | +| `Max(a)` | the largest element as a `Scalar`. | +| `MinAxis(a, dim)` / `MaxAxis(a, dim)` | the smallest and largest values along a dimension, with NaN elements never winning. | +| `ArgMin(a)` / `ArgMax(a)` | the index of the smallest and largest element of a 1-D array, with NaN elements skipped as missing. | +| `ArgMinAxis(a, dim)` / `ArgMaxAxis(a, dim)` | the same along a dimension; the result is an int array with the reduced dimension dropped. | +| `Norm(a, p, dim, keepDim)` | the Lp norm along dim, `(sum \|x\|^p)^(1/p)`, always a float64 result. p must be positive, and `math.Inf` selects the max-abs. | +| `Dot(a, b)` | the dot product of two 1-D arrays of equal length as a `Scalar`; two int arrays produce an int, float32 operands accumulate in float64, and any complex operand promotes the result. | +| `All(a)` | whether every element is non-zero; the mask semantics match `Select`. | +| `Any(a)` | whether at least one element is non-zero. | +| `CountNonzero(a)` | the number of non-zero elements. | +| `TopK(a, k, dim)` | the top-k values and their original indices along a dimension, sorted descending by value; NaN elements never rank. | + +#### Shape and indexing + +| Call | What it does | +|---|---| +| `Reshape(a, shape...)` | returns a copy with a new shape of the same element count. | +| `Flatten(a, startDim, endDim)` | returns a copy with the dimensions in `[startDim, endDim]` collapsed into one; negative indices count from the end. | +| `Squeeze(a, dim)` | returns a copy with size-1 dimensions removed: all of them when dim is -1, otherwise that one, which must have size 1. | +| `Unsqueeze(a, dim)` | returns a copy with a new size-1 dimension inserted at dim. | +| `Transpose(a)` | returns a new array with the dimensions reversed; infallible. | +| `TransposeAxes(a, dims...)` | returns a copy with the dimensions reordered according to a permutation of the axis indices. | +| `MoveAxis(a, from, to)` | moves one axis to a new position in the shape. | +| `Tile(a, reps...)` | repeats a whole array reps times per dimension; the result never aliases the input. | +| `Repeat(a, repeats, dim)` | repeats each element `repeats` times along dim. | +| `Flip(a, dims...)` | reverses along the given dimensions, all of them by default. | +| `Roll(a, shift, dim)` | shifts along dim by shift, wrapping around, with a negative shift moving elements backward. | +| `Stack(a, b, dim)` | joins b after a along a new dimension inserted at dim; the two shapes must be identical. | +| `Concat(a, b, dim)` | joins b after a along an existing dimension, every other dimension agreeing. | +| `Pad(a, pad, mode, value)` | extends an array along its trailing dimensions, pad being a flat sequence of pairs `(left, right, top, bottom, …)`; mode is `"constant"`, `"reflect"`, `"replicate"` or `"circular"`. | +| `Slice(a, dim, start, stop)` | selects the half-open range `[start, stop)` along one dimension. A contiguous selection is a read-only view sharing the source's storage, rebased so element i is still payload[i]; an interior range is copied. | +| `Take(a, indices)` | a 1-D copy whose elements are the flat-indexed `a[indices[i]]`; a negative index is an error. | +| `Gather(src, dim, index)` | a copy indexed by index along dim, the output shape being the index shape. | +| `Scatter(self, dim, index, src)` | the inverse of `Gather`, writing src values into a new array at the positions index names; unindexed positions keep the receiver's values. | +| `Row(a, i)` | a 1-D copy of row i of a 2-D array. | +| `Col(a, j)` | a 1-D copy of column j of a 2-D array. | +| `Diag(a)` | extracts the main diagonal of a 2-D array, or builds a diagonal matrix from a 1-D array. | +| `Diagonal(a, offset)` | the elements on the offset-th diagonal of a 2-D matrix as a 1-D array; zero is the main diagonal. | +| `UpperTriangle(a)` / `LowerTriangle(a)` | the triangular part of a square matrix as a copy with the other side zeroed. | +| `BroadcastTo(a, shape...)` | a new array expanded to the target shape: size-1 dimensions replicate, missing leading dimensions prepend, anything else is an error naming both shapes. | +| `BroadcastWith(a, b)` | broadcasts both operands to their common shape and returns the pair, ready for an ordinary element-wise operation. | +| `Nonzero(a)` | the multi-dimensional indices of every non-zero element, one slice per dimension; a complex array is refused. | +| `Argwhere(a)` | the coordinates of every non-zero element as an (nnz, ndim) int array. | +| `SearchSorted(haystack, needles)` | insertion positions for each needle in the sorted haystack, rightmost after equal elements. | +| `AssignBins(a, edges)` | maps every value to its bin index given ascending edges, bin k covering `[edges[k], edges[k+1])`, with values outside clamping to the outer bins. | +| `OneHot(codes, classes)` | expands int class codes into indicator vectors appended as a new last axis, so shape (…) becomes (…, classes); the result is float32. A code outside `[0, classes)` is an error. | + +#### Ordering + +| Call | What it does | +|---|---| +| `Sort(a)` | an ascending copy, NaN elements at the end; the copy carries +0.0 wherever the input held -0.0, and sorted NaNs are fresh quiet NaNs. | +| `ArgSort(a)` | the stable int permutation of indices that would sort the array ascending, NaN at the end. | +| `Reverse(a)` | a copy with the elements in the opposite order; it works for every dtype, complex included. | +| `Unique(a)` | the sorted unique values of a real array; NaN counts once. | +| `Permutation(g, n)` | a uniform permutation of `0..n-1`, Fisher-Yates over unbiased bounded draws. | +| `Shuffle(g, a)` | a shuffled copy of a; the receiver itself is never touched. | + +#### Scans, differences and quadrature + +| Call | What it does | +|---|---| +| `CumSum(a, dim)` | the cumulative sum along dim, with the same shape as the input. | +| `CumProd(a, dim)` | the cumulative product along dim. | +| `Diff(a, order, axis)` | the successive differences along one axis, taken order times; the axis shrinks by order, int arrays keep their dtype, and everything else produces float. | +| `Integrate(y, dx)` | the definite integral of y over uniform spacing dx by the trapezoidal rule, as a float64. | +| `CumulativeIntegrate(y, dx)` | the running trapezoidal integral over uniform spacing, the first element being zero. | +| `EvaluatePolynomial(coeffs, x)` | evaluates coefficients, lowest power first, at the given points. | +| `Jacobian(f, x, opts)` | the Jacobian of f at x as an (m × n) float array, entry (i, j) holding ∂f_i/∂x_j by central differences; f must map the point and every probe to a real rank-1 array of one fixed length, and the probe arrays it is handed are reused between columns. | +| `Covariance(a, b)` | the sample covariance of two equally sized 1-D samples, denominator n-1. | +| `Correlation(a, b)` | the Pearson correlation coefficient of two 1-D samples. | + +#### Core matrix operations + +| Call | What it does | +|---|---| +| `MatMul2D(a, b)` | the matrix product: 2-D×2-D, a matrix times a vector, or a vector times a matrix. The dtype follows the promotion ladder, and float16 operands are refused loudly. | +| `Kron(a, b)` | the Kronecker product of two 2-D matrices; for shapes (m, n) and (p, q) the result is (m·p, n·q). | +| `Einsum(spec, operands...)` | Einstein summation over a spec of the form `"lhs[,lhs,...]->rhs"`, where each ASCII letter names a dimension; an invalid character is an error, so a bad spec can never silently compute. Broadcasting over an ellipsis and batched products go through the same entry point. | +| `Trace(a)` | the sum along the main diagonal of a 2-D square matrix, always as a float64; a complex matrix is answered by `TraceComplex`. | +| `TraceComplex(a)` | the main-diagonal sum of a complex square matrix. | +| `CrossProduct(u, v)` | the vector cross product of two length-3 vectors. | + +#### Sparse construction + +| Call | What it does | +|---|---| +| `NewSparseCOO(indices, values, shape)` | creates a coordinate-format sparse array from explicit indices, values and shape; a nil array, a disagreement on the non-zero count, or a wrong index rank is an error. | +| `SparseFrom(dense)` | extracts a sparse array from a dense one by keeping only the non-zero elements; the values keep the dense array's dtype. | +| `SparseCOO.NNZ()` | the number of stored non-zero entries. | +| `SparseCOO.Dense()` | materialises the sparse array as a dense `Array`. | +| `SpAdd(a, b)` | the element-wise sum of two sparse arrays of the same shape, as a dense array, because addition can collapse zeros into non-zeros. | +| `SpMul(s, dense)` | the element-wise product of a sparse array and a dense array of the same shape, as a dense array. | +| `SpMatMul(s, dense)` | a sparse (n×k) matrix times a dense (k×m) matrix, as a dense n×m result; every stored coordinate is validated, so an out-of-range index is an error rather than a panic. | + +The compressed views and their solvers live in `linalg`, not here: see +[Package `linalg`](#package-linalg). + +#### Special functions + +| Call | What it does | +|---|---| +| `Gamma(a)` | the gamma function Γ(x) of each element. | +| `LnGamma(a)` | the natural logarithm of `\|Γ(x)\|`; the sign for negative arguments is dropped. | +| `LnFactorial(n)` | the natural logarithm of n! as a float64. | +| `Digamma(a)` / `Trigamma(a)` | the digamma ψ(x) and the trigamma ψ′(x) of each element. | +| `Beta(x, y)` | the Euler beta function `B(x, y) = Γ(x)Γ(y)/Γ(x+y)`, the two arrays of one shape. | +| `Erf(a)` / `Erfc(a)` | the error function and the complementary error function. | +| `BesselJ(n, x)` / `BesselY(n, x)` | the Bessel functions of the first and second kind of integer order n at one real point; Y is defined for `x > 0` and refuses a non-positive argument. | +| `BesselJRealOrder(nu, x)` | the Bessel function of the first kind of real order ν at one positive real point; series below the crossover, above it the order's fractional pair is seeded from the expansion and the recurrence runs in its stable direction. Orders within 1e-8 of a non-negative integer are served by the exact integer algorithm; a negative order or a non-positive argument is an error. | +| `BesselI0(a)` / `BesselI1(a)` / `BesselIn(n, a)` | the modified Bessel function of the first kind, of order zero, one, and integer order n. | +| `BesselK0(a)` / `BesselK1(a)` / `BesselKn(n, a)` | the modified Bessel function of the second kind; any NaN or non-positive element is an error naming it, never a silent NaN. | +| `Airy(a)` | the Airy functions of the first and second kind, returned in that order; outside `\|x\| ≤ 8` the answer is NaN rather than a silently wrong value. | +| `ExpIntegralE1(a)` | the exponential integral `E1(x)`, defined for `x > 0`. | +| `ExpIntegralEi(a)` | the exponential integral `Ei(x)` for real x ≠ 0, the negative side through `E1`. | +| `FresnelC(a)` / `FresnelS(a)` | the Fresnel cosine and sine integrals. | +| `Legendre(l, x)` | the Legendre polynomial `P_l(x)` by the Bonnet recurrence, `P_0 = 1`, `P_1 = x`, no extra scaling. | +| `LegendreAssociated(l, m, x)` | the associated Legendre function `P_l^m(x)` with the Condon-Shortley phase folded in and no extra normalisation. | +| `Hermite(n, x)` | the physicists' Hermite polynomial `H_n(x)`. | +| `Laguerre(n, alpha, x)` | the generalised Laguerre polynomial `L_n^α(x)`; the degree must be at least 0. | +| `ChebyshevT(n, x)` / `ChebyshevU(n, x)` | the Chebyshev polynomials of the first and second kind. | +| `SphericalHarmonic(l, m, theta, phi)` | the complex spherical harmonic `Y_l^m(θ, φ)` with the Condon-Shortley phase and the orthonormal convention; theta and phi must share a shape. | +| `SphericalHarmonicReal(l, m, theta, phi)` | the real spherical harmonic built from the complex one under the standard branch convention. | +| `SphericalBesselJ(l, x)` / `SphericalBesselY(l, x)` | the spherical Bessel functions of the first and second kind; every `y_l` diverges to -Inf at x = 0. | +| `EllipticK(m)` / `EllipticE(m)` / `EllipticPi(n, m)` | the complete elliptic integrals of the first, second and third kind in the parameter convention `m = k²`. | +| `EllipticKScalar(m)` / `EllipticFScalar(phi, m)` | the same first-kind integral at one point, complete and incomplete, where the incomplete form is the amplitude view behind the Jacobi functions. | +| `JacobiSN(u, m)` / `JacobiCN(u, m)` / `JacobiDN(u, m)` | the Jacobi elliptic functions of the parameter m, shared across the u array. | +| `JacobiCDScalar(u, m)` | `cd(u, m) = cn(u, m)/dn(u, m)` at one point. | +| `Hypergeometric2F1(a, b, c, x)` | the Gauss hypergeometric function `₂F₁(a, b; c; x)` element-wise over x with scalar parameters; a non-positive integer c is an error. | + +#### Interpolation + +| Call | What it does | +|---|---| +| `Interpolate(xs, ys, query)` | the piecewise-linear interpolation of the points `(xs[i], ys[i])` at each query; xs need only be non-decreasing, a query outside the range clamps to the boundary, and a NaN query or a non-finite knot is refused. | +| `InterpolateMonotone(xs, ys, query)` | the monotone piecewise cubic through the samples, passing through every knot with a slope that never exceeds twice the neighbouring secants; xs must be strictly increasing. | +| `Interpolate2D(grid, xs, ys, x0, y0, dx, dy)` | the bilinear interpolation of a regular rows × cols grid at a set of query points, exact on any field that is bilinear within a cell; a query outside the rectangle is an error, never a silent clamp, and either sign of dx or dy works. | +| `InterpolateGrid(grid, origins, steps, queries)` | multilinear interpolation on a grid of any rank, queries being an (m × rank) matrix of coordinates; a query outside the grid clamps to the boundary and a rank above 12 is an error. | + +#### Quasirandom sequences + +| Call | What it does | +|---|---| +| `HaltonPoints(n, dim, skip)` | the first n Halton points of the given dimension as an (n, dim) float64 array, skipping the leading skip points; dim must lie in `[1, 32]`, beyond which the available small primes run out. | +| `SobolPoints(n, dim, skip)` | the same for the base-2 digital Sobol sequence; dim must lie in `[1, 40]`, the width of the direction-number table. | + +Both drop the index-zero origin before the skip is applied, because it carries +no information, and both refuse a negative n or skip and require `n + skip` +below 2^32, the index budget past which the arithmetic would wrap back to the +origin. + +#### Random numbers + +| Call | What it does | +|---|---| +| `NewGenerator(seed)` | seeds a fresh generator; any int64 seed is valid, and the stream is xoshiro256++ seeded through splitmix64, stable across Go releases. | +| `Generator.Next()` | advances the generator and returns the raw 64-bit value. | +| `Generator.Unit()` | draws one uniform float in `[0, 1)` with 53-bit resolution. | +| `Generator.NormalUnit()` | draws one standard normal value by the polar Box-Muller method. | +| `Floats(g, n)` | n uniform floats in `[0, 1)` with 53-bit resolution. | +| `Float32s(g, n)` | n uniform float32 values in `[0, 1)` with 24-bit resolution. | +| `Ints(g, n, min, max)` | n uniform ints in `[min, max)`, drawn without modulo bias; min must be below max. | +| `Normal(g, n, mean, std)` | n Gaussian draws of the given mean and standard deviation, bit-stable across Go releases; a negative or NaN std is an error. | +| `TruncatedNormal(g, shape, mean, std)` | an array of the given shape drawn from a normal truncated to ±2σ, as float32; an invalid shape yields a nil array. | +| `Permutation(g, n)` | a uniform permutation of `0..n-1`. | +| `Shuffle(g, a)` | a shuffled copy of a. | +| `Splitmix64(state)` | advances the splitmix64 stream one step, returning the advanced state and the mixed output; the mixer is a bijection on uint64, so callers seeding their own streams can build on it. | +| `Substream(seed, index)` | the index-th member of the stream family one seed carries: the index is mixed through splitmix64 before it meets the seed, and both mixings are bijections, so distinct indices give provably distinct initial states. The index starts at zero; a negative one is an error. | + +Every draw of a given seed and call sequence is reproducible, which is what the +examples rely on. + +#### Execution policy + +| Call | What it does | +|---|---| +| `SetNumCPU(n)` | sets the number of goroutines the parallel kernels may use and returns the previous value; a value below 1 resets to `runtime.NumCPU()`, and a running kernel finishes with its old worker count. | +| `NumWorkers()` | returns the current worker count. | + +The default is `runtime.NumCPU()`, the logical CPUs available to the process. +The heavy loops of the library, matrix products, convolutions, element-wise +maps, axis reductions, scans and the FFT, run in parallel across this count and +reduce in a fixed order, so a given element order gives bit-identical results +run to run. + +#### Payload and element access + +| Call | What it does | +|---|---| +| `Array.RawFloats()` | returns the float64 payload directly: element i of the array sits at payload index i, views included, because the package never sets strides. Treat it as read-only; only an array the caller owns is safe to write through. | +| `Array.RawFloat32s()` | the float32 payload, same contract. | +| `Array.RawInts()` | the int64 payload, same contract. | +| `Array.RawHalves()` | the float16 payload, the raw IEEE 754 binary16 bit patterns, same contract. | +| `Array.RawComplexes()` | the complex128 payload, same contract. | +| `Array.Elements[E]()` | the flat payload converted to the caller's element type; the ladder runs int64 ← float16 ← float32 ← float64 ← complex128, so a wider E converts and a narrower one is an error naming both dtypes. The slice returned is a copy. | +| `Array.ComplexValues(name)` | the elements as complex values, sharing the payload as a read-only alias for a contiguous complex array and copied out otherwise. | +| `FloatAt(a, index...)` | the element at the given index as a float64; the array must be float. | +| `IntAt(a, index...)` | the element at the given index as an int64; the array must be int. | +| `ComplexAt(a, index...)` | the element at the given index as a complex128; the array must be complex. | +| `Item(a)` | the single element of a 1-element array as a float64, the real part for a complex array. | +| `Array.FloatAt(i)` | the numeric read primitive the derivative packages build on. | +| `Array.ComplexAt(i)` | element i as a complex128, converting a real element. | +| `Array.SetFloatAt(i, v)` | sets element i from v, converting to the array's dtype; the numeric write primitive complementing `FloatAt`. | +| `WithInt(a, v, index...)` | a new int array with v at the given index; the receiver is unchanged. | +| `WithFloat(a, v, index...)` | a new float array with v at the given index. | +| `WithComplex(a, v, index...)` | a new complex array with v at the given index. | + +The `Raw*` accessors are for kernels that must touch the payload directly; every +other path through the library reads through the accessors, so a write through a +raw slice is the caller's responsibility. + +#### The fluent chain + +| Call | What it does | +|---|---| +| `Pipe(a)` | starts the fluent chain from a, returning a `*Pipeline` that holds the current array and the first error encountered during the chain. | +| `Pipeline.Result()` | ends the chain, returning the array and that first error. | + +Every other `Pipeline` method is one of the core operations, eager and returning +the pipeline itself, so the methods are the package-level calls without the +intermediate error checks. The chain is defined in `linalg` and re-exported +here, so its full method list belongs to [Package `linalg`](#package-linalg). + +#### Domain re-exports + +Every symbol of every domain package is also a root symbol, unchanged. Each +package below has its own section further down this document: + +| Package | Reachable here as | +|---|---| +| `linalg` | `Solve`, `SVD`, `Eigen`, `SpEigen`, `Pipe` and the rest, plus `ArrayFromFloatsSafe`. See [Package `linalg`](#package-linalg). | +| `signal` | `FFT`, `DWT`, `CWT`, `Conv2D`, `KalmanFilter`, the window functions and the filter designs. See [Package `signal`](#package-signal). | +| `integrate` | `IntegrateODE`, `IntegrateHeat1D`, `IntegrateND`, `IntegrateFilon`, the FEM meshes and solvers. See [Package `integrate`](#package-integrate). | +| `stats` | `LinearRegression`, `PCA`, `KMeans`, `Histogram`, the distributions and the hypothesis tests. See [Package `stats`](#package-stats). | +| `optim` | `Minimise`, `FindRoot`, `MinimiseLBFGS`, `LevenbergMarquardt`, `LevenbergMarquardtFit`. See [Package `optim`](#package-optim). | +| `io` | `LoadCSV`, `SaveHDF5`, `LoadFITS`, the NetCDF pair and the mapping helpers. See [Package `io`](#package-io). | +| `grad` | `Tensor`, `Backward`, `Hessian`, `AdjointODE`, `SampleHMC`. See [Package `grad`](#package-grad). | + +Each package's options and result types are re-exported the same way, so +`tensor.ODEOptions`, `tensor.LBFGSOptions` and `tensor.HessianOptions` are the +types the corresponding root calls take. + +#### Options and results + +The core owns two exported structs, and both are plain option and data records +with no methods beyond those of their fields. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `SparseCOO.Indices` | `*Array` | required | the coordinates, shape (nnz, ndim) and dtype int. | +| `SparseCOO.Values` | `*Array` | required | the stored values, shape (nnz,) and the element dtype of the array. | +| `SparseCOO.Shape` | `[]int` | required | the dense shape the coordinates index into. | +| `JacobianOptions.Step` | `float64` | zero, meaning the per-column default | the absolute difference step applied to every coordinate. Zero or negative selects `sqrt(ε)·max(1, abs(x_j))` per column, the largest step whose central-difference truncation error still sits below the rounding floor. | + +#### Errors + +Every error the library returns carries the `tensor: ` prefix, whatever package +raised it. + +- Two operands of a binary element-wise operation whose shapes disagree: an + error naming both shapes; nothing is computed. +- A dimension index outside the array's rank, in any axis reduction, scan, + `MoveAxis`, `Squeeze`, `Unsqueeze`, `Concat`, `Stack`, `Diff` or + `InterpolateGrid`: an error naming the dimension and the shape. +- A reduced dimension of size zero in `CumSum` or `CumProd`: an error naming the + dimension. +- A payload length that does not exactly fill the declared shape in any `From…` + constructor: an error naming the count and the shape. +- A shape list that is empty, holds a negative dimension, or holds more elements + than fit in an index: an error, raised before any allocation. +- A dtype constant that is not one of the five: an error from `Zeros`, `Ones` + and the `Full*` family, and a nil array from `New`. +- An element read or write whose array has another dtype: `IntAt`, `FloatAt`, + `ComplexAt`, `WithInt`, `WithFloat` and `WithComplex` each name the dtype they + found. +- `Quo` or `QuoI` on a non-int array, or with a zero divisor: an error. +- `Pow` or `PowI` with a negative exponent on an int array: an error. +- A negative `n` in `Floats`, `Float32s`, `Ints`, `Normal` or `Permutation`; a + `min` at or above `max` in `Ints`; a negative or NaN `std` in `Normal`: an + error. `TruncatedNormal` is the exception, answering a nil array. +- A `Slice` range outside the dimension's extent, or a dimension outside the + rank: an error naming the range or the dimension. +- `RangeBy` with a zero step: an error. +- `HaltonPoints` with dim outside `[1, 32]`, `SobolPoints` with dim outside + `[1, 40]`, either with a negative n or skip, or either with `n + skip` at or + above 2^32: an error. +- `Take` with a negative index, `CopyRows` with an index outside the leading + dimension, `OneHot` with a code outside `[0, classes)`: an error. +- A complex array in a call that needs an ordering, among them `Sort`, + `ArgSort`, `Unique`, `Sign`, `IsNaN`, `Min`, `Max`, `MinAxis`, `MaxAxis`, + `ArgMin` and `ArgMax`: an error, because a complex lattice has no ordering. +- A complex array in `Mean`, `MeanAxis`, `Prod` or `Norm`: an error, because + there is no real-valued answer. +- An empty array in `Mean`, `Min`, `Max`, `ArgMin` or `ArgMax`: an error. +- Two operands of different length in `Dot`, or `Dot` on anything but 1-D + arrays: an error naming the shapes and the lengths. +- A float16 operand in `MatMul2D` or `Einsum`: an error naming `Astype` as the + conversion. +- `Interpolate2D` with a grid that is not rank 2 or has an extent below 2, a + zero or non-finite spacing or origin, or a query outside the grid: an error. +- `Interpolate` with a non-finite knot or a NaN query, `InterpolateMonotone` + with knots that are not strictly increasing: an error. + +#### Workflow + +```go +// Build a flat payload, reshape it, slice it, reduce it, and hand the +// result to a domain entry point. +vals := make([]float64, 16) +for i := range vals { + vals[i] = float64(i + 1) +} +flat, err := tensor.FromFloats(vals, 16) +if err != nil { + return err +} +m, err := tensor.Reshape(flat, 4, 4) +if err != nil { + return err +} + +// Whole rows of a 2-D array: a read-only view sharing the storage. +rows, err := tensor.Slice(m, 0, 1, 4) +if err != nil { + return err +} +// A partial column range: copied, because only a contiguous +// selection is a view. +sq, err := tensor.Slice(rows, 1, 0, 3) +if err != nil { + return err +} + +// Reduce the block to one value per column. +rhs, err := tensor.MeanAxis(sq, 0) +if err != nil { + return err +} + +// The domain packages take over from here; Solve reads the core +// array and returns one. +x, err := tensor.Solve(sq, rhs) +if err != nil { + return err +} +fmt.Println(sq.Shape(), x.Shape()) +``` + +### Package `linalg` + +Dense and sparse linear algebra: factorisations and their solves, the +spectral decompositions, matrix functions, polynomial fitting and +root finding, cubic splines, the compressed sparse views, and the +direct and Krylov methods that run on them. Every symbol is +re-exported by the root package, so `linalg.Solve` and `tensor.Solve` +are one function; the arrays it reads and writes are the shared core +array, re-exported by the root package as `tensor.Array`. + +The package does not choose an algorithm for the caller. Dense or +sparse, symmetric or general, direct or iterative is the decision each +call makes, and the names say which route they take. Complex data is +supported across the dense surface and on four named sparse entry +points; every other sparse routine refuses it rather than promoting +it. The `Pipeline` is eager sugar over the package-level calls, with +no lazy graph and no autograd. + +#### Constructing an input + +| Call | What it does | +|---|---| +| `ArrayFromFloatsSafe(v, n)` | builds a float vector of length n from v, copying the values so the caller keeps ownership of the slice. | + +`ArrayFromFloatsSafe` is the constructor for a caller that imports +this package on its own, without the root facade: no general array +constructors are exported here, so this one exists to build the input +to a call. The general constructors belong to the root package, where +`tensor.FromFloats`, `tensor.Zeros`, `tensor.Identity` and their +neighbours are the way in. + +#### Dense factorisations and solves + +| Call | What it does | +|---|---| +| `Solve(a, b)` | returns x with a·x = b for a square a and b a vector of length n or a matrix with n rows; one LU with partial pivoting, complex promotion when either operand is complex. | +| `Inv(a)` | returns the inverse of a square matrix, through the same LU kernel. | +| `Det(a)` | returns the determinant of a square real matrix; a singular matrix gives 0, not an error. | +| `DetComplex(a)` | the same for a square complex matrix, as a `complex128`. | +| `Cholesky(a)` | returns the lower factor L with a = L·Lᵀ for a symmetric positive definite a; a non-positive pivot is refused. | +| `CholeskyUpdate(l, x)` | returns the lower factor of A + x·xᵀ from the factor of A, by orthogonal rotations. | +| `CholeskyDowndate(l, x)` | returns the lower factor of A − x·xᵀ, by hyperbolic rotations; leaving the positive definite cone is an error. | +| `QR(a)` | returns q (m×m orthogonal) and r (m×n upper triangular) with a = q·r, for m ≥ n. | +| `LeastSquares(a, b)` | solves A·x = b in the least-squares sense for m×n a with m ≥ n and b of m rows or m×k; a rank-deficient system is refused. | +| `SolveTridiagonal(a, b, c, d)` | solves the tridiagonal system by the Thomas algorithm: a the lower diagonal of length n−1, b the main diagonal, c the upper diagonal, d the right-hand side; a zero pivot is refused. | +| `SolveCyclicTridiagonal(a, b, c, d)` | the same for the periodic system, where `a[0]` is the corner A[0][n−1] and `c[n−1]` the corner A[n−1][0], by Sherman-Morrison on the Thomas elimination. | + +#### Eigenproblems and singular values + +| Call | What it does | +|---|---| +| `Eigen(a)` | returns the eigenvalues of a real symmetric matrix, ascending, and the orthonormal eigenvectors as columns; asymmetry beyond a scale-relative 1e-12 is refused. | +| `EigenComplex(a)` | the same for a complex Hermitian matrix; values are real, vectors complex, and the same 1e-12 Hermitian tolerance applies. | +| `EigenGeneral(a)` | the eigenvalues and eigenvectors of a square matrix of any dtype; real and int inputs promote to `complex128`, values descend by magnitude, and a real matrix may carry conjugate pairs. | +| `EigenGeneralised(a, b)` | solves A·v = λ·B·v for symmetric a and symmetric positive definite b, through the Cholesky factor of b; values ascend, eigenvectors are B-orthonormal columns. | +| `SVD(a)` | the thin decomposition a = U·Σ·Vᵀ for m ≥ n: U is m×n with orthonormal columns, Σ a length-n vector descending, Vᵀ n×n orthogonal; a wide matrix is transposed first and the factors are swapped back. | +| `SVDComplex(a)` | the same shapes for a complex matrix, A = U·diag(Σ)·Vᴴ. | +| `SchurComplex(a)` | the complex Schur decomposition t·q with a = q·t·qᴴ, t upper triangular and q unitary; the diagonal of t carries the eigenvalues. | +| `Pinverse(a, eps)` | the Moore-Penrose pseudoinverse, built from the SVD by inverting singular values strictly above eps; eps ≤ 0 uses max(m, n)·max(Σ)·ε. | +| `MatrixRank(a, eps)` | the count of singular values strictly above eps, with the same default when eps ≤ 0. | +| `Cond(a, eps)` | the 2-norm condition number σmax/σmin; a singular value at or below a positive eps, or exactly zero, gives +Inf. | + +Both `SVD` routes take the singular values from the eigenvalues of a +squared matrix, BᵀB or Aᴴ·A. Their accuracy on very small singular +values is therefore that of a squared condition number: exact enough +for a rank decision, weaker for resolving near-null directions. The +same caveat applies to the rank and condition answers above, which +are read off that spectrum. + +#### Rank-revealing QR + +| Call | What it does | +|---|---| +| `RRQR(a)` | returns q, r, the column permutation `perm` and the numerical rank of an m×n matrix with m ≥ n, from A·P = Q·R with column pivoting; the rank counts \|R[i,i]\| above n·ε·max\|R\|. | +| `RRQRRank(a)` | the same rank count without forming the orthogonal factor. | +| `SolveRRQR(a, b)` | solves min ‖A·x − b‖₂ through the pivoted factorisation; on rank-deficient input it returns the minimum-norm least-squares solution. | + +#### Regularised and truncated solves + +| Call | What it does | +|---|---| +| `SolveTikhonov(a, b, lambda)` | solves min ‖A·x − b‖² + λ‖x‖² with the identity prior, scaling every singular direction by σ/(σ²+λ); a non-positive λ is an error. | +| `SolveTruncated(a, b, rank)` | keeps the `rank` largest singular values and discards the rest, the truncated pseudoinverse; a singular value that vanishes below the requested rank is an error. | + +#### Matrix functions + +| Call | What it does | +|---|---| +| `MatrixExp(a)` | returns exp(A) of a square matrix by scaling and squaring over a diagonal Padé approximant; the zero matrix gives the identity. | +| `MatrixSqrt(a)` | the principal square root of a symmetric positive semi-definite matrix by diagonalisation, with small negative eigenvalues clamped to zero and a genuinely indefinite spectrum refused; a nonsymmetric matrix runs through the complex Schur route and is refused when the principal root leaves the reals. | +| `MatrixLog(a)` | the principal logarithm of a real symmetric positive definite matrix by the same diagonalisation route; a non-positive eigenvalue is refused, and a nonsymmetric matrix takes the Schur-Parlett route. | + +#### Fitting, roots and splines + +| Call | What it does | +|---|---| +| `FitPolynomial(x, y, degree)` | returns the coefficients, lowest power first, of the degree-th polynomial through the samples, by QR least squares over the Vandermonde matrix; n < degree+1 or mismatched lengths are errors. | +| `PolynomialRoots(coeffs)` | returns the roots of the polynomial whose coefficients are lowest power first, as a complex vector descending by magnitude, through the companion matrix; trailing zeros are stripped, the zero polynomial is an error, and a nonzero constant gives an empty vector. | +| `NewCubicSpline(xs, ys)` | builds the natural cubic spline through the points; the abscissae must be strictly increasing and at least 3 points are required. | +| `CubicSpline.At(x)` | evaluates the spline at one point; outside the knot range the answer is NaN. | +| `CubicSpline.Evaluate(x)` | maps an array of query points element-wise, refusing a complex query array. | + +#### Iterative solves over an operator + +| Call | What it does | +|---|---| +| `GMRES(op, b, restart, maxIter, tol)` | solves A·x = b by restarted GMRES, with A given as the function `op` that maps a vector to A·v; restart ≤ 0 or above n means n, maxIter ≤ 0 means 20 cycles, tol ≤ 0 means 1e-10 relative residual, an exhausted budget is an error naming the best residual. `op` is handed a read-only vector the solver refills for each call, so it must not retain or modify it, and it may hand back a buffer it reuses. | | + +#### Sparse storage + +| Call | What it does | +|---|---| +| `CSRFromCOO(s)` | converts a COO matrix to compressed sparse row, summing duplicate coordinates, dropping explicit zeros and sorting each row; complex values are refused. | +| `CSCFromCOO(s)` | the same canonicalisation into compressed sparse column form. | +| `SparseCSR.ToCSC()` | converts the same matrix to the column view. | +| `SparseCSC.ToCSR()` | converts the same matrix to the row view. | +| `SparseCSR.Transpose()` | returns the transpose, in CSR form. | +| `SparseCSR.MatVec(x)` | computes A·x for a dense vector of length Cols; output rows are independent and split across workers. | +| `SparseCSC.MatVec(x)` | the same product over the column view, with a fixed scatter order. | +| `SparseCSR.MatMulDense(x)` | computes A·x for a dense matrix of shape (Cols, k), streaming the non-zeros once. | +| `SparseCSR.MatMulSparse(o)` | returns the sparse product A·B where A.Cols equals B.Rows, storing only what survives. | +| `SparseCSR.NNZ()` | the count of stored non-zeros. | +| `SparseCSC.NNZ()` | the same for the column view. | + +#### Sparse direct factorisations + +| Call | What it does | +|---|---| +| `NewSparseCholesky(a, ordering)` | factors a symmetric positive definite matrix into L with P·A·Pᵀ = L·Lᵀ; the lower triangle defines the matrix, and a stored upper entry without its equal lower counterpart is refused. One factorisation solves any number of right-hand sides. | +| `SparseCholesky.Solve(b)` | computes A⁻¹·b for a dense vector: permute, forward substitution, backward substitution, undo. | +| `SparseCholesky.Update(x)` | applies A ← A + x·xᵀ to the factor in place on the stored pattern, refusing an update that would need fill the pattern does not hold. | +| `SparseCholesky.Downdate(x)` | applies A ← A − x·xᵀ the same way, refusing a matrix that leaves the positive definite cone. | +| `SparseCholesky.Permutation()` | the elimination order, position k holding the original index factored k-th; a copy. | +| `SparseCholesky.NNZ()` | the count of stored non-zeros in the factor, diagonal included. | +| `NewSparseLU(a)` | factors a square matrix with partial pivoting into P·A = L·U; complex, rectangular and non-finite inputs are refused, and a zero pivot column means a singular matrix, which is reported. | +| `SparseLU.Solve(b)` | computes A⁻¹·b for a dense vector through L and U. | +| `SparseLU.Permutation()` | the row elimination order, position k holding the original index of the row factored k-th; a copy. | +| `SparseLU.NNZ()` | the count of stored non-zeros: L's strict columns, U's strict rows and U's diagonal. | +| `NewSparseILU(a)` | builds the ILU(0) factorisation of a square matrix in its CSR pattern, as the Krylov preconditioner; a missing diagonal entry or a zero pivot is refused. | +| `SparseILU.Apply(r)` | solves (L·U)·x = r over the stored pattern, approximating A⁻¹·r to the accuracy the dropped fill allows. | + +The fill-reducing permutation is chosen with `SparseOrdering`. It has +three values: `SparseOrderingNatural`, which eliminates in stored +order and is the reference point every ordering is measured against; +`SparseOrderingReverseCuthillMcKee`, which orders every component in +reverse breadth-first order from a pseudo-peripheral start and is the +classic choice for mesh-shaped patterns; and +`SparseOrderingMinimumDegree`, which eliminates the vertex with the +fewest remaining neighbours at each step and is the stronger choice on +irregular patterns. + +#### Sparse iterative solves + +| Call | What it does | +|---|---| +| `SpSolve(a, b, tol, maxIter, precond...)` | solves A·x = b for a real symmetric positive definite sparse A by preconditioned conjugate gradient; the default preconditioner is the Jacobi diagonal and an ILU(0) from `NewSparseILU` may replace it. | +| `SpSolveBiCGSTAB(a, b, tol, maxIter, precond...)` | the same for a general real square sparse A, by BiCGSTAB, with the same preconditioner choice. | +| `SpSolveComplexCG(a, b, tol, maxIter)` | the Hermitian positive definite complex case, by conjugate gradient with a complex Jacobi preconditioner. | +| `SpSolveComplexBiCGSTAB(a, b, tol, maxIter)` | the general non-Hermitian complex case, by BiCGSTAB. | + +All four stop on the relative residual ‖b − A·x‖₂ ≤ tol·‖b‖₂ with +tol ≤ 0 meaning 1e-10 and maxIter ≤ 0 meaning n steps. An +unconverged solve is an error naming the residual achieved, never a +silent approximation. A zero or missing diagonal entry is refused, +since the Jacobi preconditioner divides by it. + +#### Sparse least squares + +| Call | What it does | +|---|---| +| `SpLSQR(a, b, tol, maxIter, conlim)` | minimises ‖A·x − b‖₂ over a sparse overdetermined A by the Golub-Kahan recursion of Paige and Saunders, returning the solution and a `LeastSquaresInfo`. | +| `SpLSMR(a, b, tol, maxIter, conlim)` | minimises ‖Aᵀ(b − A·x)‖₂ by Fong and Saunders' LSMR, whose normal-equations residual moves monotonically; the answers agree on a consistent rank-deficient system, where both give the minimum-norm solution. | + +tol ≤ 0 means 1e-10 and feeds both the residual and the +normal-equations test; maxIter ≤ 0 means 2n steps; conlim ≤ 0 means +1e8, and cond(A) passing it stops the iteration without claiming +convergence. Running out of steps with every tolerance unmet is an +error naming the residual achieved. + +#### Sparse eigensolvers and the exponential action + +| Call | What it does | +|---|---| +| `SpEigen(s, k, gen)` | the k eigenvalues of largest magnitude of a real symmetric sparse matrix, each with its unit eigenvector, by Lanczos, values descending by magnitude. | +| `SpEigenComplex(s, k, gen)` | the same for a Hermitian sparse matrix with complex entries; a non-Hermitian matrix is refused. | +| `SpEigenGeneral(s, k, gen)` | the k eigenvalues of largest magnitude of a general real sparse matrix, as a `complex128` vector with the matching eigenvectors, by Arnoldi; symmetric input gets a cheaper answer from `SpEigen`. | +| `SpEigenGeneralComplex(s, k, gen)` | the general complex case, non-Hermitian operators included. | +| `SpExpApply(a, v, steps)` | returns exp(A)·v for a real symmetric sparse A and a vector v, by Krylov projection; steps ≤ 0 means min(n, 40), and steps ≥ n decomposes the whole space so the answer is exact. | + +The eigensolvers take a generator (`tensor.Generator`) for their start +vector; passing nil uses a fixed seed, so an unseeded call is +reproducible. The Ritz pairs are approximations whose accuracy +improves with the iteration budget, unlike the exact answers `Eigen` +computes. + +#### The pipeline + +| Call | What it does | +|---|---| +| `Pipe(a)` | starts a pipeline from a. | +| `Pipeline.Result()` | returns the current array and the first error seen during the chain, or nil. | +| `Pipeline.Add(b)`, `Sub(b)`, `Mul(b)`, `Div(b)` | element-wise array-array arithmetic. | +| `Pipeline.AddF(v)`, `SubF(v)`, `MulF(v)` | float scalar arithmetic. | +| `Pipeline.AddI(v)`, `SubI(v)`, `MulI(v)` | int scalar arithmetic. | +| `Pipeline.Neg()` | the arithmetic negation. | +| `Pipeline.Maximum(b)`, `Minimum(b)` | element-wise maximum and minimum against another array. | +| `Pipeline.ClipF(lo, hi)`, `ClipI(lo, hi)` | clamps into the closed interval. | +| `Pipeline.Abs()`, `Sqrt()`, `Exp()`, `Log()`, `Floor()` | element-wise maths. | +| `Pipeline.Tanh()`, `Sigmoid()` | the activations. | +| `Pipeline.SumAxis(dim)`, `MeanAxis(dim)` | reductions along one axis. | +| `Pipeline.ArgMaxAxis(dim)`, `ArgMinAxis(dim)` | the index of the extremum along one axis. | +| `Pipeline.TopK(k, dim)` | the top-k values along dim, discarding the matching indices; the package-level `TopK` is the call when the indices matter. | +| `Pipeline.Reshape(shape...)`, `Flatten(startDim, endDim)`, `Squeeze(dim)`, `Unsqueeze(dim)` | shape manipulation. | +| `Pipeline.Transpose()`, `TransposeAxes(dims...)` | the transpose and the axis permutation. | +| `Pipeline.MatMul2D(b)` | the matrix product. | +| `Pipeline.Inv()` | the inverse of the current square matrix. | +| `Pipeline.Solve(b)` | the solve of the current square matrix against b. | + +Every step returns the pipeline, so the chain reads as one +expression. A step that fails records its error and the steps after +it become no-ops; `Result` reports the first one. + +#### Options and results + +**`LeastSquaresInfo`**: what a sparse least-squares iteration +achieved and which stopping test ended it. Returned by `SpLSQR` and +`SpLSMR`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Iterations` | `int` | 0 | the number of Golub-Kahan steps folded into the answer. | +| `Criterion` | `string` | `""` | the test that ended the iteration: `LeastSquaresResidual`, `LeastSquaresNormal` or `LeastSquaresCondition`; empty when the exact answer x = 0 was returned without a step. | +| `ResidualNorm` | `float64` | 0 | the achieved ‖b − A·x‖₂, recomputed from the returned x. | +| `NormalResidual` | `float64` | 0 | the achieved ‖Aᵀ(b − A·x)‖₂, recomputed the same way. | +| `MatrixNorm` | `float64` | 0 | the iteration's running estimate of ‖A‖. | +| `Condition` | `float64` | 0 | the iteration's running estimate of cond(A). | +| `Converged` | `bool` | `false` | true when a residual or normal-equations test fired; a condition stop leaves it false, so a caller must treat that answer as a warning rather than a solution. | + +The three criterion names are the exported constants +`LeastSquaresResidual` (`"residual"`), `LeastSquaresNormal` +(`"normal"`) and `LeastSquaresCondition` (`"condition"`). The +iteration stops when any of the three fires. + +**`SparseCSR`**: the compressed sparse row view of a COO matrix, +the format the iterative solvers and the Lanczos eigensolver run on. +Built by `CSRFromCOO`, `SparseCSC.ToCSR` and `SparseCSR.Transpose`. +All fields are exported, so a caller may build or inspect the view +directly, but the constructors above are what guarantee the +row-major sorted canonical form the methods assume. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `RowStart` | `[]int` | nil | the offset of each row's entries; length Rows+1. | +| `ColIdx` | `[]int` | nil | the column index of each stored entry, ascending within a row. | +| `Values` | `[]float64` | nil | the stored entries, aligned with `ColIdx`. | +| `Rows` | `int` | 0 | the row count. | +| `Cols` | `int` | 0 | the column count. | + +**`SparseCSC`**: the compressed sparse column view, the format the +direct factorisations run on, where a column at a time is eliminated +and the fill of one column extends the entries below the diagonal. +Built by `CSCFromCOO` and `SparseCSR.ToCSC`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `ColStart` | `[]int` | nil | the offset of each column's entries; length Cols+1. | +| `RowIdx` | `[]int` | nil | the row index of each stored entry, ascending within a column. | +| `Values` | `[]float64` | nil | the stored entries, aligned with `RowIdx`. | +| `Rows` | `int` | 0 | the row count. | +| `Cols` | `int` | 0 | the column count. | + +#### Errors + +Every message carries the `tensor: ` prefix and names the call. +The conditions this package reports: + +- a shape that is not square, not 2-D, or wider than tall where the + algorithm needs m ≥ n: the shape is named in the error. +- a complex input to a routine with no complex form: `Det`, `QR`, + `Cholesky`, `LeastSquares`, `Eigen`, `SVD`, `Pinverse`, + `MatrixRank`, `Cond`, `RRQR`, the `CSRFromCOO` and `CSCFromCOO` + conversions, and every sparse solver except the four complex entry + points. The complex routes exist and are named: `DetComplex`, + `EigenComplex`, `SVDComplex`, `SchurComplex`, + `SpSolveComplexCG`, `SpSolveComplexBiCGSTAB`, `SpEigenComplex` and + `SpEigenGeneralComplex`. +- a singular matrix: `Solve` and `Inv` report it; `Det` and + `DetComplex` return 0 instead, and `NewSparseLU` reports a zero + pivot column. +- a matrix that is not positive definite: `Cholesky` on a non-positive + pivot, `CholeskyDowndate` and `SparseCholesky.Downdate` when a + diagonal can no longer dominate, and `NewSparseCholesky` on an + input whose stored upper triangle contradicts its lower one. +- a zero or missing diagonal entry in `SpSolve`, `SpSolveBiCGSTAB`, + `SpSolveComplexCG` and `SpSolveComplexBiCGSTAB`: the Jacobi + preconditioner divides by it. +- a rank-deficient system: `LeastSquares` refuses what it cannot + back-substitute safely, while `RRQR` and `RRQRRank` report the rank + they found and `SolveRRQR` and `Pinverse` answer through it rather + than failing. +- an exhausted iteration budget: `GMRES`, `SpSolve`, + `SpSolveBiCGSTAB`, `SpSolveComplexCG`, `SpSolveComplexBiCGSTAB`, + `SpLSQR` and `SpLSMR` return the error naming the residual achieved + with no estimate, and the two least-squares solvers report a + condition stop through `LeastSquaresInfo.Converged` instead. +- a modification that needs fill the factor does not hold: + `SparseCholesky.Update` and `Downdate` refuse it and leave the + factor as it was, and `CholeskyUpdate` refuses a rank-1 factor, a + mismatched vector or a complex input. +- a complex or wrong-length right-hand side to a sparse solve, and a + preconditioner built for another dimension. +- non-finite input to `NewSparseLU`, `NewSparseILU` and the sparse + Cholesky updates; the two general sparse eigensolvers instead + answer NaN Ritz pairs. + +#### Workflow + +The flagship sequence: one dense solve, the factorisation that serves +the right-hand sides after it, the spectral decompositions, and the +sparse route through a view, a factorisation and an iterative solve. +`a` and `b` are a square matrix and a matching right-hand side, `coo` +the same system as a `tensor.SparseCOO`. + +```go +func workflows(a, b *tensor.Array, coo *tensor.SparseCOO) ( + dense, factored, damped, sparseDirect, sparseIterative, piped *tensor.Array, err error) { + + // Dense: one call, the LU with partial pivoting. + dense, err = linalg.Solve(a, b) // x = A⁻¹·b + if err != nil { + return + } + + // A factorisation for reuse: a = L·Lᵀ, symmetric positive definite. + factored, err = linalg.Cholesky(a) + if err != nil { + return + } + + // The spectral decompositions. + _, _, err = linalg.Eigen(a) // ascending values, orthonormal columns + if err != nil { + return + } + _, _, _, err = linalg.SVD(a) // thin, singular values descending + if err != nil { + return + } + + // Ill-posed instead of square: damp the small singular directions. + damped, err = linalg.SolveTikhonov(a, b, 1e-3) + if err != nil { + return + } + + // Sparse: a factorisation under a fill-reducing ordering, then a solve. + f, err := linalg.NewSparseCholesky(coo, linalg.SparseOrderingReverseCuthillMcKee) + if err != nil { + return + } + sparseDirect, err = f.Solve(b) + if err != nil { + return + } + + // The same system iteratively, with an ILU(0) preconditioner. + ilu, err := linalg.NewSparseILU(coo) + if err != nil { + return + } + sparseIterative, err = linalg.SpSolve(coo, b, 0, 0, ilu) + if err != nil { + return + } + + // A sequence of element-wise steps read as one expression. + piped, err = linalg.Pipe(a).AddF(1).Sqrt().MulF(2).Result() + return +} +``` + +The same calls exist in the root package under one import: +`tensor.Solve`, `tensor.Cholesky`, `tensor.SpSolve` and the rest. + +### Package `signal` + +Transforms, spectra, filters and stencils. The package covers the Fourier, cosine, +sine and wavelet transforms, the spectral estimators, IIR filter design and +application, sample-rate conversion, the Hilbert envelope, convolution and pooling, +rank filters, spatial stencils, state-space estimation and time-series models. + +Every call takes whole arrays and returns whole arrays, and an array is an immutable +value, so no call modifies its input. The package holds no state between calls: a +design returns the direct-form coefficients `b` and `a`, and the caller hands those +to `FilterApply` for one pass or to `Filtfilt` for a zero-phase forward and backward +sweep. There is no streaming filter object, no transform plan to reuse, and no +batched or multi-channel entry point: `WelchPSD`, `STFT`, `Spectrogram`, the +resamplers, the filter application and the wavelet transforms take one rank-1 +signal per call, so a stack of channels is looped by the caller. The estimators and +resamplers also refuse a complex array; the Fourier transforms are the complex +path. The spectral estimators take `fs`, a sample rate in hertz, and so does every +filter design, whose cutoff or band edges must lie strictly inside `(0, fs/2)`. A +spectral estimate is one-sided unless the call says otherwise, and a wavelet +transform packs its coefficients as `[A_levels, D_levels, …, D_1]`, the deepest +approximation first and the finest detail last. + +#### Fourier transforms + +| Call | What it does | +|---|---| +| `FFT(a)` | Returns the forward DFT of a rank-1 array as a complex array of the same length, reading real, integer and complex input alike; refuses an empty array and any rank above 1. | +| `IFFT(a)` | The inverse of `FFT`, scaled by `1/n`; refuses an empty array and any rank above 1. | +| `FFT2(a)` | Returns the 2-D DFT of an `(H, W)` array as an `(H, W)` complex array, each row transformed then each column; refuses an empty array or a rank other than 2. | +| `IFFT2(a)` | The inverse 2-D DFT, scaled by `1/(H·W)`. | +| `FFT3(a)` | Returns the 3-D DFT over shape `(D, H, W)`; refuses a rank other than 3. | +| `IFFT3(a)` | The inverse 3-D DFT, scaled by `1/(D·H·W)`. | +| `FFTN(a, dims)` | Returns the N-D DFT along the dimensions in `dims`, or along every dimension when `dims` is empty; a dimension listed twice is transformed twice. Refuses out-of-range or zero-length dimensions. | +| `IFFTN(a, dims)` | The inverse N-D DFT over the same `dims`, dividing by the product of the transformed extents. | +| `RFFT(a)` | The real-input DFT: a rank-1 real array in, the non-redundant `n/2+1` complex half-spectrum out; refuses a complex array and an empty one. | +| `IRFFT(a, n)` | The inverse of `RFFT`: the half-spectrum back to a real signal of length `n`, where `n = 0` means `2·(len−1)` (or 1 for a lone DC bin). Refuses a real input, a spectrum that is not rank-1, and a length that does not match the spectrum. | +| `FFTFreq(n, d)` | Returns the `n` DFT sample frequencies for spacing `d` in the standard FFT order, `[0, 1/n, …, −1/2, …, −1/n]/d`, with `d = 0` read as 1 and a non-positive `n` answering an empty array. | + +#### Cosine and sine transforms + +The orthogonal transforms of types I to IV, all orthonormal, so +`IDCT(DCT(x, k), k)` and `IDST(DST(x, k), k)` restore `x`. + +| Call | What it does | +|---|---| +| `DCT(x, kind)` | Returns the orthonormal discrete cosine transform of type `kind` (1 to 4) of a rank-1 vector; refuses a kind outside 1 to 4, a complex or non-vector input, an empty vector, and a type I of fewer than two points. | +| `IDCT(x, kind)` | The inverse orthonormal DCT, the transpose partner of the forward transform of the same kind; refuses a kind outside 1 to 4 and the same input checks. | +| `DST(x, kind)` | Returns the orthonormal discrete sine transform of type `kind` (1 to 4) of a rank-1 vector, under the `DCT` input checks. | +| `IDST(x, kind)` | The inverse orthonormal DST of the same kind. | + +#### Non-uniform transform + +| Call | What it does | +|---|---| +| `NUFFTType1(x, c, n)` | Computes `f_k = Σ_j c_j·e^{2πi·k·x_j}` for `k = 0…n−1` by Gaussian gridding, so the answer carries the gridding error of a few digits rather than the exactness of the direct sum. `x` holds the non-uniform coordinates in `[−1/2, 1/2)` and `c` the complex values; refuses an out-of-range or NaN coordinate, mismatched lengths, a non-positive output size and non-real coordinates. | + +#### Spectral estimation + +| Call | What it does | +|---|---| +| `STFT(x, opts)` | Returns the complex frames of the short-time Fourier transform as a `(frames × segment)` array, row `t` holding the spectrum of `x[t·hop : t·hop+segment]` under the window, `hop = segment − overlap`; refuses a segment outside `[2, n]`, an overlap outside `[0, segment)`, a complex signal and an unknown window name. | +| `Spectrogram(x, fs, opts)` | Returns the one-sided power spectrogram as a `(frames × segment/2+1)` array in units of `x²/Hz`, each frame the periodogram of its windowed segment under the same doubling and window-power normalisation Welch's estimate uses, so averaging the frames reproduces `WelchPSD` on the same geometry; refuses a non-positive or infinite `fs` on top of the `STFT` checks. | +| `WelchPSD(x, fs, segment, overlap, window)` | Returns the frequency grid and the averaged one-sided power spectral density of a rank-1 real signal: `segment/2+1` bins from the `(n−overlap)/(segment−overlap)` segments the signal fills, the doubling applied away from DC and Nyquist, the window's power normalising the scale, so white noise of variance `σ²` estimates `σ²` across the band. Refuses a segment outside `[2, n]`, an overlap outside `[0, segment)`, a non-positive or infinite `fs`, a complex signal and an unknown window name. | +| `LombScargle(times, values, minFreq, maxFreq, nFreq)` | Returns the frequency grid and the normalised Lomb-Scargle periodogram of unevenly sampled observations over `nFreq` frequencies evenly spaced from `minFreq` to `maxFreq` inclusive, with the classical normalisation that peaks near `A²·n/(4·var(values))` for a pure sinusoid of amplitude `A`. Refuses a time base shorter than three points, a length mismatch, an all-equal time base, a non-positive variance, a complex array and a range that does not satisfy `0 < minFreq ≤ maxFreq`. | + +#### Poisson solvers + +| Call | What it does | +|---|---| +| `SolvePoissonPeriodic(f, lx, ly)` | Solves `−Δu = f` on the period square `[0, Lx] × [0, Ly]` on the grid carried by the shape of `f`, returning a float64 array of the same shape with the zero Fourier mode set to zero. Refuses a non-float64 `f`, a rank other than 2, a dimension below 2, a non-positive side and a nonzero mean. | +| `SolvePoissonDirichlet(f, lx, ly)` | Solves `−Δu = f` with `u = 0` on the whole boundary, over the interior of the grid; the boundary entries of `f` play no role. Refuses a non-float dtype, a grid under 3×3 and a non-positive side length. | +| `SolvePoissonNeumann(f, lx, ly)` | Solves `−Δu = f` with zero normal derivative on the whole boundary, every grid sample an unknown and the constant mode fixed to zero. Refuses a non-float dtype, a grid under 3×3, a non-positive side and a source whose trapezoidal-weighted sum does not vanish, the operator's own compatibility condition. | + +#### Filter design + +Each design returns the direct-form coefficients `b` (numerator) and `a` +(denominator, `a[0] = 1`) at a sample rate `fs`, with the passband edge or the +band edges prewarped to the bilinear axis and the answer exact at the mapped +frequencies. All of them refuse an order below 1, a non-positive or infinite `fs`, +and an edge outside `(0, fs/2)`; the band forms additionally require +`0 < edge1 < edge2 < fs/2`. + +| Call | What it does | +|---|---| +| `ButterworthLowPass(order, fs, cutoff)` | The maximally flat low-pass: −3 dB at `cutoff`, unity at DC. | +| `ButterworthHighPass(order, fs, cutoff)` | The high-pass mirror, sharing the low-pass poles with its zeros at `z = 1`: −3 dB at the same cutoff, unity at Nyquist. | +| `ButterworthBandPass(order, fs, edge1, edge2)` | The maximally flat band-pass spanning `edge1` to `edge2`; the prototype's order doubles through the band move. | +| `ButterworthBandStop(order, fs, edge1, edge2)` | The maximally flat band-stop over the same band. | +| `ChebyshevLowPass(order, fs, cutoff, rippleDB)` | The type I equiripple low-pass: the passband oscillates between 0 and `−rippleDB`, the edge is the last touch of `−rippleDB`; refuses a non-positive `rippleDB`. | +| `ChebyshevHighPass(order, fs, cutoff, rippleDB)` | The type I mirror above the edge. | +| `ChebyshevBandPass(order, fs, edge1, edge2, rippleDB)` | The type I band-pass; the order doubles through the band move. | +| `ChebyshevBandStop(order, fs, edge1, edge2, rippleDB)` | The type I band-stop. | +| `InverseChebyshevLowPass(order, fs, cutoff, stopbandDB)` | The type II low-pass: flat through the edge, the stopband bottoming out at `−stopbandDB` and equiripple beyond it; refuses a non-positive `stopbandDB`. | +| `InverseChebyshevHighPass(order, fs, cutoff, stopbandDB)` | The type II mirror above the edge. | +| `InverseChebyshevBandPass(order, fs, edge1, edge2, stopbandDB)` | The type II band-pass. | +| `InverseChebyshevBandStop(order, fs, edge1, edge2, stopbandDB)` | The type II band-stop. | +| `CauerLowPass(order, fs, cutoff, rippleDB, stopbandDB)` | The elliptic low-pass: equiripple within `rippleDB` in the passband and not above `−stopbandDB` in the stopband, with the narrowest transition of any design at the order and its zeros finite and on the unit circle. Refuses a `rippleDB` that is not positive or a `stopbandDB` not above it. | +| `CauerHighPass(order, fs, cutoff, rippleDB, stopbandDB)` | The elliptic mirror above the edge. | +| `CauerBandPass(order, fs, edge1, edge2, rippleDB, stopbandDB)` | The elliptic band-pass. | +| `CauerBandStop(order, fs, edge1, edge2, rippleDB, stopbandDB)` | The elliptic band-stop. | + +The designs hand back direct-form coefficients, and direct-form filtering loses +digits as the order climbs; past roughly order eight the caller is expected to +split the design into second-order sections. + +#### Filtering a signal + +| Call | What it does | +|---|---| +| `FilterApply(b, a, x)` | Runs the direct-form-II-transposed recursion `y[n] = Σ b_m·x[n−m] − Σ a_m·y[n−m]` over a rank-1 real signal, the output as long as the input. The coefficient lists may differ in length (the shorter is zero-padded), `a[0]` must be non-zero and is normalised away, float32 and float64 keep their dtype and other real dtypes widen to float64. Refuses a complex or higher-rank signal, an empty coefficient list and a zero `a[0]`. | +| `Filtfilt(b, a, x)` | Filters the rank-1 real signal forwards and backwards with the same transfer function and returns a result the length of `x` whose magnitude response is the square of the one-pass filter's and whose phase is zero, so a passband sinusoid comes out aligned with its input. Each end is extended by an even reflection of `pad = 3·(nfilt−1)` samples, the edge sample not repeated, and both pad regions are cropped away afterwards. Refuses a signal no longer than the pad, an empty coefficient list and a zero `a[0]`. | + +#### Windows + +Nine tapers, each in the symmetric (`periodic = false`) and the periodic +(`periodic = true`) convention: the symmetric form divides its argument by `n−1` +so the first and last samples coincide, the periodic form divides by `n` so the +sequence is one exact period of its underlying shape. Every builder returns a +fresh `n`-sample `[]float64`, every one refuses `n` below 1, and a one-sample +window is the single value 1 in both conventions. `WindowKaiser` additionally +refuses a `beta` that is negative, NaN or infinite. + +| Call | What it does | +|---|---| +| `WindowHann(n, periodic)` | The raised cosine `0.5 − 0.5·cos(2πx)`, zero at the edges in both conventions, −6 dB per octave sidelobe roll-off. | +| `WindowHamming(n, periodic)` | The raised cosine on a pedestal, `0.54 − 0.46·cos(2πx)`: the pedestal cancels Hann's first sidelobe at the cost of a floor the outer sidelobes never drop below. | +| `WindowBlackman(n, periodic)` | The exact Blackman window `0.42 − 0.5·cos(2πx) + 0.08·cos(4πx)`: sidelobes below −58 dB, a main lobe twice Hann's. | +| `WindowBlackmanHarris(n, periodic)` | The four-term Blackman-Harris window: sidelobes below −92 dB, a main lobe three Hann lobes wide, the choice when dynamic range matters more than resolution. | +| `WindowBartlett(n, periodic)` | The triangle `1 − \|2x − 1\|`, the Fejér kernel of the box, non-negative everywhere and zero at both edges. | +| `WindowKaiser(n, beta, periodic)` | The modified Bessel taper `I0(beta·sqrt(1 − r²))/I0(beta)` over `r = 2x − 1`, the adjustable compromise between main lobe width and sidelobe height; `beta` 0 is the box, near 5 the sidelobes sit around −30 dB, near 9 around −60 dB. | +| `WindowFlatTop(n, periodic)` | The five-term generalised cosine whose main lobe is flat to within a hundredth of a decibel, so a spectral line's amplitude reads true wherever it falls between bins; its edge samples are slightly negative, so it is for amplitude metrology, not for filtering. | +| `WindowCosine(n, periodic)` | The cosine (sine) window `sin(πx)`, one positive half-period. | +| `WindowBox(n, periodic)` | The untapered box: `n` ones. The `periodic` flag changes nothing here and exists only for signature uniformity. | + +`STFT` and `Spectrogram` read an empty window name as `"hann"`; `WelchPSD` has no +default, so an empty name is an error there. All three accept only `"hann"`, +`"hamming"` and `"box"`, resolved through the periodic forms of the catalogue. + +#### Generators + +| Call | What it does | +|---|---| +| `Chirp(n, f0, f1, rate)` | `n` samples of a linear frequency sweep from `f0` through `f1` at `rate` samples per unit of time, reaching `f1` exactly at the last sample. Every phase comes from the closed form at that sample's own time, never from a recursive oscillator, so no drift accumulates and any single sample stands alone. Both sweep edges must stay strictly below the Nyquist frequency `rate/2`, magnitudes counted for negative edges. Refuses `n` below 1, a non-positive or non-finite rate and a non-finite edge; the one-sample chirp is the single zero. | + +#### Resampling and the envelope + +| Call | What it does | +|---|---| +| `Decimate(data, factor, taps)` | Reduces the sample rate by the integer `factor` behind a Kaiser tapered FIR whose passband ends at nine tenths of the new Nyquist and whose stopband floor of about 80 dB is reached before it, compensating the filter's group delay and keeping every `factor`-th sample. `taps` sets the filter length, a non-positive value meaning `32·factor+1`; the result holds about `(n − taps)/factor` samples and is empty rather than read past the end when the tap count leaves no sample with full context. Refuses a `factor` below 2, an empty series and a tap count at or above the length. | +| `Resample(data, up, down, taps)` | Converts the rate by the rational factor `up/down`, filtering on the up-sampled grid at the tighter Nyquist with the up gain folded in and keeping every `down`-th compensated sample. `taps` is the kernel length on the up-sampled grid, a non-positive value meaning `32·max(up, down)+1`. Refuses the identity `1/1`, a factor below 1, an empty series and a tap count above `up·n`. | +| `ResampleFourier(data, size)` | Resamples to exactly `size` samples by the Fourier (band-limited) definition: the spectrum's bins are kept and zero-padded or truncated at the fold, scaled so the amplitudes carry over. Exact for series band-limited below the new Nyquist and treats the input as one period, so energy past the new Nyquist is lost. Refuses a non-vector or empty series, a complex one, and a target size below 1. | +| `AnalyticSignal(data)` | Builds the analytic signal `z = x + i·H(x)`: negative frequencies removed, positive ones doubled, the DC and Nyquist bins left alone. The Fourier definition treats the series as one period, so it is exact for an integer number of tones. Refuses a complex or empty rank-1 signal. | +| `Envelope(data)` | The instantaneous amplitude of a series, the modulus of its analytic signal: for a narrowband series this traces the curve a peak detector would find, without the smoothing lag. | + +#### Wavelets + +The discrete transforms pack `[A_levels, D_levels, …, D_1]` and the Haar pair +conserves energy exactly. + +| Call | What it does | +|---|---| +| `DWT(x, levels)` | Returns the Haar transform of a rank-1 real signal over `levels` scales with the periodic boundary; `levels` must be at least 1 and at most `log2(n)`, and Parseval holds for the orthonormal Haar basis. | +| `IDWT(coef, levels)` | Inverts `DWT` over the same level count and layout, widening the coefficients through the float accessor, so every real dtype `DWT` accepts inverts here as well. | +| `DaubechiesDWT(x, family, levels, mode)` | Returns the `dbN` transform over `levels` scales in the same packed layout. `DWTPeriodic` refuses a length that does not leave every live block a multiple of the `2N`-tap filter; `DWTZeroPad` pads the tail to the next such length first, so any length transforms. The periodic transform conserves energy exactly, a zero-padded one conserves the energy of the padded signal. `levels` must be at least 1. | +| `DaubechiesIDWT(coef, family, levels, mode)` | Inverts `DaubechiesDWT` over the same family, level count and mode, undoing the packed layout deepest band first. The coefficient vector must satisfy the length contract the forward transform produced; the mode argument gates exactly that contract, since the inverse bank itself is the periodic one. | +| `CWT(x, wavelet, scales, dt)` | Returns the continuous wavelet transform of a rank-1 real signal sampled at spacing `dt` as a `(len(scales), n)` complex array, row `k` the transform at `scales[k]`. The convolution runs through the FFT zero-padded to `2n` and the wavelet is L1-normalised per scale, so amplitudes stay comparable across scales. Morlet yields the full complex transform, MexicanHat a real result with zero imaginary part. Refuses a signal that is not rank-1 real or is empty, a `dt` that is not positive and finite, an empty scale list, a scale that is not positive and finite, and a `CWTWavelet` other than the two constants. | + +#### Correlation and time series + +| Call | What it does | +|---|---| +| `Autocorrelate(x, maxLag)` | Returns the sample autocorrelation at lags `0…maxLag` with the mean removed and every value normalised by the lag-zero sum, the biased estimator the Durbin-Levinson recursion and the Bartlett-band theory are written for. `maxLag` must lie in `[0, n−1]`; the result starts at 1 unless the signal is constant, whose zero power makes every lag 0, lag zero included. | +| `CrossCorrelate(x, y)` | Returns the raw cross-correlation of two equal-length real signals at every lag as `2n−1` entries, index `i` carrying lag `i−(n−1)` and the value `Σ_t x[t+lag]·y[t]` over the overlapping `t`. No mean removal and no normalisation: correlation as the inner product of shifted copies. | +| `PartialAutocorrelate(x, maxLag)` | Returns the partial autocorrelation at lags `1…maxLag` by the Durbin-Levinson recursion over the biased autocorrelation, `pacf[k]` being the last coefficient of the order-`k` AR fit. `maxLag` must lie in `[1, n−2]`. | +| `EstimateAR(x, order)` | Fits the order-`p` autoregression `x_t − μ = Σ_j φ_j·(x_{t−j} − μ) + e_t` by the Yule-Walker equations over the biased autocovariance, solved by Durbin-Levinson. The result carries the coefficients, the innovation variance, the concentrated Gaussian likelihood over the whole series and both information criteria. Refuses an order below 1, a series too short for the lag, non-finite or complex data, and autocovariances that do not describe a stationary process. | +| `EstimateARMA(x, p, q, opts)` | Fits `ARMA(p, q)` by the Hannan-Rissanen innovations method: an order-`m` proxy autoregression whose residuals stand in for the unobserved innovations, then one least-squares regression of the demeaned series on its own lags and the lagged residuals. A `q` of 0 is routed to `EstimateAR`, which solves the pure AR case exactly. The estimates carry more sampling noise than a maximum-likelihood fit would, so tolerances set against them should be looser. Refuses negative orders, `p + q = 0`, a proxy order below `p+q+1`, a criterion that is neither `"aic"` nor `"bic"` and a series too short to leave a regression window worth fitting. | +| `SelectARMA(x, maxAR, maxMA, opts)` | Searches the order grid `0…maxAR × 0…maxMA` for the cell the chosen criterion ranks best, fitting every cell with `EstimateARMA`, skipping the `(0, 0)` cell and any cell whose fit fails, and reporting the last failure when nothing on the grid could be fitted. Refuses negative grid bounds and a grid that holds no candidate. | +| `ARMASpectrum(res, nFreq)` | Returns the fitted model's theoretical one-sided spectrum `σ²·\|Θ(e^{−2πif})\|² / \|Φ(e^{−2πif})\|²` at `nFreq` frequencies evenly spaced from 0 to the Nyquist frequency 0.5 inclusive, on the same convention `WelchPSD` prints, so model and data line up bin for bin at `fs = 1` without a scale fudge. Refuses fewer than two frequencies, a nil model, a non-positive innovation variance or a non-finite coefficient, and names the frequency of a pole of `Φ` that sits on the unit circle instead of answering infinities. | + +#### State-space estimation + +| Call | What it does | +|---|---| +| `KalmanFilter(z, transition, observation, opts)` | Runs the linear Kalman filter over the measurement stack `z` under `x_{t+1} = F·x_t + w_t`, `z_t = H·x_t + v_t`, predicting with `F` and correcting with a Joseph-form update, and accumulating the exact Gaussian log-likelihood of the innovations. A rank-1 `z` holds `n` scalar observations, a rank-2 `z` holds `n` rows of `m` channels. Refuses a non-square transition, a shape mismatch, an asymmetric noise covariance, a singular measurement noise, a non-finite value anywhere and an innovation covariance that loses positive definiteness, naming the step. | +| `ExtendedKalmanFilter(z, transition, observation, opts)` | The same recursion on the nonlinear model `x_{t+1} = f(x_t) + w_t`, `z_t = h(x_t) + v_t`: the mean propagates through `f`, the covariance through the linearisation `F = ∂f/∂x`, and the correction uses `H = ∂h/∂x`. The Jacobians come from the options when supplied and from central differences otherwise. Being the linear filter on local linear models, it is blind to the curvature of `f` and `h`, so a strongly bent observation map wants the unscented filter. Refuses a state dimension it cannot fix from the options and reports a failing callback or Jacobian together with the step it failed at. | +| `UnscentedKalmanFilter(z, transition, observation, opts)` | Carries the state distribution through `f` and `h` to second order by the deterministic sigma-point set `x̂ ± sqrt(d+λ)·L[:, i]`, `L` the Cholesky factor of `P`, whose weighted moments reconstruct the predicted mean and covariance. The update carries no Joseph form, so `P` leaves it as `P⁻ − K·S·Kᵀ` mirrored into its symmetric average, and the next predict's Cholesky factorisation is what enforces positive definiteness, naming the step when it fails. The log-likelihood accumulates exactly as in the linear filter, and on a linear model the sigma transforms are exact, so the filter degenerates to the Kalman answer to rounding. | + +#### Rank filters and smoothing + +| Call | What it does | +|---|---| +| `MedianFilter(x, window)` | Returns the running median of a rank-1 signal over an odd window, removing isolated spikes whole while holding monotone ramps and edges still. The window must be odd, at least 3 and no longer than the signal, and a NaN in a window propagates into that output sample. | +| `MedianFilter2D(img, window)` | The same over a square odd window of a rank-2 image, the standard impulse-noise cleaner: salt-and-pepper dots vanish while steps between regions keep their corners. The window must be odd, at least 3 and fit both dimensions. | +| `RankFilter(x, window, k)` | Returns the `k`-th order statistic (ascending, 0-based) of each window of a rank-1 signal: rank 0 is the running minimum, `window−1` the running maximum, the middle rank the median `MedianFilter` takes. Windows near a boundary are truncated to the samples that exist, so the output is complete without padding, and the requested rank scales to the truncation. The window must be odd, at least 3, no longer than the signal, and `k` must lie in `[0, window)`. | +| `RankFilter2D(img, window, k)` | The same over a square window of a rank-2 image: rank 0 is erosion, `window²−1` is dilation, the middle rank the median. The window must be odd and at least 3, must fit both dimensions, and `k` must lie in `[0, window²)`. | +| `SavitzkyGolay(data, window, order)` | Smooths a rank-1 signal with a Savitzky-Golay polynomial filter, `window` the odd number of samples per fit and `order` the polynomial degree; a polynomial of degree at most `order` passes through unchanged. Every edge point owns a truncated window fitted at the sample's own position, so the result is as long as the input without padding. Refuses a negative order, an order above `window/2` (the truncated edge fits would be underdetermined), a window that is not odd or is below 3, a window longer than the signal, a complex signal and weights that overflow. | + +#### Stencils + +| Call | What it does | +|---|---| +| `Gradient1D(y, dx)` | Returns the central-difference derivative of a rank-1 uniform grid signal of spacing `dx`, interior points on the second-order stencil and the two endpoints on the first-order one-sided stencil, so the result is as long as the input. Refuses a rank other than 1, fewer than two points, a zero spacing and a complex signal. | +| `Laplacian(a, spacings…)` | Returns the second-derivative Laplacian of a grid signal: one spacing for rank 1, two (`dx`, `dy`) for a row-major rank-2 grid on the 5-point stencil, three for rank 3 on the 7-point stencil. Boundary points copy their nearest interior value, the boundary condition being the caller's. Refuses a rank outside 1 to 3, the wrong spacing count, a zero spacing and a degenerate extent (every axis at least 2 points, every axis of a 2-D or 3-D grid at least 3). | + +#### Convolutions + +| Call | What it does | +|---|---| +| `Conv1D(input, kernel, bias, stride, padding, dilation)` | The 1-D forward pass: `(N, C_in, L)` input, `(C_out, C_in, kL)` kernel, `(N, C_out, L_out)` output, NCL layout, `bias` optional and applied per output channel. Refuses a complex input or kernel, a channel mismatch, a rank other than 3, a non-positive stride, a negative padding, a kernel that does not fit, an empty output and a dilation other than 1, which is not implemented. | +| `Conv2D(input, kernel, bias, stride, padding)` | The 2-D forward pass: `(N, C_in, H, W)` input, `(C_out, C_in, kH, kW)` kernel, `(N, C_out, H_out, W_out)` output. The call is `Conv2DGroups` with `groups = 1`, under the same refusals. | +| `Conv2DGroups(input, kernel, bias, stride, padding, groups)` | The same pass with `groups`: a grouped convolution for `groups > 1` and a depthwise one when `groups == C_in`, the kernel's in-channels being `C_in/groups` per group. Refuses a group count below 1, an input or output channel count not divisible by it, a kernel in-channel count that does not match `C_in/groups`, a complex input or kernel, a rank other than 4, a non-positive stride, a negative padding, a kernel that does not fit and an empty output. | +| `Conv3D(input, kernel, bias, stride, padding, dilation)` | The 3-D forward pass: `(N, C_in, D, H, W)` input, `(C_out, C_in, kD, kH, kW)` kernel, `(N, C_out, D_out, H_out, W_out)` output, NCDHW layout. `padding` and `dilation` are per-axis triples; a dilation other than 1 is refused, not implemented. | +| `ConvTranspose2D(input, kernel, bias, stride, padding)` | The transposed convolution: `(N, C_in, H, W)` input, `(C_in, C_out, kH, kW)` kernel, `(N, C_out, H_out, W_out)` output with `H_out = (H−1)·stride − 2·padding + kH` and the same for `W`. Refuses a complex input, a rank other than 4, a kernel in-channel count that does not match the input, a non-positive stride, a negative padding and an empty output. | + +#### Pooling + +The maxima propagate a NaN in a window; the averages divide by the full kernel +size when `countIncludePad` is true and by the number of non-padded elements +otherwise. Padding is refused at or above the kernel, where a window would hold +no data. + +| Call | What it does | +|---|---| +| `MaxPool1D(input, kernel, stride, padding)` | The maximum over each window of an `(N, C, L)` tensor. | +| `MaxPool2D(input, kernel, stride, padding)` | The maximum over each window of an `(N, C, H, W)` tensor. | +| `MaxPool3D(input, kernel, stride, padding)` | The same over a `(N, C, D, H, W)` tensor, `kernel`, `stride` and `padding` being per-axis triples. | +| `AvgPool1D(input, kernel, stride, padding, countIncludePad)` | The average over each window of an `(N, C, L)` tensor. | +| `AvgPool2D(input, kernel, stride, padding, countIncludePad)` | The average over each window of an `(N, C, H, W)` tensor. | +| `AvgPool3D(input, kernel, stride, padding, countIncludePad)` | The same over a `(N, C, D, H, W)` tensor with per-axis triples. | +| `AdaptiveMaxPool1D(input, outputL)` | Pools an `(N, C, L)` tensor to `outputL` by taking the maximum over the window from `floor(o·L_in/L_out)` to `ceil((o+1)·L_in/L_out)`, so every input sample lands in some window. | +| `AdaptiveMaxPool2D(input, outputH, outputW)` | The same floor-start, ceil-end convention on an `(N, C, H, W)` tensor; both output sizes must be at least 1. | +| `AdaptiveMaxPool3D(input, outputD, outputH, outputW)` | The same on a `(N, C, D, H, W)` tensor. | +| `AdaptiveAvgPool1D(input, outputL)` | The average over the same windows of an `(N, C, L)` tensor. | +| `AdaptiveAvgPool2D(input, outputH, outputW)` | The average over the same windows of an `(N, C, H, W)` tensor. | +| `AdaptiveAvgPool3D(input, outputD, outputH, outputW)` | The same on a `(N, C, D, H, W)` tensor. | +| `GlobalMaxPool1D(input)` | Reduces an `(N, C, L)` tensor to `(N, C, 1)`. | +| `GlobalMaxPool2D(input)` | Reduces an `(N, C, H, W)` tensor to `(N, C, 1, 1)`, the global maximum of a feature map. | +| `GlobalMaxPool3D(input)` | Reduces a `(N, C, D, H, W)` tensor to `(N, C, 1, 1, 1)`. | +| `GlobalAvgPool1D(input)` | Reduces an `(N, C, L)` tensor to `(N, C, 1)`. | +| `GlobalAvgPool2D(input)` | `AdaptiveAvgPool2D` reduced to `(1, 1)`, the global average pooling of modern CNNs. | +| `GlobalAvgPool3D(input)` | Reduces a `(N, C, D, H, W)` tensor to `(N, C, 1, 1, 1)`. | + +Every one of them refuses a complex array, the rank that does not match its +dimensionality, an empty spatial dimension, a kernel or stride below 1, a negative +padding, a padding at or above the kernel, a window that does not fit the extent +and an output of zero extent; the two adaptive families additionally refuse an +output size below 1. + +#### Accumulation + +| Call | What it does | +|---|---| +| `SumKahan(a)` | Returns the sum of all elements by Kahan's compensated summation, which carries more of the small terms than a plain left-to-right accumulation when the magnitudes are mixed. Refuses a complex array. | + +#### Options and results + +**`STFTOptions`** tunes `STFT` and `Spectrogram`. The frame count follows from +the geometry: `(n − overlap)/(segment − overlap)`, at least one. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Segment` | `int` | required, must lie in `[2, n]` | The length of one frame in samples. A value outside that range is an error. | +| `Overlap` | `int` | `0` | The samples two neighbouring frames share; `segment/2` is the usual choice. Must lie in `[0, segment)`. | +| `Window` | `string` | `"hann"` for `""` | The taper: `"hann"`, `"hamming"` or `"box"`; any other name is an error. | + +**`ARMAOptions`** tunes the Hannan-Rissanen estimation and the `SelectARMA` +grid search. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `HighAROrder` | `int` | `0` means `max(16, p+q+8)`, bounded by the data length | The order of the proxy autoregression whose residuals supply the MA stage; below `p+q+1` it is an error. | +| `Criterion` | `string` | `""` means `"aic"` | The information criterion `SelectARMA` minimises: `"aic"` or `"bic"`. Any other value is an error. | + +**`ARMAResult`** is one fitted time-series model, in the convention +`x_t − μ = Σ_j φ_j·(x_{t−j} − μ) + e_t + Σ_k θ_k·e_{t−k}`, the innovations white +with variance `InnovationVariance`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `AR` | `[]float64` | empty only for a pure MA model | The coefficients `φ_1…φ_P`. | +| `MA` | `[]float64` | empty only for a pure AR model | The coefficients `θ_1…θ_Q`. | +| `InnovationVariance` | `float64` | set by the fit | The estimated variance `σ²` of `e_t`. A fit whose innovation variance does not come out positive is refused rather than returned. | +| `Mean` | `float64` | set by the fit | The sample mean the fit removed and reports back. | +| `LogLikelihood` | `float64` | set by the fit | The concentrated Gaussian likelihood of the innovations over the whole series, `−n/2·(log 2πσ² + 1)`, so the criteria of different models on one series sit on a common footing. | +| `AIC` | `float64` | set by the fit | `−2·LogLikelihood + 2k` with `k = P + Q + 1`, the variance counted as a parameter. | +| `BIC` | `float64` | set by the fit | `−2·LogLikelihood + k·log n`, the same `k` against the innovation count's logarithm. | +| `P`, `Q` | `int` | set by the fit | The orders the result carries. | + +**`KalmanOptions`** carries the initial condition, the noise levels and the +filter-specific knobs. An unset array field takes its default; the extended and +unscented filters cannot infer the state dimension from their callbacks, so at +least one of `InitialState` and `InitialCovariance` must be set for them, and the +initial covariance must be symmetric positive definite. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `InitialState` | `*core.Array` | nil means a zero state of the inferred dimension | The initial mean `x̂_0`; must be a vector of one fixed length, finite. | +| `InitialCovariance` | `*core.Array` | nil means the identity | The initial covariance `P_0`, a symmetric positive-definite `d×d` matrix. | +| `ProcessNoise` | `*core.Array` | nil means zeros | The process noise covariance `Q`, a finite symmetric `d×d` matrix. | +| `MeasurementNoise` | `*core.Array` | nil means the identity | The measurement noise covariance `R`, a finite symmetric `m×m` matrix; it is factored every step, so a singular one is refused. | +| `TransitionJacobian` | `JacobianFunc` | nil means central differences, two evaluations per partial per step | The analytic `∂f/∂x` of the extended filter's transition. | +| `ObservationJacobian` | `JacobianFunc` | nil means central differences | The analytic `∂h/∂x` of the extended filter's observation. | +| `SigmaAlpha` | `float64` | `0` means `0.001` | The unscented spread `α`, read by `UnscentedKalmanFilter` only; a negative or NaN value, or a combination with `SigmaKappa` that leaves no positive scale, is an error. | +| `SigmaBeta` | `float64` | `0` means `2` | The unscented prior correction `β`, the standard 2 for a Gaussian prior; a negative or NaN value is an error. | +| `SigmaKappa` | `float64` | `0` | The unscented secondary scaling `κ`; NaN is an error. | + +**`KalmanResult`** is one filtering pass over the measurement stack, `n` +measurements of `m` channels and a state of dimension `d`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `States` | `*core.Array` | set by the pass | The `(n × d)` stack of filtered means `x̂_{t|t}`, one row per measurement. | +| `Covariances` | `*core.Array` | set by the pass | The `(n × d × d)` stack of filtered covariances `P_{t|t}`, one symmetric positive-definite block per measurement. | +| `Innovations` | `*core.Array` | set by the pass | The `(n × m)` stack of one-step prediction errors `z_t − h(x̂_{t|t−1})`. | +| `InnovationCovariances` | `*core.Array` | set by the pass | The `(n × m × m)` stack of the innovation covariances `S_t` the likelihood reads. | +| `LogLikelihood` | `float64` | set by the pass | `Σ_t log N(z_t; h(x̂_{t|t−1}), S_t)`, the exact Gaussian likelihood of the measurement sequence under the model and the quantity noise and parameter estimation maximises. | + +**`CWTWavelet`** names the analysing wavelet of `CWT`. The zero value names no +wavelet; pass one of the two constants. + +| Constant | Value | Meaning | +|---|---|---| +| `Morlet` | `"morlet"` | The complex Morlet with `ω₀ = 5`, the time-frequency standard. | +| `MexicanHat` | `"mexicanhat"` | The real Ricker wavelet, the zero-mean second derivative of a Gaussian. | + +**`DWTMode`** picks the boundary treatment of `DaubechiesDWT` and +`DaubechiesIDWT`. The zero value is `DWTPeriodic`; a value that is neither +constant is an error. + +| Constant | Value | Meaning | +|---|---|---| +| `DWTPeriodic` | `0` | Treats the signal as one period of a periodic sequence. Every block keeps its exact energy, and the length must offer every level a multiple of the filter length to work on. | +| `DWTZeroPad` | `1` | Extends the signal with zeros at the tail to the next length the level tree needs, so any length transforms; the coefficients past the signal's own span carry the response of the step to zero at the seam. | + +**`Daubechies`** names a member of the Daubechies wavelet family for +`DaubechiesDWT` and `DaubechiesIDWT`; `dbN` carries `N` vanishing moments and a +`2N`-tap filter pair. Haar is the `db1` case and keeps its own exact transform in +`DWT` and `IDWT`. Any other value is an error. + +| Constant | Value | Meaning | +|---|---|---| +| `DB2` | `"db2"` | The 4-tap Daubechies wavelet with 2 vanishing moments. | +| `DB3` | `"db3"` | The 6-tap wavelet with 3 vanishing moments. | +| `DB4` | `"db4"` | The 8-tap wavelet with 4 vanishing moments. | +| `DB5` | `"db5"` | The 10-tap wavelet with 5 vanishing moments. | +| `DB6` | `"db6"` | The 12-tap wavelet with 6 vanishing moments. | +| `DB7` | `"db7"` | The 14-tap wavelet with 7 vanishing moments. | +| `DB8` | `"db8"` | The 16-tap wavelet with 8 vanishing moments. | + +**`StateFunc`** is the type `func(x *core.Array) (*core.Array, error)`, the shape +of a model callback: the `f` of `x_{t+1} = f(x_t) + w_t` or the `h` of +`z_t = h(x_t) + v_t`. The input array is private to the call and may be kept +until the call returns; the output must be a real rank-1 array of one fixed +length. + +**`JacobianFunc`** is the type `func(x *core.Array) (*core.Array, error)`, +returning the Jacobian of a `StateFunc` at `x` as an `(m × n)` float array, row +`i` holding the partials of output `i`. It is the type of +`KalmanOptions.TransitionJacobian` and `KalmanOptions.ObservationJacobian`; an +analytic Jacobian is always the better instrument, and a nil one is replaced by +central differences. + +#### Errors + +Every error carries the library's `tensor: ` prefix and names the call. + +- A rank or shape the call does not take: refused with the shape it received + (`needs a 1-D array, got shape (2, 3)`, `input must be 4-D (N, C_in, H, W)`, + `the signal must be a vector`). +- An empty input where a transform needs samples: refused (`an empty array has no + FFT`, `the input must not be empty`, `the signal must not be empty`). +- A complex array where the call reads real values: refused by the estimators, + the resamplers, the filter application, the pooling and convolution kernels and + the Kalman filters. +- A transform parameter outside its range: an order below 1, a `fs` that is not + positive and finite, an edge at or beyond `fs/2`, a band whose second edge does + not exceed its first, a non-positive `rippleDB` or `stopbandDB`, a `stopbandDB` + not above `rippleDB`, a DCT or DST kind outside 1 to 4. +- A window parameter that cannot be honoured: a `segment` outside `[2, n]`, an + `overlap` outside `[0, segment)`, a window name that is not `"hann"`, + `"hamming"` or `"box"`, an odd-window filter given an even, below-3 or + oversized window, a rank outside `[0, window)` (or `[0, window²)`). +- A filter that cannot run: an empty coefficient list or `a[0] = 0` in + `FilterApply` and `Filtfilt`, a complex or non-vector signal, and a signal no + longer than `3·(nfilt−1)` in `Filtfilt`, which leaves nothing to return once + both transient regions are excluded. +- A resampling that cannot stand: a `factor` below 2 in `Decimate`, an identity + `up/down` in `Resample`, and a tap count that leaves nothing after the filter + delay. +- A wavelet length the mode cannot serve: `DWTPeriodic` refuses a length that does + not leave every live block a multiple of the `2N`-tap filter, `DWT` and + `DaubechiesDWT` refuse a `levels` below 1, and `DWT` refuses one above + `log2(n)`. +- A solve that has no solution: a nonzero mean in `SolvePoissonPeriodic`, a + nonzero trapezoidal mean in `SolvePoissonNeumann`. +- A time-series fit that carries no information: an order below 1, an order above + what the series can support, a non-positive innovation variance, a series too + short for the Hannan-Rissanen regression window, a criterion that is neither + `"aic"` nor `"bic"`, and a grid where no model could be fitted (the last + failure is reported). +- A Kalman pass that cannot continue: a non-square transition, a shape mismatch, + an asymmetric noise covariance, a singular measurement noise, an innovation + covariance that loses positive definiteness (naming the step), a failing model + callback or Jacobian (naming the step) and a state dimension the options do not + fix. + +A numeric overflow in the Savitzky-Golay normal equations, a pole of a fitted +ARMA model's `Φ` on the unit circle in `ARMASpectrum`, and a coordinate outside +`[−1/2, 1/2)` in `NUFFTType1` are refused the same way, each naming the value. + +#### Workflow + +Design a low-pass, filter a two-tone series both ways for a zero-phase result, and +read the spectrum of the answer back out. + +```go +package main + +import ( + "fmt" + "log" + "math" + + "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/signal" +) + +// A 5 Hz tone a design must pass, a 45 Hz tone it must attenuate: one +// filter pass lags the tone, the forward-and-backward sweep does not. +func main() { + const n, fs, pass, stop = 400, 100.0, 5.0, 45.0 + raw := make([]float64, n) + for i := range raw { + t := float64(i) / fs + raw[i] = math.Sin(2*math.Pi*pass*t) + 0.4*math.Sin(2*math.Pi*stop*t) + } + x, err := tensor.FromFloats(raw, n) + if err != nil { + log.Fatal(err) + } + + // The design: order 4, sample rate 100 Hz, cutoff 20 Hz. b and a + // are the direct-form coefficients, a[0] = 1. + b, a, err := signal.ButterworthLowPass(4, fs, 20) + if err != nil { + log.Fatal(err) + } + + // One pass, then the zero-phase two-pass sweep. + once, err := signal.FilterApply(b, a, x) + if err != nil { + log.Fatal(err) + } + y, err := signal.Filtfilt(b, a, x) + if err != nil { + log.Fatal(err) + } + + // The residual against a pure 5 Hz reference: the one-pass result + // carries its phase lag, the zero-phase sweep only what the + // stopband let through. + leak := func(v *tensor.Array) float64 { + worst := 0.0 + for i := range v.Len() { + ref := math.Sin(2 * math.Pi * pass * float64(i) / fs) + worst = math.Max(worst, math.Abs(v.FloatAt(i)-ref)) + } + return worst + } + fmt.Printf("residual, one pass: %.4f\n", leak(once)) + fmt.Printf("residual, filtfilt: %.4f\n", leak(y)) + + // The spectrum of the filtered series, averaged over Hann segments + // of 100 samples: at fs = 100 the bin spacing is 1 Hz, so the 5 Hz + // tone has a bin of its own. + freqs, psd, err := signal.WelchPSD(y, fs, 100, 50, "hann") + if err != nil { + log.Fatal(err) + } + peak := 0 + for k := range psd.Len() { + if psd.FloatAt(k) > psd.FloatAt(peak) { + peak = k + } + } + fmt.Printf("peak of the filtered spectrum: %.0f Hz over %d bins\n", + freqs.FloatAt(peak), psd.Len()) +} +``` + +### Package `integrate` + +Differential equations, quadrature, PDE evolution and finite elements. +The package keeps no state between calls: every solver takes the problem as +arguments and returns the answer. + +It deliberately stops at the point where a solver would have to guess. It +offers no dense-output object (an event time is narrowed by re-integrating +the accepted step, and `IntegrateODEPath` and `IntegrateODESteps` return +exactly the states they name), no adaptive step size in the symplectic +family (adaptivity would destroy the property those methods exist for), no +projection of an inconsistent start onto a DAE's constraint manifold and no +order above one there, and no element order above P1 or mesh that is not +conforming. The gradient of a trajectory with respect to its parameters is +the `grad` package's to compute. + +An ordinary differential equation is `y' = f(t, y)` with `f` returning the +derivative of the state. The state is rank 1, and the ODE family widens its +elements 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. Returned trajectories are +freshly allocated `float64` arrays, and inputs are never written to. + +#### Initial value problems + +| Call | What it does | +|---|---| +| `IntegrateODE(f, t0, t1, y0, opts)` | Integrates to `t1` with the adaptive Dormand-Prince 4(5) pair and returns `y(t1)`; the default route for a non-stiff problem. | +| `IntegrateRK4(f, t0, t1, y0, steps)` | The classical fixed-step fourth-order Runge-Kutta scheme over `steps` equal steps; returns `y(t1)`. No options, no adaptivity. | +| `IntegrateBackwardEuler(f, t0, t1, y0, steps, opts)` | Fully implicit Euler over `steps` equal steps, each solved by Newton over a numerical Jacobian and the library's LU solver; the entry-level stiff scheme. | +| `IntegrateBDF2(f, t0, t1, y0, opts)` | Variable-step BDF2; second order with adaptive step control, for stiff problems where an explicit pair is pinned by the fast mode. | +| `IntegrateBDFVar(f, t0, t1, y0, opts)` | Variable-step, variable-order BDF of orders one through five, with the run's counters reported through `BDFVarOptions.Stats`; the stiff workhorse. | +| `IntegrateROS4(f, t0, t1, y0, opts)` | The four-stage L-stable Rosenbrock-Wanner scheme ROS4, with the Jacobian taken by central differences once per step. The fourth order holds for autonomous systems: a genuinely time-dependent `f` loses the second-order local terms the tableau's time-derivative weights would carry and can degrade to first order globally, while the adaptive controller still holds its tolerance at the cost of more steps. | +| `IntegrateODEPath(f, t0, t1, y0, nSamples, opts)` | Returns the trajectory sampled at `nSamples` evenly spaced times, endpoints included; each interval is integrated on its own, so the step control never has to align with the sampling grid. | +| `IntegrateODESteps(f, t0, t1, y0, opts)` | Returns the trajectory at every accepted solver step, starting at `(t0, y0)` and ending at `(t1, y(t1))`; those nodes are where the local error was judged within tolerance, which is what post-processing and adjoint passes want. | + +Every solver in this group accepts `t1 < t0` and integrates in the negative +direction, the step carrying the sign of the span. None of them returns a +truncated or silently wrong trajectory: an exhausted step budget, a collapsed +step size or an `f` that returns an array of the wrong shape is an error +naming the call, and the fixed-step `IntegrateRK4` additionally refuses a +non-finite derivative, having no step control to catch one. + +#### Events along a trajectory + +| Call | What it does | +|---|---| +| `IntegrateODEEvents(f, t0, t1, y0, watches, opts)` | Integrates exactly like `IntegrateODE` and, alongside the final state, returns every zero crossing of the watches, sorted by time, as `ODEEventHit` values. | + +A watch is an `ODEWatch`: a scalar function of `(t, y)` and a `Direction` +filter. Direction `0` records both crossings, `+1` only a rising one (the +watch going from negative to non-negative) and `-1` only a falling one. The +watch is evaluated at the boundaries of every accepted step, and the first +accepted step is seeded with the watch value at its start state, so a +crossing inside it is detected like any other; the crossing itself is then +narrowed by bisection that re-integrates that step's interval. Watches are +compared only across accepted-step boundaries, so a watch that touches zero +and returns to its sign inside one step goes unnoticed, and a sign change is +searched strictly after `t0`. + +#### Differential-algebraic equations + +| Call | What it does | +|---|---| +| `IntegrateDAE(f, m, t0, t1, y0, steps, opts)` | Integrates the mass-matrix system `M·y' = f(t, y)` over `steps` equal implicit Euler steps and returns `y(t1)`. | + +The mass matrix must be square and singular, with its rank deficiency +carried by whole zero rows matched by an equal count of whole zero columns: +the zero rows are the algebraic constraints, the zero columns the algebraic +variables. The scheme is first order at a fixed step, and consistent initial +values are the caller's contract: the solver verifies the initial residual on +the algebraic rows and refuses when it sits beyond the tolerance, but it does +not project a general start onto the constraint manifold. The index is +certified at `t0` by factoring the Jacobian of the algebraic rows against +the algebraic variables, which refuses an index-3 system such as the +Cartesian pendulum with multipliers. + +#### Boundary value problems + +| Call | What it does | +|---|---| +| `IntegrateBoundary(f, t0, t1, y0, bc, nSamples, opts)` | Solves `y' = f(t, y)` under `bc` by shooting and returns the trajectory at `nSamples` evenly spaced times; `states[0]` carries the initial state with the shooting unknowns replaced by the values that satisfy the end conditions. | +| `SolveBoundaryCollocation(f, t0, t1, y0, bc, opts)` | Solves the same problem by three-point Lobatto IIIA collocation on an adaptively refined mesh, returning a `CollocationSolution` with the mesh, the nodal states and the nodal slopes. | + +`BoundaryConditions` states which components are prescribed where: +`len(Start) + len(End)` must equal the state length, the `Start` values are +read from `y0`, the `End` values come from `EndValues`, and a component may +carry a condition at both ends. The remaining components are the shooting +unknowns. The shooting root find is local: a start whose basin holds no +matching trajectory, or a trial that blows up on the way to `t1`, reports +the failure. Backward integration works. + +#### Symplectic integrators + +| Call | What it does | +|---|---| +| `IntegrateVerlet(accel, t0, t1, q0, p0, steps)` | Integrates a separable Hamiltonian system with unit masses by velocity Verlet (kick-drift-kick) over `steps` equal steps, returning the positions and momenta at `t0 + s·h`. | +| `IntegrateYoshida4(accel, t0, t1, q0, p0, steps)` | The same contract at fourth order, by Yoshida's composition of three leapfrog sub-steps per step. | +| `IntegrateMidpoint(gradH, t0, t1, q0, p0, steps, opts)` | Integrates the general Hamiltonian flow `dz/dt = J·∇H(z)`, `z = (q, p)`, by the implicit midpoint rule, each step's implicit stage solved by `optim.FindRootSystem`. | + +`accel` returns the acceleration `-∂V/∂q` at a position and `gradH` returns +the stacked gradient `(∂H/∂q, ∂H/∂p)`. The step size stays fixed by design; +`positions[0]` is `q0` and `momenta[0]` is `p0`. + +#### Quadrature and cubature + +| Call | What it does | +|---|---| +| `IntegrateFunction(f, a, b, opts)` | Returns the definite integral of a scalar `f` over `[a, b]` and an estimate of the absolute error. Infinite bounds are accepted under a rational substitution, and a reversed interval (`a > b`) integrates in the negative direction. | +| `IntegrateFilon(f, a, b, k, opts)` | Returns `∫ f(x)·cos(kx) dx` and `∫ f(x)·sin(kx) dx` over `[a, b]` by a Filon-type scheme: the amplitude is interpolated through Gauss-Legendre nodes on each panel and the product with the carrier is carried exactly through per-panel weights, so the error tracks the smoothness of `f` alone and the cost does not grow with the frequency. A reversed interval negates both parts; `k = 0` degenerates to the plain integral with a zero sine part. | +| `GaussLegendreNodes(n)` | Returns the nodes (ascending) and weights of the `n`-point Gauss-Legendre rule over `[-1, 1]`, exact for polynomials up to degree `2n-1`; `n` must be between 1 and 128. | +| `IntegrateND(f, lower, upper, opts)` | Returns the integral of `f` over the hyperrectangle `[lower, upper]` element-wise, by globally adaptive bisection with product Gauss-Legendre rules (orders 3 and 5 per axis), always bisecting the worst box along its longest edge; no error estimate comes back, only the value. | + +The slices `GaussLegendreNodes` returns are a shared cache and must be +treated as read-only, because a write would poison every later quadrature run +on the same node count. Every sampled scheme has one blind spot: a feature +entirely inside the gaps of the first rule's nodes, say a peak far narrower +than `(b-a)/n`, produces small values everywhere it samples and is missed +with a small error estimate, so known sharp features belong in their own +`IntegrateFunction` calls. + +#### PDE evolution in one dimension + +| Call | What it does | +|---|---| +| `IntegrateHeat1D(u0, kappa, dx, tFinal, dt, samples, boundL, boundR)` | Evolves `u_t = κ·u_xx` over the interior grid of `u0` (`n = u0.Len()`, `dx = L/(n+1)`) with the Dirichlet ends `boundL` and `boundR`, by Crank-Nicolson in steps of at most `dt`; returns the `(samples, n)` array of interior states evenly spaced in time, endpoints included. | +| `IntegrateWave1D(u0, v0, c, dx, tFinal, dt, samples)` | Evolves `u_tt = c²·u_xx` with zero Dirichlet ends and the initial velocity `v0`, by velocity Verlet at fixed step `dt`; same return contract as `IntegrateHeat1D`. | +| `IntegrateUpwindAdvection1D(u0, a, dx, tFinal, dt, samples, boundL, boundR)` | Evolves `u_t + a·u_x = 0` with the plain first-order upwind flux: monotone under CFL ≤ 1 and diffuse, the baseline the limited scheme is measured against. | +| `IntegrateAdvection1D(u0, a, dx, tFinal, dt, samples, boundL, boundR)` | The same equation with the Koren-limited upwind flux: total variation diminishing under CFL ≤ 1, third order at smooth faces and first order next to the inflow boundary. | +| `IntegrateAdvectionDiffusion1D(u0, a, kappa, dx, tFinal, dt, samples, boundL, boundR)` | Evolves `u_t + a·u_x = κ·u_xx`: the limited advection flux advanced explicitly over the Crank-Nicolson diffusion step, first order in time and second in space. With `a = 0` it reduces exactly to `IntegrateHeat1D`. | + +Crank-Nicolson is stable for any `dt`, but accuracy wants `dt` of a few +`dx²/κ`; the wave and the advection solvers have genuine explicit CFL +budgets (`|c·dt/dx| ≤ 1` and `|a·dt/dx| ≤ 1`), which are enforced as errors +because the explicit stencils have no honest answer past them. + +#### PDE evolution on a rectangle + +| Call | What it does | +|---|---| +| `IntegrateHeat2D(u0, kappa, dx, dy, tFinal, dt, samples, boundBottom, boundTop, boundLeft, boundRight)` | Evolves `u_t = κ·Δu` on the rank-2 grid of `u0` by Peaceman-Rachford alternating direction implicit steps, unconditionally stable and second order in space and time; returns a `(samples, rows, cols)` array with the final state forced into the last sample. | +| `IntegrateWave2D(u0, v0, c, dx, dy, tFinal, dt, samples)` | Evolves `u_tt = c²·Δu` with the boundary ring held at zero, by the explicit central-difference stencil, with the CFL budget `c·dt·sqrt(1/dx² + 1/dy²) ≤ 1` enforced as an error; same return contract as `IntegrateHeat2D`. | + +The grid is the rank-2 shape of the initial state: row `r` samples +`y = r·dy` and column `c` samples `x = c·dx`, the boundary ring is held +fixed and the interior carries the dynamics. The grid must be at least 3×3 +to hold interior points. + +#### Meshes + +| Call | What it does | +|---|---| +| `GridTriangleMesh2D(x0, y0, width, height, m, n)` | Builds the structured triangulation of the axis-aligned rectangle with `m` by `n` cells, two triangles each; `m` and `n` must both be positive. | +| `NewTriangleMesh2D(vertices, triangles)` | Builds a mesh from a vertex table with two columns and a triangle table with three columns of vertex indices; an out-of-range index or a degenerate (collinear) triangle is an error. | +| `BoxTetraMesh3D(x0, y0, z0, width, height, depth, m, n, p)` | Builds the structured tetrahedralisation of the box with `m` by `n` by `p` cells, six positively oriented tetrahedra per cell (the Kuhn subdivision), conforming across cell faces. | +| `NewTetraMesh3D(vertices, tetrahedra)` | Builds a mesh from a vertex table with three columns and a tetrahedron table with four columns; a zero-volume or negatively oriented tetrahedron is an error naming the element and its vertices. | + +| Method | What it does | +|---|---| +| `mesh.Vertices2()` | The vertex count of a `TriangleMesh2D`. | +| `mesh.Triangles3()` | The triangle count of a `TriangleMesh2D`. | +| `mesh.BoundaryEdges()` | The boundary edges as flat pairs of vertex indices, sorted; an edge is on the boundary when exactly one triangle carries it. | +| `mesh.Vertices3()` | The vertex count of a `TetraMesh3D`. | +| `mesh.Tetrahedra4()` | The tetrahedron count of a `TetraMesh3D`. | +| `mesh.BoundaryFaces()` | The boundary faces as flat triples of vertex indices, sorted lexicographically; a face is on the boundary when exactly one tetrahedron carries it. | + +A triangle's orientation does not matter; a tetrahedron's does, because the +signed volume has to be positive. + +#### Finite element Poisson solvers + +| Call | What it does | +|---|---| +| `SolvePoissonFEM2D(mesh, f, opts)` | Solves `-∇·(κ∇u) = f` on the triangular mesh with P1 elements and returns the vertex values; `f` may be nil for the homogeneous equation. | +| `SolvePoissonFEM3D(mesh, f, opts)` | The tetrahedral counterpart of the same problem, on P1 elements over a `TetraMesh3D`. | + +Both assemble the stiffness matrix per element (the conductivity evaluated at +the centroids when it varies), integrate Neumann fluxes on the prescribed +boundary edges or faces, eliminate Dirichlet values by lifting, and hand the +reduced system to the sparse Cholesky factorisation in `linalg`. The 2-D load +is lumped at the vertices from `f` at the centroids; the 3-D load is +integrated per tetrahedron with the 3×3×3 collapsed Gauss rule, exact through +degree 5. At least one Dirichlet node is required, because a purely Neumann +problem has no unique solution. + +#### Options and results + +**`ODEOptions`**: tunes the ODE drivers that take it (`IntegrateODE`, +`IntegrateBackwardEuler`, `IntegrateBDF2`, `IntegrateROS4`, +`IntegrateODEPath`, `IntegrateODESteps`, `IntegrateODEEvents` and +`IntegrateBoundary`); the defaults are applied inside the solver when a field +is not positive. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `RelTol` | `float64` | `1e-6` | Relative part of the local error tolerance. Values `≤ 0` mean the default. | +| `AbsTol` | `float64` | `1e-9` | Absolute part of the local error tolerance. Values `≤ 0` mean the default. | +| `MaxSteps` | `int` | `100000` | Cap on attempted steps (a rejected step spends the budget like an accepted one) before the run is refused. Values `≤ 0` mean the default. | + +**`ODEWatch`**: one scalar quantity to watch along a trajectory, passed to +`IntegrateODEEvents`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Function` | `func(t float64, y *tensor.Array) (float64, error)` | none, and a nil function is an error | The scalar `g(t, y)` whose zero crossings are reported. | +| `Direction` | `int` | `0` | `+1` records rising crossings, `-1` falling ones, `0` both. | + +**`ODEEventHit`**: one recorded zero crossing, returned by +`IntegrateODEEvents` sorted by time. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Time` | `float64` | set by the solver | When the crossing happened. | +| `State` | `*tensor.Array` | set by the solver | The state at the crossing. | +| `Watch` | `int` | set by the solver | Index into the `watches` slice of the watch that fired. | +| `Rising` | `bool` | set by the solver | True for a rising crossing, false for a falling one. | + +**`DAEOptions`**: tunes the per-step Newton solves of `IntegrateDAE`; +`RelTol` and `AbsTol` scale the Newton tolerance. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `RelTol` | `float64` | `1e-6` | Relative scale of the Newton tolerance. Values `≤ 0` mean the default. | +| `AbsTol` | `float64` | `1e-9` | Absolute scale of the Newton tolerance. Values `≤ 0` mean the default. | + +**`BDFVarOptions`**: tunes `IntegrateBDFVar`; it carries the `ODEOptions` +defaults and one extra output pointer. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `RelTol` | `float64` | `1e-6` | Relative part of the local error tolerance. Values `≤ 0` mean the default. | +| `AbsTol` | `float64` | `1e-9` | Absolute part of the local error tolerance. Values `≤ 0` mean the default. | +| `MaxSteps` | `int` | `100000` | Cap on attempted steps. Values `≤ 0` mean the default. | +| `Stats` | `*BDFVarStats` | `nil` | When not nil, receives the run's counters. | + +**`BDFVarStats`**: what a variable-order run did, written through +`BDFVarOptions.Stats`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Steps` | `int` | set by the solver | Accepted steps. | +| `Rejected` | `int` | set by the solver | Rejected steps. | +| `MaxOrder` | `int` | set by the solver | Highest order the driver reached, between 1 and 5. | + +**`BoundaryConditions`**: fixes the state of a boundary value problem at +the two ends, read by `IntegrateBoundary` and `SolveBoundaryCollocation`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Start` | `[]int` | none; may be empty | Components prescribed at `t0`, their values read from the initial state. | +| `End` | `[]int` | none, and an empty `End` is an error | Components prescribed at `t1`, parallel to `EndValues`. | +| `EndValues` | `[]float64` | none | The prescribed values at `t1`, one per entry of `End`. | + +**`CollocationOptions`**: tunes `SolveBoundaryCollocation`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `RelTol` | `float64` | `1e-6` | Together with `AbsTol`, scales the mesh-refinement estimate and floors the Newton convergence, an order of magnitude below it. Values `≤ 0` mean the default. | +| `AbsTol` | `float64` | `1e-9` | The absolute half of the same pair. Values `≤ 0` mean the default. | +| `InitialNodes` | `int` | `10` | Interval count of the uniform starting mesh. Values `≤ 0` mean the default. | +| `MaxNodes` | `int` | `256` | Cap on the refined mesh; the Newton matrix is factored by the dense LU, so the cap also bounds the per-round cost. Values `≤ 0` mean the default. | +| `MaxIterations` | `int` | `40` | Cap on the damped Newton rounds on each mesh. Values `≤ 0` mean the default. | + +**`CollocationSolution`**: the solved problem, returned by +`SolveBoundaryCollocation`. The piecewise cubic Hermite through +`(Mesh, Values, Slopes)` is the collocation solution itself, so interpolating +between the nodes on that data is exact to the solver's tolerance. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Mesh` | `[]float64` | set by the solver | The node times. | +| `Values` | `[]*tensor.Array` | set by the solver | `Values[k]` is the state at `Mesh[k]`. | +| `Slopes` | `[]*tensor.Array` | set by the solver | `Slopes[k]` is the derivative `y' = f(t, y)` at `Mesh[k]`. | + +**`QuadratureOptions`**: tunes `IntegrateFunction`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `RelTol` | `float64` | `1e-10` | Relative part of the target for the summed error estimate. Values `≤ 0` mean the default. | +| `AbsTol` | `float64` | `1e-12` | Absolute part of the same target. Values `≤ 0` mean the default. | +| `MaxIntervals` | `int` | `256` | Cap on subintervals; reaching it with the tolerance unmet is an error. Values `≤ 0` mean the default. | + +**`FilonOptions`**: tunes `IntegrateFilon`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Panels` | `int` | `≤ 0` means automatic | The number of equal panels the interval splits into. Automatic keeps each panel at most about `Nodes` half-wavelengths of the carrier, the range the weights are built exact in. A forced count whose panels carry more than 4096 half-wavelengths is refused. | +| `Nodes` | `int` | `≤ 0` means `16` | The Gauss-Legendre node count per panel; the amplitude interpolant's degree is `Nodes−1`. Values outside `[2, 32]` are refused. | + +**`CubatureOptions`**: tunes `IntegrateND`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Tolerance` | `float64` | `1e-10` | Bounds the global sum of box error estimates. Values `≤ 0` mean the default. | +| `MaxEvals` | `int` | `2000000` | Bounds the function evaluations; an exhausted budget is an error naming the achieved estimate. Values `≤ 0` mean the default. | + +**`MidpointOptions`**: tunes the per-step implicit stage of +`IntegrateMidpoint`, which is a root find. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Tolerance` | `float64` | `1e-13` | Convergence tolerance of the stage solve, deliberately tight because the energy band wants it. Values `≤ 0` mean the default. | +| `MaxIterations` | `int` | `100` | Cap on the root find's iterations. Values `≤ 0` mean the default. | + +**`FEMPoissonOptions`**: everything `SolvePoissonFEM2D` needs beside the +mesh and the source. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Kappa` | `float64` | none, and it must be positive when `KappaFunc` is nil | The constant conductivity used when `KappaFunc` is nil; a non-finite value is refused either way. | +| `KappaFunc` | `func(x, y float64) float64` | `nil` | The conductivity at a point, evaluated at the triangle centroids; a non-positive or non-finite value names the triangle. | +| `DirichletNodes` | `[]int` | none, and an empty slice is an error | The vertices with prescribed values. | +| `DirichletValues` | `[]float64` | none | The prescribed values, parallel to `DirichletNodes`. | +| `NeumannEdges` | `[]int` | none | Boundary edges as flat pairs of vertex indices; each receives half of `length·flux` at its midpoint into both endpoints. | +| `NeumannFlux` | `func(x, y float64) float64` | `nil`, which means zero flux | The flux `κ∂u/∂n` along each edge's outward normal. | +| `Ordering` | `linalg.SparseOrdering` | zero value, the natural order | Fill-reducing permutation for the sparse Cholesky factorisation; meshes usually want `SparseOrderingReverseCuthillMcKee`. | + +**`FEMPoisson3DOptions`**: the same for `SolvePoissonFEM3D`, with faces +instead of edges. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Kappa` | `float64` | none, and it must be positive when `KappaFunc` is nil | The constant conductivity used when `KappaFunc` is nil; a non-finite value is refused either way. | +| `KappaFunc` | `func(x, y, z float64) float64` | `nil` | The conductivity at a point, evaluated at the tetrahedron centroids; a non-positive or non-finite value names the element. | +| `DirichletNodes` | `[]int` | none, and an empty slice is an error | The vertices with prescribed values. | +| `DirichletValues` | `[]float64` | none | The prescribed values, parallel to `DirichletNodes`. | +| `NeumannFaces` | `[]int` | none | Boundary faces as flat triples of vertex indices; each face's integral comes from the degree-2 edge-midpoint rule. | +| `NeumannFlux` | `func(x, y, z float64) float64` | `nil`, which means zero flux | The flux `κ∂u/∂n` along each face's outward normal. | +| `Ordering` | `linalg.SparseOrdering` | zero value, the natural order | Fill-reducing permutation for the sparse Cholesky factorisation. | + +**`TriangleMesh2D`**: a conforming triangular mesh, built by +`GridTriangleMesh2D` or `NewTriangleMesh2D` and consumed by +`SolvePoissonFEM2D`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Vertices` | `[]float64` | none, filled by the constructor | The vertex coordinates as x,y pairs, two entries per vertex. | +| `Triangles` | `[]int64` | none, filled by the constructor | Three vertex indices per triangle. | + +**`TetraMesh3D`**: a conforming tetrahedral mesh, built by +`BoxTetraMesh3D` or `NewTetraMesh3D` and consumed by `SolvePoissonFEM3D`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Vertices` | `[]float64` | none, filled by the constructor | The vertex coordinates as x,y,z triples, three entries per vertex. | +| `Tetrahedra` | `[]int64` | none, filled by the constructor | Four vertex indices per tetrahedron, in positive orientation. | + +#### Errors + +The package exports no error values and no error types: every failure comes +back as a plain non-nil error carrying the library's `tensor: ` prefix and, +in almost every case, the name of the call that refused. The conditions, by +family: + +- **Any ODE solver**: a state that is not a rank-1 non-empty array, or is + complex; an `f` that returns an array of the wrong shape; an exhausted + `MaxSteps`; a step size that has collapsed below the resolution of `t`. + Backward integration is supported, not an error. +- **Non-finite values**: `IntegrateRK4`, `IntegrateROS4`, `IntegrateDAE`, + `SolveBoundaryCollocation` and the symplectic family refuse a derivative, + stage, residual or acceleration that is not finite; `IntegrateODE`, the BDF + drivers and `IntegrateBoundary` do not scan `f`'s output for one. +- **`IntegrateODEPath` and `IntegrateBoundary`**: `nSamples` below 2. +- **`IntegrateODEEvents`**: no watches, a watch with a nil `Function`, a + watch that returns an error, or a watch value that is not finite. +- **`IntegrateDAE`**: a `steps` count below 1; a mass matrix that is wrongly + shaped, complex, non-finite, nonsingular, has unequal zero-row and + zero-column counts, or carries rank deficiency outside whole zero rows; + an initial residual on an algebraic row beyond the consistency tolerance; + a singular algebraic block at `t0`, the index-1 refutation; a Newton + iteration that cannot converge. +- **`IntegrateBoundary`**: a condition count that does not equal the state + length; an out-of-range or repeated index in `Start` or in `End`; + `EndValues` of the wrong length; an exhausted step budget in any trial + trajectory; a shooting root find that cannot converge. +- **`SolveBoundaryCollocation`**: an inconsistent `BoundaryConditions` set, + a non-positive interval, a starting mesh past `MaxNodes`, refinement that + would grow past `MaxNodes`, a singular Newton matrix, or an iteration that + cannot converge. +- **The symplectic family**: `steps` below 1; a position and momentum that + are not rank-1 arrays of equal non-zero length; a complex or `int` state; + a non-finite state entry; an acceleration or gradient of the wrong shape + or a non-finite one; a root find that cannot converge (`IntegrateMidpoint`). +- **`GaussLegendreNodes`**: `n` outside 1 to 128; a Newton iteration that + fails to converge at rounding level. +- **`IntegrateFunction`**: NaN bounds; a non-finite integrand value; a + tolerance that cannot be met within `MaxIntervals` subintervals. An empty + interval (`a == b`) returns zero with a nil error. +- **`IntegrateFilon`**: NaN or infinite bounds; a NaN or infinite frequency; + `Nodes` outside `[2, 32]`; a forced `Panels` whose panels carry more than + 4096 half-wavelengths of the carrier; a span that overflows the float64 + range; a frequency whose span product leaves no representable panel count; + an amplitude that fails or returns a non-finite value. A reversed interval + negates both parts and `a == b` returns zeros with a nil error. +- **`IntegrateND`**: empty or unequal-length bounds; a non-positive bound + (an edge runs from `lower[d]` to `upper[d]`, so it must be increasing); a + non-finite integrand value; a dimension whose root box or single bisection + already exceeds `MaxEvals`; an exhausted evaluation budget. +- **The 1-D PDE solvers**: an initial state that is not rank-1, is empty, + is complex, or holds a non-finite value; non-positive `dx`, `tFinal` or + `dt`; a step count above `1e12`; `samples` below 2; a non-positive + diffusivity; a non-finite boundary or ghost value; a non-finite transport + speed; a CFL violation in `IntegrateWave1D`, `IntegrateAdvection1D`, + `IntegrateAdvectionDiffusion1D` or `IntegrateUpwindAdvection1D`. +- **`IntegrateHeat2D` and `IntegrateWave2D`**: an initial state that is not + rank 2, a grid below 3×3, non-positive spacings, `tFinal` or `dt`; + `samples` below 2; a non-finite state or velocity; a missing or mismatched + velocity; a non-positive wave speed; a CFL violation in `IntegrateWave2D`. +- **The meshers and mesh constructors**: a non-positive cell count; a + non-finite or non-positive extent; a vertex table or element table of the + wrong rank or column count; a table that does not hold integers; fewer + than three vertices (2-D) or four (3-D); an empty element table; a + non-finite coordinate; an out-of-range index; a degenerate element; a + negatively oriented tetrahedron. +- **The FEM solvers**: a nil mesh (3-D); a non-positive conductivity, or a + non-finite conductivity from `KappaFunc`; a `DirichletNodes`/`DirichletValues` + length mismatch; no Dirichlet node at + all; an out-of-range or non-finite Dirichlet entry; a Neumann index list + that is not a whole number of edges or faces, names an out-of-range or + repeated vertex, or repeats a vertex within a face; a non-finite source + value or Neumann flux; a degenerate element + found during assembly; a factorisation failure from `linalg`. + +#### Workflow + +```go +package main + +import ( + "fmt" + "log" + "math" + + tensor "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/integrate" +) + +func main() { + // The oscillator y″ = −y as a first-order system, y(0) = (1, 0). + 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 plain initial value problem. + end, err := integrate.IntegrateODE(f, 0, 1, y0, integrate.ODEOptions{}) + if err != nil { + log.Fatal(err) + } + fmt.Printf("y(1) = %.6f\n", end.FloatAt(0)) + + // The same trajectory with the level y = 0.5 watched in both + // directions, the crossing times reported alongside the final state. + level := func(t float64, y *tensor.Array) (float64, error) { return y.FloatAt(0) - 0.5, nil } + watches := []integrate.ODEWatch{ + {Function: level, Direction: -1}, // falling + {Function: level, Direction: +1}, // rising + } + hits, end, err := integrate.IntegrateODEEvents(f, 0, 7, y0, watches, integrate.ODEOptions{}) + if err != nil { + log.Fatal(err) + } + for _, h := range hits { + fmt.Printf("crossing at t = %.4f, y = %.4f\n", h.Time, h.State.FloatAt(0)) + } + + // The same problem as a two-point boundary value problem: + // y(0) = 0 with y(π/2) = 1 fixes the free initial slope. + shoot := func(t float64, y *tensor.Array) (*tensor.Array, error) { + return tensor.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2) + } + start, _ := tensor.FromFloats([]float64{0, 0.5}, 2) + bc := integrate.BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}} + times, states, err := integrate.IntegrateBoundary(shoot, 0, math.Pi/2, start, bc, 3, + integrate.ODEOptions{RelTol: 1e-10, AbsTol: 1e-13}) + if err != nil { + log.Fatal(err) + } + fmt.Printf("y'(0) = %.6f, y(π/2) = %.6f\n", states[0].FloatAt(1), states[len(times)-1].FloatAt(0)) + + // Quadrature and cubature. + tail, 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("∫ e^(-x²) over [0, ∞) = %.6f ± %.1e\n", tail, errEst) + + // A PDE: heat flow from sin(πx) on [0, 1] with Dirichlet ends, and + // a finite element Poisson solve on the unit square. + const n = 399 + u0 := make([]float64, n) + for i := range u0 { + u0[i] = math.Sin(math.Pi * float64(i+1) / float64(n+1)) + } + state, _ := tensor.FromFloats(u0, n) + history, err := integrate.IntegrateHeat1D(state, 1, 1.0/float64(n+1), 0.1, 1e-4, 5, 0, 0) + if err != nil { + log.Fatal(err) + } + fmt.Printf("u(1/2, 0.1) = %.6f\n", history.FloatAt((5-1)*n+n/2)) + + mesh, err := integrate.GridTriangleMesh2D(0, 0, 1, 1, 16, 16) + 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] + if x == 0 || x == 1 || y == 0 || y == 1 { + 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 := (16/2)*(16+1) + 16/2 + fmt.Printf("u(1/2, 1/2) = %.4f\n", u.FloatAt(centre)) +} +``` + +Each piece of that program is also a runnable example in the package's +test file, with the printed output checked by `go test`. + +### Package `stats` + +Distributions, descriptive summaries, classical inference and the statistical models that fit a response to a design. The distribution surface covers the normal, exponential, gamma, chi-square, Student t, Poisson, binomial, negative binomial, Weibull, lognormal and Pareto laws, the Dirichlet, the multivariate normal, and the noncentral chi-square, F and t families. Each univariate law is carried as far as the law goes: a density where one exists, a CDF and a quantile always, and a generator for the laws the library draws from. On that foundation sit the descriptives, the hypothesis tests, the contingency table analyses, the linear, generalised linear, regularised, robust and quantile regressions, the linear mixed model, principal components, clustering, hidden Markov models, kernel density estimation and Gaussian-process regression. + +The package stops where the estimator ends. Nothing here is Bayesian, and every draw takes an explicit generator: no entry point reads a global source, so a fit is reproducible from its seed alone. The one search it deliberately will not run is the Gaussian-process hyperparameter fit: `MarginalLogLikelihood` is exposed as the objective, and the caller drives it from outside with the house minimiser, because the package has no edge to `optim` in the dependency graph. Complex-valued input and non-finite observations are refused rather than propagated, and a sample beyond the exactness contract of `TheilSenRegression` is refused rather than approximated. + +#### Distribution foundations + +| Call | What it does | +|---|---| +| `GammaLower(a, x)` | The regularised lower incomplete gamma `P(a, x)`, the CDF of a Gamma(shape `a`, rate 1) draw. Refuses `a ≤ 0` or `x < 0`, and errors when the power series or the continued fraction has not converged inside its budget of `1000 + 20·√a` rounds. | +| `GammaUpper(a, x)` | The regularised upper incomplete gamma `Q(a, x) = 1 - P(a, x)`, on the same foundation and with the same refusals as `GammaLower`. | +| `BetaIncomplete(x, a, b)` | The regularised incomplete beta `I_x(a, b)`, the CDF of a Beta(`a`, `b`) draw, by continued fraction with the symmetry switch at `x > (a+1)/(a+b+2)`. Refuses `a ≤ 0`, `b ≤ 0` or `x` outside `[0, 1]`. | + +#### Normal and multivariate normal + +| Call | What it does | +|---|---| +| `NormalCDF(x)` | Φ(x), the standard normal CDF, through one `Erfc`. Returns a plain `float64`: no error, for any `x`. | +| `NormalQuantile(q)` | The q-quantile of the standard normal, the inverse of `NormalCDF`, by a bracketed Newton walk through the density, with a reflected-tail bisection route below `2⁻⁵³`. Refuses `q` outside `[0, 1]`, NaN included, and answers an error at `q = 0` and `q = 1`, which have no finite quantile. | +| `MultivariateNormalLogDensity(mean, cov, x)` | The log density of N(`mean`, `cov`) at one point, through the Cholesky factor. `mean` and `x` rank 1, `cov` rank 2 symmetric positive definite. Refuses complex input, every non-finite entry, a mirror pair of the covariance that disagrees beyond a relative `1e-12`, and a non-positive pivot, naming the row. | +| `MultivariateNormalDraws(g, n, mean, cov)` | `n` draws from N(`mean`, `cov`) as an (n × d) array, one draw per row, from standard normals coloured by the Cholesky factor. Deterministic for a given generator state. Refuses a nil generator, `n < 1`, a shape or dimension mismatch, non-finite input, an asymmetric covariance and a non-positive pivot. | + +#### Exponential + +| Call | What it does | +|---|---| +| `ExponentialCDF(x, rate)` | `P(X ≤ x)` for X ~ Exponential(`rate`), evaluated through `Expm1` so the left tail keeps its digits. Refuses `rate ≤ 0` and a NaN `x`; `x ≤ 0` answers 0. | +| `ExponentialQuantile(q, rate)` | The q-quantile of Exponential(`rate`), by a bracketed Newton walk through the density. Refuses `q` outside `[0, 1]`, the open ends `q = 0` and `q = 1`, and a non-positive rate. | +| `ExponentialDraws(g, n, rate)` | `n` draws from Exponential(`rate`), mean 1/`rate`, by inverse CDF. Refuses `n < 1` or `rate ≤ 0`. | + +#### Gamma + +| Call | What it does | +|---|---| +| `GammaCDF(x, shape, rate)` | `P(X ≤ x)` for X ~ Gamma(`shape`, `rate`), the same parametrisation `GammaDraws` samples, through `GammaLower`. Refuses `shape ≤ 0` or `rate ≤ 0` and a negative `x`. | +| `GammaQuantile(q, shape, rate)` | The q-quantile of Gamma(`shape`, `rate`), by a bracketed Newton walk through the density, seeded at the mean `shape/rate`. Refuses the parameter contract of `GammaCDF` and the q contract of every quantile. | +| `GammaDraws(g, n, alpha, beta)` | `n` draws from Gamma(shape α, rate β) by Marsaglia-Tsang for α ≥ 1 and a boost with the exponential for α < 1. Refuses `n < 1` or a non-positive shape or rate. | + +#### Chi-square + +| Call | What it does | +|---|---| +| `ChiSquareCDF(x, df)` | `P(X ≤ x)` for X ~ χ²(`df`). Refuses `df < 1`. | +| `ChiSquareQuantile(q, df)` | The q-quantile of χ²(`df`), the critical value tables quote. Refuses `df < 1` and the q contract. | +| `ChiSquareDraws(g, n, df)` | `n` draws from χ²(`df`), the gamma(`df`/2, 2) law. Refuses `df < 1`, and `n < 1` through `GammaDraws`. | + +#### Student t + +| Call | What it does | +|---|---| +| `StudentTCDF(t, df)` | `P(T ≤ t)` for T ~ Student t(`df`), in closed form through the incomplete beta. Refuses `df < 1`. | +| `StudentTQuantile(q, df)` | The q-quantile of Student t(`df`) on the signed axis, with the same reflected-tail route the normal quantile takes. Refuses `df < 1`, `q` outside `[0, 1]`, and the open ends. | +| `StudentTDraws(g, n, df)` | `n` draws from Student t(`df`) as N(0,1) over the square root of χ²(`df`)/`df`; a χ² draw that underflowed to zero becomes a signed infinity rather than a NaN. Refuses `df < 1`. | + +#### Poisson + +| Call | What it does | +|---|---| +| `PoissonCDF(k, lambda)` | `P(N ≤ k)` for N ~ Poisson(`lambda`), through `ΓUpper(k+1, λ)`. Refuses `lambda ≤ 0`; `k < 0` answers 0. | +| `PoissonQuantile(q, lambda)` | The smallest `k` with `P(N ≤ k) ≥ q`, by doubling the bound then bisecting the integer grid. Refuses `lambda ≤ 0`, `q` outside `[0, 1]` and a bracket that never reaches `q`. `q = 0` answers 0. | +| `PoissonDraws(g, n, lambda)` | `n` draws from Poisson(`lambda`): the Knuth multiplication method for λ < 30, the normal approximation above. Refuses `n < 1` or `lambda < 0`; λ = 0 answers an array of zeros. | + +#### Binomial + +| Call | What it does | +|---|---| +| `BinomialCDF(k, trials, p)` | `P(X ≤ k)` for X ~ Binomial(`trials`, `p`), through `I_{1-p}(n-k, k+1)`. Refuses `trials < 1` or `p` outside `(0, 1)`; `k < 0` answers 0 and `k ≥ trials` answers 1. | +| `BinomialQuantile(q, p, trials)` | The smallest `k` with `P(X ≤ k) ≥ q`. Refuses `p` outside `(0, 1)` and the shared q contract. | +| `BinomialDraws(g, n, trials, p)` | `n` draws from Binomial(`trials`, `p`) by the exact per-trial uniform loop, so it costs O(`trials`) uniforms per draw. Refuses `n < 1`, `trials < 1` or `p` outside `[0, 1]`, the closed ends included. | + +#### Negative binomial + +| Call | What it does | +|---|---| +| `NegativeBinomialPMF(k, r, p)` | `P(X = k)` for the number of failures `X` before the `r`-th success, assembled in log space with `lgamma`. Refuses `r < 1` or `p` outside `(0, 1)`; `k < 0` answers 0. | +| `NegativeBinomialCDF(k, r, p)` | `P(X ≤ k)` by direct summation of the PMF terms under their multiplicative recurrence, so the sum carries no cancellation. Same refusals as the PMF. | +| `NegativeBinomialQuantile(q, p, r)` | The smallest `k` with `P(X ≤ k) ≥ q`, through the discrete bracketed search. Refuses `r < 1`, `p` outside `(0, 1)` and the shared q contract. | + +#### Weibull + +| Call | What it does | +|---|---| +| `WeibullDensity(x, k, lambda)` | The Weibull(k, λ) density, `(k/λ)(x/λ)^{k-1}e^{-(x/λ)^k}`, at `x ≥ 0`, and 0 below the support and at +∞. Refuses a shape or scale that is not finite and positive, and a NaN `x`. At `x = 0` the formula speaks: 0 for k > 1, 1/λ for k = 1, +∞ for k < 1. | +| `WeibullCDF(x, k, lambda)` | `P(X ≤ x)` for X ~ Weibull(k, λ), the closed form `1 - e^{-(x/λ)^k}` through `Expm1`. Same parameter refusals; `x ≤ 0` answers 0. | +| `WeibullQuantile(q, k, lambda)` | The q-quantile of Weibull(k, λ), the closed form `λ(-ln(1-q))^{1/k}`. Refuses the parameter contract and the open ends of the q range. | + +#### Lognormal + +| Call | What it does | +|---|---| +| `LognormalDensity(x, mu, sigma)` | The lognormal density with location μ and log-scale σ at `x > 0`, assembled in log space so a subnormal `x` still answers 0 instead of a NaN. Refuses a non-finite μ, a σ that is not finite and positive, and a NaN `x`; 0 at `x ≤ 0` and at +∞. | +| `LognormalCDF(x, mu, sigma)` | `P(X ≤ x)` for X ~ lognormal(μ, σ), the normal CDF at `(ln x - μ)/σ`. Same refusals; `x ≤ 0` answers 0. | +| `LognormalQuantile(q, mu, sigma)` | The q-quantile of lognormal(μ, σ) through the normal quantile, `e^{μ + σ·Φ⁻¹(q)}`. Refuses the location and scale contract and the q contract of `NormalQuantile`. | + +#### Pareto + +| Call | What it does | +|---|---| +| `ParetoDensity(x, xm, alpha)` | The Pareto density `α·x_m^α/x^{α+1}` at `x ≥ x_m`, and 0 below the support and at +∞. Refuses a scale or tail index that is not finite and positive, and a NaN `x`. | +| `ParetoCDF(x, xm, alpha)` | `P(X ≤ x)` for X ~ Pareto(x_m, α), the closed form `1 - (x_m/x)^α` through `Expm1`, so the answers just above the support keep their digits. Same parameter refusals; `x < x_m` answers 0. | +| `ParetoQuantile(q, xm, alpha)` | The q-quantile of Pareto(x_m, α), `x_m(1-q)^{-1/α}`. Refuses the parameter contract and the open ends of the q range. | + +#### Dirichlet + +| Call | What it does | +|---|---| +| `DirichletDensity(alpha, x)` | The Dirichlet density at the simplex point `x`, `Π x_i^{α_i-1}/B(α)`. Refuses fewer than two components, mismatched lengths, a concentration that is not finite and positive, a non-finite or negative `x`, and a point that does not sum to 1 within `1e-9`. A boundary `x_i = 0` answers +∞ below α_i = 1, contributes 1 at α_i = 1 and 0 above. | +| `DirichletMean(alpha)` | The mean vector `α_i/α₀`. Refuses fewer than two components and a concentration entry that is not finite and positive. | +| `DirichletMode(alpha)` | The interior mode `(α_i-1)/(α₀-k)`. Refuses the concentration contract and any `α_i ≤ 1`, which pushes the mode onto the boundary. | +| `DirichletDraws(g, n, alpha)` | `n` draws as an (n, k) array whose rows sum to one, each row scaling the independent gamma(α_i, 1) draws `GammaDraws` uses. Refuses `n < 1`, fewer than two components and the concentration contract. A row whose every gamma draw underflowed falls back to the uniform row. | + +#### Noncentral chi-square, F and t + +| Call | What it does | +|---|---| +| `NoncentralChiSquareCDF(x, df, lambda)` | `P(X ≤ x)` for X ~ χ²(ν, λ), the Poisson(λ/2) mixture of central χ²(ν + 2i) CDFs. Refuses `df < 1`, a negative or non-finite λ, a NaN `x`, and a noncentrality whose Poisson weight peak sits past the term budget (the mixture answers up to roughly 2·10⁵). λ = 0 answers through `ChiSquareCDF` exactly. | +| `NoncentralChiSquareDensity(x, df, lambda)` | The χ²(ν, λ) density, the same mixture with central densities. Refuses `df < 1`, a negative or non-finite λ, a NaN `x`, and the term budget. 0 below `x = 0`; at `x = 0` it is +∞ for `df = 1`, `e^{-λ/2}/2` for `df = 2` and 0 past that. | +| `NoncentralChiSquareQuantile(q, df, lambda)` | The q-quantile of χ²(ν, λ) by bisection, seeded near the mean ν + λ. Refuses `df < 1`, a negative or non-finite λ, and the shared q contract. | +| `NoncentralFCDF(x, df1, df2, lambda)` | `P(X ≤ x)` for X ~ F(ν₁, ν₂, λ), the Poisson(λ/2) mixture of the scaled central pieces, summed through `BetaIncomplete`. Refuses `df1 < 1` or `df2 < 1`, a negative or non-finite λ, a NaN `x`, and the term budget. λ = 0 is the central F exactly. | +| `NoncentralFQuantile(q, df1, df2, lambda)` | The q-quantile of F(ν₁, ν₂, λ) by bisection, seeded at 1. Refuses `df1 < 1` or `df2 < 1`, a negative or non-finite λ, and the shared q contract. | +| `NoncentralTCDF(t, df, delta)` | `P(T ≤ t)` for T ~ t(ν, δ), through Lenth's even and odd series. Refuses `df < 1`, a non-finite δ, and a series that misses its term budget (the walk answers δ up to about 440). δ = 0 answers through `StudentTCDF` and `t = 0` through Φ(-δ). | +| `NoncentralTQuantile(q, df, delta)` | The q-quantile of t(ν, δ) by bisection on the signed axis, the bracket growing both ways from the seed. Refuses `df < 1`, a non-finite δ and the shared q contract. | + +#### Descriptives + +| Call | What it does | +|---|---| +| `Median(a)` | The median as a `float64`, averaging the two middle values on even length (in integer arithmetic for an int sample, so the exact answer survives past 2⁵³). Refuses a complex or empty array and any non-finite sample. | +| `Std(a)` | The population standard deviation, ddof = 0. Refuses a complex or empty array. | +| `Var(a)` | The population variance, ddof = 0. Refuses a complex or empty array. | +| `VarSample(a)` | The unbiased sample variance, ddof = 1. Refuses a complex array and fewer than two elements. | +| `MedianAbsoluteDeviation(a)` | The median of `|x - median(x)|`, the robust scale that breaks down only when nearly half the sample is wild. Multiply by 1.4826 to read it as a standard deviation on Gaussian data. Refuses whatever `Median` refuses. | +| `TrimmedMean(a, fraction)` | The mean after dropping `fraction` of the samples from each tail, floored to whole samples. Refuses a fraction outside `[0, 0.5)`, a complex or empty sample, a non-finite sample, and a trim that leaves nothing. | +| `Quantile(a, qs)` | The quantiles in `qs`, each in `[0, 1]`, by linear interpolation on the sorted values (on the exact integer difference for an int sample). Refuses a complex or empty array, a non-finite sample, and any `q` outside `[0, 1]`. | +| `Histogram(a, bins)` | The counts and the edges of `bins` equal-width bins over `[min, max]`: an int count array of length `bins` and a float edge array of length `bins + 1`. The top bin absorbs the maximum; an all-equal sample widens to `[v-0.5, v+0.5]`. Refuses a complex or empty array, `bins < 1`, more than 1048576 bins, a non-finite sample, and a sample spanning more than the float64 range. | +| `BinCounts(a, bins)` | The count array of `Histogram` alone, with the same refusals. | +| `Histogram2D(x, y, xBins, yBins)` | The (xBins × yBins) count matrix and the two edge vectors of the paired samples. Bins are closed on the left and the last bin absorbs the maximum, exactly as `Histogram`'s are. Refuses unequal lengths, empty samples, complex input, a bin count below 1, a product above 1048576, and a non-finite sample. | +| `RollingMean(a, window)` | The mean of every window of the series: `n - window + 1` elements, element `i` summarising samples `[i, i+window)`. Refuses a non-vector or complex series and a window outside `[1, n]`. | +| `RollingSum(a, window)` | The total of every window, with the shape and refusals of `RollingMean`. | +| `RollingMin(a, window)` | The smallest sample of every window. A NaN never wins a comparison and an all-NaN window answers NaN. | +| `RollingMax(a, window)` | The largest sample of every window, with the same NaN rule as `RollingMin`. | + +#### Inference + +| Call | What it does | +|---|---| +| `CovarianceMatrix(a)` | The sample covariance matrix of an (n, p) observation array, with the `1/(n-1)` normalisation. Refuses a non-2-D array, fewer than two observations, complex observations and a non-finite observation. | +| `CorrelationMatrix(a)` | The Pearson correlation matrix, the covariance normalised by each column's sample standard deviation, with the diagonal exactly 1. Adds a refusal of a zero-variance column to the covariance contract. | +| `WelchTTest(a, b)` | Welch's t-test on two independent samples: the statistic, the Welch-Satterthwaite degrees of freedom (generally fractional) and the two-sided p-value from the closed-form t tail. Needs at least two observations per sample, real and finite, and refuses two samples of zero variance. | +| `ANOVAOneWay(groups)` | The one-way analysis of variance: the F statistic and the upper tail of F on (k-1, N-k) degrees of freedom. Needs at least two real, non-empty, finite groups with N > k, and refuses groups with no within-group variance or observations that all share one value. Distinct means with no within-group noise answer +∞ with p = 0. | +| `MannWhitneyU(a, b)` | The rank-based Mann-Whitney U test: U and the two-sided p-value from the normal approximation with the continuity correction and the tie-corrected variance. Refuses an empty, complex or non-finite sample, and data in which every observation is tied. | +| `KolmogorovSmirnovTest(a, b)` | The largest vertical distance between two empirical distribution functions and the asymptotic Kolmogorov p-value, accurate from a few dozen observations up. Refuses an empty, complex or non-finite sample. | +| `ChiSquareGoodnessOfFit(observed, expected)` | Pearson's test of an observed frequency table against expected ones: the statistic, the degrees of freedom (bins minus one) and the χ² upper tail. Needs equal lengths and at least two bins, every expected entry positive, and refuses non-finite input. | +| `BootstrapCI(data, statistic, level, resamples, seed)` | The percentile bootstrap of a statistic: the `alpha/2` and `1-alpha/2` quantiles of the bootstrap distribution, `alpha = 1 - level`, with the resampling on a generator seeded by `seed`. The statistic receives a fresh array per resample and may return its own error. Refuses empty or complex data, a level outside `(0, 1)` and fewer than two resamples. | +| `SpearmanRho(x, y)` | Spearman's rank correlation, computed as the Pearson correlation of the mid-ranks, which is exactly the tie-corrected form. Needs equal-length, real, finite pairs, at least two of them, and refuses a sample whose ranks all agree. | +| `KendallTau(x, y)` | Kendall's τ-b, the tie-corrected rank correlation, `±1` on every perfectly monotone pairing. Same input contract as `SpearmanRho`, and refuses a constant sample, which leaves the denominator zero. | + +#### Contingency tables + +| Call | What it does | +|---|---| +| `FisherExactTest(table, alternative)` | Fisher's exact test on a 2×2 table of counts, the answer for a cross-classification too sparse for the χ² approximation. Conditions on the margins and sums the hypergeometric probabilities of the tables at least as extreme as the observed one, exactly and in log space; the two-sided sum collects every table whose probability does not exceed the observed one's. Returns the p-value and the sample odds ratio (`+Inf` where a zero cell sits against a full one, `NaN` with p = 1 where a zero margin leaves only one table). The alternative is `TwoSided`, `Less` or `Greater`. Refuses a table that is not 2×2, complex input, a non-finite, negative or fractional count, a count past 2^51 where float64 no longer holds the integers stepwise, and margins spanning more than 4000000 support points, for which `ChiSquareIndependence` is the honest test. | +| `ChiSquareIndependence(table)` | Pearson's χ² test of independence on an r×c table: the expected count is the row total times the column total over the grand total, the statistic is Σ(O−E)²/E and the p-value its tail on (r−1)(c−1) degrees of freedom. Refuses fewer than two rows or columns, complex input, a non-finite, negative or fractional count, a zero row or column total, and an all-zero table. | +| `McNemarTest(table)` | The exact McNemar test on paired counts: under the null the two off-diagonal disagreements split evenly, so the smaller one follows the binomial law at p = ½, and the two-sided p-value doubles the smaller tail through the incomplete beta, costing the same for millions of pairs as for a dozen. No discordant pair answers 1. The input contract is `FisherExactTest`'s own. | +| `CramersV(table)` | Cramér's V, the χ² statistic rescaled into [0, 1] by the sample size and the smaller margin: √(χ²/(n·(min(r,c)−1))). Zero means the counts sit exactly on independence. The input contract is `ChiSquareIndependence`'s own. | + +#### Multiple testing + +| Call | What it does | +|---|---| +| `Bonferroni(p)` | The Bonferroni-adjusted p-values, each scaled by the vector length and clamped at 1. Refuses an empty vector, a non-finite entry and an entry outside `[0, 1]`. | +| `Holm(p)` | The Holm step-down adjusted p-values, with the running maximum enforcing the non-decreasing order. Same refusals as `Bonferroni`. | +| `BenjaminiHochberg(p)` | The Benjamini-Hochberg step-up adjusted p-values, the q-values of the false discovery rate literature, with the running minimum enforcing the order. Same refusals as `Bonferroni`. | + +#### Linear and generalised linear models + +| Call | What it does | +|---|---| +| `LinearRegression(x, y)` | Ordinary least squares `y = X·β` with the full classical inference in a `LinearRegressionResult`. The intercept is supplied by the caller as a constant column. Refuses a design that is not rank 2, a response that is not rank 1, a length mismatch, complex or non-finite input, `n ≤ p`, and a rank-deficient or near-collinear design. | +| `WeightedLinearRegression(x, y, w)` | Weighted least squares with the positive weight `w[i]` on observation `i`: the normal equations run on the sqrt-weighted system, so every statistic is the weighted-theory one while `Fitted` and `Residuals` stay in the original units. Refuses a weight that is not finite and positive, then validates as `LinearRegression` does. | +| `LogisticRegression(x, y)` | The binary response `y = Bernoulli(sigmoid(X·β))` by maximum likelihood: Newton-Raphson to a coefficient movement below `1e-10`, at most 100 steps, with Wald inference from the inverse Fisher information. `y` must hold only 0 and 1. Perfectly separable data has no finite optimum and is reported as an error, as is a singular Fisher information or a near-collinear design. | +| `PoissonRegression(x, y)` | The count response `y = Poisson(exp(X·β))` by maximum likelihood on the log link, Newton-Raphson with steps halved while the likelihood does not rise, to the same tolerance and budget. `y` must hold non-negative integers. A singular Fisher information, a near-collinear design, non-finite or negative or fractional responses, and an exhausted iteration budget are all errors. | +| `QuantileRegression(x, y, tau)` | The tau-th conditional quantile by the Frisch-Newton interior-point method on the dual of the check-loss program, started from the least squares fit, at most 100 iterations with the barrier parameter shrinking by `0.25` per iteration down to a floor of `1e-16` and a coefficient tolerance of `1e-14`. Refuses `tau` outside the open interval `(0, 1)`, and validates the design as `LinearRegression` does. | + +#### Regularised and robust models + +| Call | What it does | +|---|---| +| `Lasso(x, y, lambda)` | The pure lasso at one `lambda`: `ElasticNet` with alpha = 1. The slopes and only the slopes are penalised. | +| `ElasticNet(x, y, lambda, alpha)` | The elastic net, `(1/2n)·Σ(y - β₀ - x·β)² + λ·(α·Σ|β| + (1-α)/2·Σβ²)`, by coordinate descent on the design standardised to unit population standard deviation, at most 10000 cycles to a coefficient tolerance of `1e-8`, with the standardisation inverted before the result returns. Refuses a negative or NaN `lambda`, an `alpha` outside `[0, 1]`, fewer than two observations, a design column with no standard deviation to divide by, and complex or non-finite input. A rank-deficient design and more columns than rows are legitimate here. | +| `LassoPath(x, y, alpha)` | The whole regularisation path at the given mixing, warm-started along 100 log-spaced lambdas from the largest penalty that zeroes every slope down to its thousandth. Refuses an `alpha` outside `[0, 1]` and the same input contract as `ElasticNet`. | +| `HuberRegression(x, y)` | The Huber M-estimate of the linear model at the default tuning constant `DefaultHuberTuning`. | +| `HuberRegressionTuned(x, y, tuning)` | The Huber M-estimate at the given tuning constant: a quadratic loss inside the band `|r| ≤ tuning·σ` and a linear one outside it, fitted by iteratively reweighted least squares, at most 100 rounds to a coefficient tolerance of `1e-10`, with the robust scale σ re-estimated each round as 1.4826 times the median absolute deviation of the residuals. Refuses a tuning constant that is not finite and positive, a design that is not `n > p` or full rank, and complex or non-finite input. A robust scale that collapses to zero stops the iteration as an exact fit. | +| `TheilSenRegression(x, y)` | The Theil-Sen estimate of the simple model `y = a + b·x`: the median of the pairwise slopes over all pairs with distinct predictors, then the median of `y_i - b·x_i` at that slope. Needs at least three observations with at least two distinct predictors, all finite and real. Refuses more than `TheilSenMaxObservations` observations with the cost named, since the median is taken over all n(n-1)/2 slopes. | + +#### Linear mixed models + +| Call | What it does | +|---|---| +| `LinearMixedModel(y, x, z, groups)` | The Gaussian linear mixed model `y = X·β + Z·b + ε`: a fixed-effect design shared by every row plus a random-effect design whose coefficients vary by group, each group's `b_g` drawn from one shared unstructured q×q covariance Σ, the residual from N(0, σ²I). Σ and σ² are estimated by residual maximum likelihood through expectation maximisation over the random-effect posterior, with the fixed effects' own posterior covariance folded into the M step, so the fixed point is the REML optimum and not the ML one; β̂ follows by generalised least squares at the fitted components and carries standard errors from (XᵀV⁻¹X)⁻¹. The starting point is the data's own least squares, so no generator enters and the fit is deterministic. Convergence is the REML log likelihood settling under `1e-10` relative, within 500 sweeps; `Converged` names which happened. Refuses complex or non-finite input, a response shorter than two rows, a design of fewer than one column, a row count mismatch, at least as many fixed coefficients as observations, a negative group label, a group whose marginal covariance fails to factor, a singular fixed design, and a design whose REML criterion has no interior optimum (a random design that spans the fixed one under an unstructured covariance drives the components up a ridge without a summit; the fit refuses with the condition named). | + +#### Principal components + +| Call | What it does | +|---|---| +| `PCA(a)` | The decomposition of an (n, p) observation array onto its principal components, by the cyclic Jacobi eigensolver over the package's own covariance matrix, into a `PCAResult`. Needs at least two real, finite observations, and refuses a sample of no variance at all and a covariance that decomposed to a negative eigenvalue beyond the rounding scale. A rank-deficient covariance decomposes normally; only the whitening transforms are withheld. | +| `(r *PCAResult) Whiten(x)` | The unit-covariance representation of an (n, p) array on the fit's own variables: the centred rows expressed on the components and scaled by each component's standard deviation. Refuses a nil fit, a rank-deficient fit, a rank or width mismatch, and complex or non-finite input. | +| `(r *PCAResult) Unwhiten(z)` | The inverse of `Whiten`, recovering the observations in their own coordinates to rounding. Same refusals as `Whiten`. | + +#### Clustering + +| Call | What it does | +|---|---| +| `KMeans(g, x, k)` | `k` clusters over the rows of an (n, d) sample by Lloyd's iteration, seeded by k-means++ over the house generator, so the fit is deterministic for a generator state. Convergence is declared once no centre moves by more than `1e-10` scaled by the largest absolute coordinate in the sample, and at most 300 sweeps are run; `Converged` names which happened. An empty cluster takes the sample farthest from its centre, and the run is an error if the sample holds fewer distinct points than `k`. Refuses a nil generator, a non-2-D or complex or non-finite sample, `k < 1` and `k > n`. | +| `GaussianMixture(g, x, components)` | A mixture of multivariate normals over the rows of a sample by expectation maximisation, with the component densities through the Cholesky factors of their covariances, initialised from the k-means partitions of the same generator. The E step runs in log space with a max-guarded log-sum-exp normalisation, and each M-step covariance is floored by `1e-6` of the sample's mean per-dimension variance, so a component that collapsed onto one sample keeps a factorable covariance. Convergence is the log likelihood settling under `1e-9` relative, within 200 sweeps. Refuses a nil generator, fewer than two samples, a non-2-D or complex or non-finite sample, and a component count outside `[1, n]`. | +| `GaussianMixtureBIC(g, x, maxComponents)` | Fits the mixtures of 1 to `maxComponents` components and keeps the one with the lowest BIC, `-2·logL + p·ln n` with p the free parameters of the mixture. The grid travels in `BICGrid` and the winner in `Components`; a tie resolves to the smaller model. Refuses a largest component count outside `[1, n]` and the sample contract of `GaussianMixture`. | +| `HierarchicalClustering(x, method)` | The agglomerative dendrogram of the sample's rows over Euclidean distances. Every row starts as its own cluster and the closest pair merges repeatedly, the merged distances rewriting through the linkage's Lance-Williams update; the naive sweep costs O(n³) time and O(n²) memory and is deterministic, ties resolving toward the lowest indices. The method is `SingleLinkage`, `CompleteLinkage`, `AverageLinkage`, `CentroidLinkage` or `WardLinkage`; the centroid and Ward updates run on squared distances with the heights rooted, and the centroid rule is not monotone, so its heights may invert. Refuses a non-rank-2 or complex or non-finite sample, fewer than two rows, more than `HierarchicalMaxObservations` rows, and an unknown linkage. | +| `(d *Dendrogram) Cut(k)` | The flat partition into `k` clusters left by undoing the last k−1 merges, labelling every row 0 to k−1 by each cluster's smallest row, so the sample's first row always lands in cluster 0. Refuses a nil dendrogram and a `k` outside `[1, n]`. | +| `(d *Dendrogram) CutHeight(height)` | The partition left by applying merges in order while they sit at or below `height`: below the first merge the singletons, at or above the last one cluster. Labels as `Cut` does. Refuses a nil dendrogram and a non-finite height. | + +#### Hidden Markov models + +| Call | What it does | +|---|---| +| `NewHiddenMarkovModel(initial, transition, emission)` | Validates and copies a discrete hidden Markov model: `initial` the state distribution at the first step, `transition` the states×states one-step conditionals row-major, `emission` the states×symbols observation conditionals row-major. Every entry must be finite and in [0, 1] and every row must sum to 1 within 1e-9, so a model value is always safe to evaluate. Refuses an empty state set, a transition or emission matrix whose shape does not fit, and any row that fails those checks. | +| `(m *HiddenMarkovModel) Forward(observations)` | The scaled forward recursion: the filtered state posteriors, one row per step, and the sequence's log likelihood P(observations | model). The per-step scaling keeps a thousand-step sequence from underflowing. Refuses a nil model, an empty sequence and any observation outside the model's symbol set. | +| `(m *HiddenMarkovModel) Smooth(observations)` | The forward-backward smoothed posteriors P(state at t | all observations), one row per step, with the sequence's log likelihood. The refusals are `Forward`'s own. | +| `(m *HiddenMarkovModel) Viterbi(observations)` | The most likely state path through the sequence, decoded in log space with ties broken toward the lowest state index, with its log probability log P(path, observations | model). The refusals are `Forward`'s own. | +| `FitHiddenMarkovModel(g, observations, states, symbols)` | Baum-Welch expectation maximisation of a model over the given state and symbol counts: every starting row drawn from the flat Dirichlet through the house generator, so the fit is deterministic for the generator state, though expectation maximisation only climbs to a local maximum and a different state may land elsewhere. The re-estimation floors every parameter at 1e-12 and renormalises the rows, so no symbol or transition is silenced by one sweep. Convergence is the log likelihood settling under `1e-6` relative, within 500 sweeps, the tolerance stopping the crawl along the flat ridge the objective leaves rather than waiting for the parameters to stop moving. Refuses a nil generator, a count below one, an empty sequence and an out-of-range observation. | + +#### Density estimation and Gaussian processes + +| Call | What it does | +|---|---| +| `KernelDensity(sample, bandwidth, points)` | The Gaussian-kernel density estimate of the sample at every point, each sample contributing a unit-variance normal of width `bandwidth` averaged over the sample. A non-positive `bandwidth` asks for Silverman's rule, `0.9·min(σ, IQR/1.34)·n^{-1/5}`, with the σ fallback when the interquartile range is degenerate. Refuses a non-vector, complex or non-finite sample or points, a sample of fewer than two points, and a bandwidth that resolved to a non-positive or NaN width. | +| `Kernel` | The covariance function of a Gaussian process, one method: `Covariance(x, y)` returns the prior covariance `k(x, y)` of two points of equal length. Every house kernel carries unit amplitude, `k(x, x) = 1`, and its constructor enforces the parameter contract, so a `Kernel` value is always safe to evaluate. | +| `SquaredExponentialKernel(lengthScale)` | The squared-exponential (RBF) kernel `exp(-r²/(2·lengthScale²))` over the Euclidean distance, the smooth prior. Refuses a length scale that is not finite and positive. | +| `Matern32Kernel(lengthScale)` | The Matérn kernel with ν = 3/2, `(1 + √3·r/ℓ)·exp(-√3·r/ℓ)`, once differentiable. Refuses a length scale that is not finite and positive. | +| `Matern52Kernel(lengthScale)` | The Matérn kernel with ν = 5/2, twice differentiable. Refuses a length scale that is not finite and positive. | +| `PeriodicKernel(lengthScale, period)` | The periodic kernel `exp(-2·sin²(π·r/period)/lengthScale²)`, the prior over functions that repeat exactly with the period. Refuses a length scale or a period that is not finite and positive. | +| `GaussianProcessRegression(kernel, trainX, trainY, noiseVariance, testX)` | The Gaussian-process posterior at the test points: the posterior mean, the full posterior covariance between them and its diagonal in a `GaussianProcessResult`, conditioning on the training rows through the Cholesky factor of `K = k(X, X) + noiseVariance·I`. A zero noise variance is a legitimate noiseless fit and interpolates the training data exactly; duplicated training rows are then a singular Gram matrix and are refused. The inputs must be real and finite, the training and test designs of equal width, the response of the training length, and the noise variance finite and non-negative. | +| `MarginalLogLikelihood(kernel, trainX, trainY, noiseVariance)` | The log marginal likelihood of the observations under the prior the kernel defines, `-½·yᵀK⁻¹y - Σ ln L_ii - (n/2)·ln 2π`, the evidence of the hyperparameters. It is the objective a hyperparameter fit maximises; the package has no edge to `optim`, so the house minimiser drives this function from outside over the kernel parameters and the noise variance. Refuses a nil kernel, a negative or non-finite noise variance, and the training contract of `GaussianProcessRegression`. | + +#### Constants + +| Constant | Value | Meaning | +|---|---|---| +| `DefaultHuberTuning` | `1.345` | The tuning constant `HuberRegression` passes to `HuberRegressionTuned`: the literature's standard choice, 95 percent asymptotic efficiency at the Gaussian with the influence of an outlier bounded at 1.345 times a residual inside the band. | +| `TheilSenMaxObservations` | `4096` | The exactness contract of `TheilSenRegression`: the largest sample whose pairwise slopes are all medianed exactly. Beyond the cap the estimator refuses with the cost named rather than approximate. | +| `HierarchicalMaxObservations` | `4096` | The sample cap of `HierarchicalClustering`: the working distance matrix is quadratic in the sample, and past the cap the refusal names the cost rather than handing gigabytes to the allocator. | + +#### Options and results + +The package carries no options structs; every entry point takes its parameters as arguments. The result types below are the whole of its exported state. + +**`LinearRegressionResult`** : the least-squares fit and its inference, returned by `LinearRegression` and `WeightedLinearRegression`. Each coefficient slice is indexed by column of the design, in order. + +| Field | Type | Meaning | +|---|---|---| +| `Coefficients` | `[]float64` | The estimates β̂. | +| `StandardErrors` | `[]float64` | The estimated standard deviations of the coefficient estimators, σ̂²(XᵀX)⁻¹ on the diagonal. | +| `TStatistics` | `[]float64` | β̂/SE per coefficient. | +| `PValues` | `[]float64` | The two-sided p-values of the t-tests. | +| `ResidualVariance` | `float64` | σ̂² = RSS/(n - p). | +| `RSquared` | `float64` | The coefficient of determination; 1 by convention when the response is constant and reproduced exactly. | +| `AdjustedRSquared` | `float64` | R² adjusted for the degrees of freedom of the total and the residual. | +| `FStatistic` | `float64` | The model F test of every coefficient being zero. | +| `DModel` | `int` | The model degrees of freedom: p - 1 with a constant column in the design, p without one. | +| `DResidual` | `int` | The residual degrees of freedom, n - p. | +| `FPValue` | `float64` | The upper tail of the F distribution at `FStatistic`; 1 for an intercept-only design, which has no model term to test. | +| `Fitted` | `[]float64` | The fitted value of every design row. In `WeightedLinearRegression` these are in the original, unweighted units. | +| `Residuals` | `[]float64` | The residual of every row, aligned with the design. | + +**`LogisticRegressionResult`** : the fit of a binary response, returned by `LogisticRegression`. + +| Field | Type | Meaning | +|---|---|---| +| `Coefficients` | `[]float64` | The maximum-likelihood estimates β̂ on the logit scale. | +| `StandardErrors` | `[]float64` | The Wald standard errors, from the inverse Fisher information at the optimum. | +| `ZStatistics` | `[]float64` | β̂/SE per coefficient; an exact fit reports a signed infinity beside a zero standard error. | +| `PValues` | `[]float64` | The two-sided normal-tail probabilities. | +| `Fitted` | `[]float64` | The predicted probability per sample, clamped into `[1e-12, 1 - 1e-12]` exactly as the fitting loop clamps it. | +| `LogLikelihood` | `float64` | The maximised Bernoulli log likelihood. | +| `Iterations` | `int` | The Newton steps taken. | +| `Converged` | `bool` | Whether the coefficient update fell under the tolerance. | + +**`PoissonRegressionResult`** : the fit of a count response, returned by `PoissonRegression`. + +| Field | Type | Meaning | +|---|---|---| +| `Coefficients` | `[]float64` | The maximum-likelihood estimates β̂ on the log scale. | +| `StandardErrors` | `[]float64` | The Wald standard errors from the inverse Fisher information at the optimum. | +| `ZStatistics` | `[]float64` | β̂/SE per coefficient. | +| `PValues` | `[]float64` | The two-sided normal-tail probabilities. | +| `Fitted` | `[]float64` | The predicted mean count per sample, clamped into `[1e-12, 1e300]` exactly as the fitting loop clamps it. | +| `LogLikelihood` | `float64` | The maximised Poisson log likelihood, evaluated on the clamped `Fitted` values. | +| `Iterations` | `int` | The Newton steps taken. | +| `Converged` | `bool` | Whether the coefficient update fell under the tolerance. | + +**`ElasticNetResult`** : one regularised fit at a single lambda, returned by `Lasso` and `ElasticNet`. + +| Field | Type | Meaning | +|---|---|---| +| `Intercept` | `float64` | The intercept on the original scale of the design as supplied, the standardisation inverted. | +| `Coefficients` | `[]float64` | One slope per design column, in order. | +| `Fitted` | `[]float64` | The fitted value of every design row. | +| `Residuals` | `[]float64` | The residual of every row, aligned with the design. | +| `ColumnMeans` | `[]float64` | The column means the standardisation subtracted, recorded so the fit can be reproduced. | +| `ColumnScales` | `[]float64` | The column population standard deviations it divided by. | +| `Lambda` | `float64` | The penalty the fit ran at. | +| `Alpha` | `float64` | The elastic net mixing it ran with. | +| `Iterations` | `int` | The full coordinate-descent cycles taken. | +| `Converged` | `bool` | Whether no slope moved more than the tolerance in the last cycle. An exhausted budget returns the fit found so far with `Converged` false: coordinate descent on this convex objective cannot diverge. | + +**`LassoPathResult`** : the warm-started regularisation path, returned by `LassoPath`. + +| Field | Type | Meaning | +|---|---|---| +| `Alpha` | `float64` | The mixing the path ran with. | +| `Lambdas` | `[]float64` | The penalty values in descending order: 100 log-spaced values from the largest penalty that zeroes every slope down to its thousandth. For alpha = 0 the pure-lasso grid defines the same path. | +| `Intercepts` | `[]float64` | One intercept per lambda, on the original scale, in the order of `Lambdas`. | +| `Coefficients` | `[][]float64` | One slope vector per lambda, in the same order. | +| `Iterations` | `[]int` | The coordinate-descent cycles each lambda needed, the evidence the warm start earns its keep. | +| `Converged` | `bool` | Whether every fit on the path converged within the iteration budget. | + +**`HuberRegressionResult`** : the Huber M-estimate of the linear model, returned by `HuberRegression` and `HuberRegressionTuned`. + +| Field | Type | Meaning | +|---|---|---| +| `Coefficients` | `[]float64` | The M-estimates β̂, one per design column, in the design's order. | +| `StandardErrors` | `[]float64` | The asymptotic standard errors from σ²·(XᵀWX)⁻¹ with the final weights and the robust scale σ in place of the residual standard deviation. | +| `Weights` | `[]float64` | The final IRLS weights, one per observation: exactly 1 inside the band and tapering as `band/|r|` outside it. The contaminated observations sit at the bottom of the list. | +| `Scale` | `float64` | The final robust scale σ, 1.4826 times the median absolute deviation of the residuals. | +| `Fitted` | `[]float64` | The fitted value of every design row. | +| `Residuals` | `[]float64` | The residual of every row, aligned with the design. | +| `Iterations` | `int` | The reweighting steps taken. | +| `Converged` | `bool` | Whether the coefficient updates fell under the tolerance. | + +**`QuantileRegressionResult`** : a quantile regression fit, returned by `QuantileRegression`. + +| Field | Type | Meaning | +|---|---|---| +| `Coefficients` | `[]float64` | The quantile estimates β̂, one per design column; an intercept column is estimated like any other coefficient. | +| `Fitted` | `[]float64` | The fitted value of every design row. | +| `Residuals` | `[]float64` | The residual of every row, aligned with the design. | +| `Tau` | `float64` | The quantile the fit minimises the check loss for. | +| `CheckLoss` | `float64` | The minimised check loss `Σ ρ_τ(r)` at the fit. | +| `Objective` | `[]float64` | The best check loss seen after every interior-point iteration, from the least squares start on. Monotone non-increasing by construction. | +| `Iterations` | `int` | The interior-point iterations taken. | +| `Converged` | `bool` | Whether the run settled by its own stopping rules. | + +**`PCAResult`** : the decomposition of an observation array, returned by `PCA`. + +| Field | Type | Meaning | +|---|---|---| +| `Mean` | `[]float64` | The column means of the observations the fit ran on. | +| `Loadings` | `*tensor.Array` | The (p, p) rotation: entry (j, k) is the loading of variable j on component k, columns ordered by falling explained variance and orthonormal as columns. Every column's largest-magnitude loading is positive, the first index winning a tie, so two runs on the same data agree sign for sign. | +| `Scores` | `*tensor.Array` | The (n, p) coordinates of the observations on the components: the centred observations times the loadings. | +| `ExplainedVariance` | `[]float64` | Each component's eigenvalue of the covariance, in the loadings' order. | +| `ExplainedVarianceRatio` | `[]float64` | `ExplainedVariance` divided by the total variance, so the ratios sum to one. | +| `Whitening` | `*tensor.Array` | The (p, p) transform taking a centred row to the unit-covariance representation. Nil when the covariance is rank deficient, where no such transform exists. | +| `Unwhitening` | `*tensor.Array` | The (p, p) inverse of `Whitening`. Nil on a rank-deficient fit. | + +**`KMeansResult`** : the fit of k-means over a sample, returned by `KMeans`. + +| Field | Type | Meaning | +|---|---|---| +| `Centres` | `[][]float64` | The k fitted centroids, one row of d coordinates each. | +| `Labels` | `[]int` | The cluster of every sample row, 0 to k-1. | +| `Inertia` | `float64` | The within-cluster sum of squared distances to the centres, the objective the Lloyd loop minimises. | +| `Iterations` | `int` | The Lloyd sweeps taken. | +| `Converged` | `bool` | Whether the centre movement fell under the tolerance before the iteration budget ran out. | + +**`GaussianMixtureResult`** : the fit of a Gaussian mixture, returned by `GaussianMixture` and `GaussianMixtureBIC`. + +| Field | Type | Meaning | +|---|---|---| +| `Components` | `int` | The fitted component count. | +| `Weights` | `[]float64` | The mixture weights, in component order, summing to 1. | +| `Means` | `[][]float64` | One mean vector per component. | +| `Covariances` | `[][]float64` | One d-by-d covariance per component, row-major, in component order. | +| `Responsibilities` | `[]float64` | The final E-step posteriors, n rows of `Components` entries each, row-major: the posterior probability of component j given sample i. | +| `LogLikelihood` | `float64` | The maximised observed-data log likelihood. | +| `BIC` | `float64` | `-2·LogLikelihood + p·ln n` at the fitted parameters, p the count of free parameters. | +| `Iterations` | `int` | The EM sweeps taken. | +| `Converged` | `bool` | Whether the log likelihood settled under the tolerance before the budget. | +| `BICGrid` | `[]float64` | Filled by `GaussianMixtureBIC` only: the BIC of every fit on the component grid 1 to `len(BICGrid)`, in grid order. The plain fit leaves it nil. | + +**`GaussianProcessResult`** : the posterior of a Gaussian process at the test points, returned by `GaussianProcessRegression`. + +| Field | Type | Meaning | +|---|---|---| +| `Mean` | `[]float64` | The posterior mean function at the test points. | +| `Covariance` | `[]float64` | The posterior covariance between the test points, m-by-m row-major in the order the test points were given. | +| `Variance` | `[]float64` | The diagonal of `Covariance`, the marginal posterior variance per test point, clamped at zero. | +| `LogLikelihood` | `float64` | The log marginal likelihood of the training observations under the prior. | + +**`LinearMixedModelResult`** : the fit of a linear mixed model, returned by `LinearMixedModel`. + +| Field | Type | Meaning | +|---|---|---| +| `Coefficients` | `[]float64` | The fitted fixed effects β̂, one per column of the fixed design. | +| `StandardErrors` | `[]float64` | Their estimated standard deviations, the square roots of the diagonal of the GLS covariance (XᵀV⁻¹X)⁻¹ at the fitted components. | +| `RandomEffects` | `[][]float64` | The posterior mean b̂_g per group, each the length of a row of the random design, in `GroupLabels` order. | +| `GroupLabels` | `[]int` | The distinct group labels in the order the fit met them. | +| `RandomCovariance` | `[]float64` | The fitted between-group covariance Σ̂ of the random effects, q×q row-major. | +| `ResidualVariance` | `float64` | The fitted σ̂². | +| `LogLikelihood` | `float64` | The maximised REML log likelihood. | +| `Fitted` | `[]float64` | The conditional fit X·β̂ + Z·b̂ per row. | +| `Residuals` | `[]float64` | The response less `Fitted`, aligned with the rows. | +| `Iterations` | `int` | The EM sweeps taken. | +| `Converged` | `bool` | Whether the log likelihood settled under the tolerance before the budget. | + +**`HiddenMarkovFitResult`** : the Baum-Welch fit of a hidden Markov model, returned by `FitHiddenMarkovModel`; the fitted `HiddenMarkovModel` carries `Initial` (states), `Transition` (states×states) and `Emission` (states×symbols), all row-major distributions. + +| Field | Type | Meaning | +|---|---|---| +| `Model` | `*HiddenMarkovModel` | The fitted model. | +| `LogLikelihood` | `float64` | The training sequence's log likelihood under the fitted model. | +| `Iterations` | `int` | The Baum-Welch sweeps taken. | +| `Converged` | `bool` | Whether the log likelihood settled under the tolerance before the budget. | + +**`Dendrogram`** : the merge record of an agglomerative clustering, returned by `HierarchicalClustering`. + +| Field | Type | Meaning | +|---|---|---| +| `Left` | `[]int` | The smaller cluster id of each merge; leaves are the sample's rows and the merge at step t creates cluster n+t. | +| `Right` | `[]int` | The larger cluster id of each merge. | +| `Heights` | `[]float64` | The merge distances in merge order, non-decreasing except under `CentroidLinkage`, where an inversion is legitimate. | +| `Sizes` | `[]int` | The row count of each merge's cluster. | + +#### Errors + +Every entry point returns its value with an error, the sole exception being `NormalCDF`, which returns a bare `float64`. Every error carries the library's `tensor: ` prefix and names the entry point that raised it. + +- An empty or otherwise unusable sample: an error naming what is missing, for example `Quantile: empty array has no quantiles`. +- A complex array where a real one is required: an error naming the entry point, for example `Median: complex arrays have no median`. +- A non-finite observation, NaN or ±Inf: an error naming the value, for example `Histogram: sample 3 is not finite (NaN)`. The estimation entries check this before any arithmetic, so a NaN never reaches a result. +- A length, rank or shape mismatch: an error naming both counts or the offending shape, for example `LinearRegression: the design has 6 rows but the response 5`. +- A parameter outside its documented domain, such as a non-positive rate or shape, a probability outside `(0, 1)`, a `q` outside `[0, 1]`, an `alpha` outside `[0, 1]` or a non-positive bandwidth: an error naming the parameter and the value received. +- A quantile at `q = 0` or `q = 1` of a continuous law: an error, since neither has a finite quantile. +- A design that is rank deficient, near-collinear or not `n > p`: an error naming the condition, for example `LinearRegression: the design is rank deficient`. +- A fit that would not settle: `LogisticRegression` reports a perfectly separable response, `PoissonRegression` and `HuberRegression` report an exhausted iteration budget, and the incomplete gamma and beta functions report an iteration that missed its convergence budget. +- An out-of-range search: a discrete quantile whose bracket never reaches `q`, and a sample larger than `TheilSenMaxObservations`, both of which name the cost rather than answer approximately. + +#### Workflow + +The flagship sequence: build the design, fit it, and quote an interval using the package's own distribution quantile. + +```go +// y on the design, the intercept carried as the constant first column. +design, _ := tensor.FromFloats([]float64{ + 1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, +}, 6, 2) +y, _ := tensor.FromFloats([]float64{4.1, 7.0, 9.9, 13.2, 15.9, 19.1}, 6) + +fit, err := stats.LinearRegression(design, y) +if err != nil { + return err +} +// The two-sided 95 percent critical value on the residual degrees of +// freedom. +critical, err := stats.StudentTQuantile(0.975, fit.DResidual) +if err != nil { + return err +} +slope := fit.Coefficients[1] +se := fit.StandardErrors[1] +fmt.Printf("slope %.3f, 95%% CI [%.3f, %.3f], R² %.4f\n", + slope, slope-critical*se, slope+critical*se, fit.RSquared) +``` + +The same shape carries every model of the package: `LinearRegression`, `LogisticRegression`, `PoissonRegression`, `ElasticNet`, `HuberRegression` and `QuantileRegression` all take the design first with the intercept as a constant column, return a result struct with the coefficients and their uncertainty, and leave the caller to choose a test or an interval from the distribution functions. + +### Package `optim` + +Fitting, minimisation, programming and root finding: Levenberg-Marquardt +nonlinear least squares, the Nelder-Mead simplex and limited-memory BFGS for +local minima, the augmented Lagrangian for constrained problems, the two-phase +revised simplex and a primal active-set method for linear and quadratic +programs, three derivative-free global searchers, and scalar and vector root +finders. Every entry point hands the caller's function the candidate point as a +rank-1 `*Array` and returns a fresh array the caller owns; an objective, +residual or gradient callback receives a copy, never a view of a reused +buffer. + +The package is real-valued: a complex starting point, cost, right-hand side, +bound, constraint matrix or callback payload is refused with an error rather +than read through a nil payload. It also leaves the modelling to the caller. +`MinimiseLinear` requires the standard form `A·x = b`, `x ≥ 0` exactly and +`MinimiseLinearRows` is the wrapper that converts the house two-sided rows into +it; `MinimiseDifferentialEvolution` clamps to its box and requires one, while +`MinimiseCMAES` has no bounds at all; a local minimum is the minimum of the +basin the solver started in, and multistart or a global searcher is what finds +the others. + +#### Curve fitting + +| Call | What it does | +|---|---| +| `LevenbergMarquardt(residual, p0, opts)` | Minimises `‖r(p)‖²` over the parameter vector by LM damping of the Gauss-Newton step, returning the parameters and the residual sum of squares. Refuses an underdetermined problem, a residual whose length changes mid-fit and a damping that collapses without meeting the tolerance; a singular solve and an exhausted budget keep the historical error contract here. | +| `LevenbergMarquardtFit(residual, p0, opts)` | The same fit with the full report: a `FitResult` carrying the parameters, the (weighted) χ², the `FitStatus` (`FitConverged`, `FitStalled`, `FitBudget`) and, when requested, the parameter covariance. A singular solve and a collapsed damping come back as `FitStalled` on the best point reached and a spent budget as `FitBudget`, never as an error; the errors here are the model's own fault and, with `RequestCovariance`, a rank-deficient Jacobian at the answer. The `GradTol` and `StepTol` options converge the fit on the gradient norm and on the step size, the exits that reach the flat optimum the χ² tolerance alone never leaves. | + +#### Local minimisation + +| Call | What it does | +|---|---| +| `Minimise(f, x0, opts)` | Returns the point and value of a local minimum of `f` near `x0` by the Nelder-Mead simplex, which needs no derivatives. | +| `MinimiseLBFGS(f, grad, x0, opts)` | Returns the point and value of a local minimum by limited-memory BFGS with a backtracking Armijo line search; `grad` may be nil, and the box walls in `opts` project the iterate, so the objective is never evaluated outside them. Refuses a vanished or non-descent search direction, a stalled line search and an exhausted budget. | + +#### Constrained minimisation + +| Call | What it does | +|---|---| +| `MinimiseConstrained(f, grad, x0, cons, opts)` | Returns the point and value of a local minimum of `f` subject to the box walls carried by `opts` and the rows `l ≤ A·x ≤ u` of `cons`, each outer round minimising the augmented Lagrangian and then ascending the row multipliers. A start outside the feasible set is fine; failure to reach feasibility is an error naming the worst remaining violation. | +| `MinimiseNonlinearConstrained(f, grad, x0, cons, opts)` | Returns the point, the value of `f` and the row multipliers of a local minimum subject to the functional rows of `cons` (equalities `g(x) = 0` first, then inequalities `h(x) ≤ 0`, in declaration order). With both slices empty it delegates to `MinimiseLBFGS` and returns nil multipliers. | + +#### Linear and quadratic programming + +| Call | What it does | +|---|---| +| `MinimiseLinear(c, a, b, opts)` | Returns the point and value of the minimum of `c·x` over the standard-form polytope `A·x = b`, `x ≥ 0` by the two-phase revised simplex under Bland's rule. Refuses an infeasible problem with the phase-1 evidence, an unbounded objective with the profitable column, and an exhausted pivot budget. | +| `MinimiseLinearRows(c, cons, opts)` | The same, over the two-sided rows `l ≤ A·x ≤ u` of `cons`, converting them to the standard form mechanically (each variable splits into two non-negative columns, each finite row side gains a slack). Refuses exactly as `MinimiseLinear` does. | +| `MinimiseQP(h, c, cons, x0, opts)` | Returns the point, the value `½xᵀHx + c·x` and one multiplier per row of the minimum of the strictly convex quadratic over the two-sided rows. Refuses a non-symmetric or non-positive-definite `H` at entry, naming the failed pivot, and an infeasible `x0`. | + +#### Global minimisation + +| Call | What it does | +|---|---| +| `MinimiseDifferentialEvolution(f, lower, upper, opts)` | Returns the best point and value of the global minimum over the box `[lower, upper]` by rand/1/bin differential evolution with clamping to the bounds. Refuses mismatched or degenerate bounds and a population below four. The generation budget is a tuning parameter, so running it out returns the best point found without an error. | +| `MinimiseCMAES(f, x0, opts)` | Returns the best point and value the covariance matrix adaptation strategy found, one run without restarts. No bounds: the caller who needs a box reparametrises. An exhausted generation budget is an error unless `AllowBudgetExit` is set. | +| `MinimiseSimulatedAnnealing(f, x0, opts)` | Returns the best point and value a Metropolis walk with a Gaussian proposal and geometric cooling found. Finds basins, not minima to full precision; running `MinimiseLBFGS` from the returned point is the standard composition. A schedule that ended while the best value was still improving is an error unless `AllowBudgetExit` is set. | + +#### Scalar root finding + +| Call | What it does | +|---|---| +| `FindRoot(f, a, b, tol)` | Returns a root of `f` in the bracket `[a, b]` by Brent's method; `f(a)` and `f(b)` must be finite with opposite signs. The budget is fixed at 200 iterations. | +| `FindRootBrent(f, a, b, opts)` | The same method with the tolerance and the budget under the caller's control, in the vocabulary of `RootSystemOptions`: the tolerance thresholds `|f|` at the returned point, and a run costs at most `MaxIterations + 2` evaluations. An exhausted budget is an error, never a silent guess. | +| `FindRootNewton(f, df, x0, tol, maxIter)` | Returns a root near `x0` by Newton's iteration with the supplied derivative. A vanishing derivative or an exhausted budget is an error. | + +#### Systems of equations + +| Call | What it does | +|---|---| +| `FindRootSystem(f, x0, opts)` | Solves `r(x) = 0` for a vector function `r` of an `n`-vector, returning the solution and the residual infinity norm at it. The residual must be a vector of the same length as `x0`. A singular Jacobian falls back to the steepest-descent direction of the residual norm; a residual that cannot be reduced by damping, and an exhausted iteration budget, are errors. | + +#### Options and results + +**`LMOptions`**: tunes `LevenbergMarquardt` and `LevenbergMarquardtFit`. The +damping factor is scaled by 0.3 after every accepted step and by 10 after every +rejected one; a damping that passes 1e20 ends the run as a collapse. The +tolerance is the relative χ² improvement threshold. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `MaxIterations` | `int` | `≤ 0` means `200` | The number of Gauss-Newton steps the fit may take. | +| `Tolerance` | `float64` | `≤ 0` means `1e-10` | A step is accepted as converged when the χ² improvement falls below `Tolerance·(1 + χ²)`. | +| `Lambda` | `float64` | `≤ 0` means `1e-3` | The initial damping factor. | +| `Jacobian` | `func(p *core.Array) (*core.Array, error)` | nil means central differences, at two residual evaluations per parameter per iteration | The analytic Jacobian of the residual, an observations × parameters matrix with row `i` holding `∂r_i/∂p_j`. A matrix of the wrong shape is refused. | +| `AllowBudgetExit` | `bool` | `false` | Reports the last point with a nil error when the iteration budget runs out instead of refusing it. | +| `ParallelJacobian` | `bool` | `false` | Lets the central-difference Jacobian sweep its columns on several goroutines. Setting it is the caller's consent that the residual callback may run concurrently from more than one goroutine; the default keeps every evaluation on the caller's goroutine. The columns are independent, so the fit is bit-identical either way. No effect while `Jacobian` supplies the analytic matrix. | +| `GradTol` | `float64` | `≤ 0` disables | Converges the fit once the gradient norm `‖Jᵀr‖∞` falls to it, the test that reaches the flat optimum where χ² still falls in slivers while the step directions carry no information. | +| `StepTol` | `float64` | `≤ 0` disables | Converges the fit once an accepted step's infinity norm falls to `StepTol·(‖p‖∞ + StepTol)`, the relative test that stops a fit whose parameters have stopped moving meaningfully. | +| `Sigma` | `*core.Array` | nil | The measurement covariance: a vector of one positive variance per residual, or an exactly symmetric positive definite `nR×nR` matrix. The fit whitens the residuals and the Jacobian through the factor once, χ² becomes `rᵀC⁻¹r` and a requested covariance becomes `(JᵀC⁻¹J)⁻¹`. | +| `RequestCovariance` | `bool` | `false` | Fills `FitResult.Covariance` with `(JᵀJ)⁻¹`, or `(JᵀC⁻¹J)⁻¹` under `Sigma`, at the returned point, at the cost of one more Jacobian there. A rank-deficient Jacobian at the answer has no covariance to report and the run fails naming it. | + +**`MinimiseOptions`**: tunes `Minimise`. Both convergence tests are absolute +in the objective's own scale: the value spread is compared against +`Tolerance·max(1, |f|)` and the simplex diameter against +`Tolerance·max(1, |x|)`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `MaxIterations` | `int` | `≤ 0` means `2000` | The number of simplex rounds. | +| `Tolerance` | `float64` | `≤ 0` means `1e-10` | The threshold on both the value spread and the simplex diameter, which must be met together for the run to count as converged. | +| `InitialStep` | `float64` | `≤ 0` means `1` | The offset of each simplex vertex from `x0`, scaled by `max(1, |x_i|)`. | +| `AllowBudgetExit` | `bool` | `false` | Reports the best vertex with a nil error when the budget runs out instead of refusing it. | + +**`LBFGSOptions`**: tunes `MinimiseLBFGS` and the inner solves of +`MinimiseConstrained` and `MinimiseNonlinearConstrained`. The tolerance is the +L∞ norm of the projected gradient, which is the KKT residual of the box +problem, and it is absolute in the gradient's own scale. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `MaxIterations` | `int` | `≤ 0` means `10000` | The number of quasi-Newton steps. | +| `Tolerance` | `float64` | `≤ 0` means `1e-8` | The threshold on the projected gradient's infinity norm. | +| `Memory` | `int` | `≤ 0` means `10` | The number of correction pairs kept; capped at the number of variables. | +| `Lower` | `[]float64` | nil opens every lower side | Coordinate-wise lower walls, one entry per variable; an infinite entry opens that side. | +| `Upper` | `[]float64` | nil opens every upper side | Coordinate-wise upper walls, one entry per variable. A NaN wall or a crossed pair is an error; `x0` is projected onto the box rather than refused. | +| `AllowBudgetExit` | `bool` | `false` | Reports the best point with a nil error when the budget runs out instead of refusing it. `MinimiseConstrained` sets it for its inner solves. | + +**`LinearConstraints`**: carries the rows of `l ≤ A·x ≤ u` for +`MinimiseConstrained`, `MinimiseLinearRows` and `MinimiseQP`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `A` | `*core.Array` | required | The `r×n` constraint matrix over the `n` variables. A nil matrix, a complex one, a rank other than 2, a second dimension other than `n` and a non-finite coefficient are all errors. | +| `Lower` | `[]float64` | required, one entry per row | The lower side of each row; an infinite entry opens that side. A NaN entry or a pair with `Lower > Upper` is an error. | +| `Upper` | `[]float64` | required, one entry per row | The upper side of each row. A row with `Lower = Upper` is an equality constraint. | + +**`NonlinearConstraints`**: carries the functional rows for +`MinimiseNonlinearConstrained`. Each callback receives the candidate point as a +rank-1 array and must return a finite value; an error it returns is fatal for +the run. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Equalities` | `[]func(*core.Array) (float64, error)` | empty | The rows `g(x) = 0`. Their signed multipliers come first in the returned slice. | +| `Inequalities` | `[]func(*core.Array) (float64, error)` | empty | The rows `h(x) ≤ 0`. Their non-negative multipliers follow the equality ones. A row that is slack at the answer has an estimate that decays to zero. With both slices empty the call delegates to `MinimiseLBFGS` and returns nil multipliers. | + +**`LinearProgramOptions`**: tunes `MinimiseLinear` and `MinimiseLinearRows`. +The tolerance prices reduced costs and separates ratio-test ties, and it is +absolute in the scale the caller's costs and rows carry. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `MaxIterations` | `int` | `≤ 0` means `10000` | The number of simplex pivots. | +| `Tolerance` | `float64` | `≤ 0` means `1e-9` | The threshold on reduced costs and on the ratio-test ties. | + +**`QPOptions`**: tunes `MinimiseQP`. The tolerance is the threshold on the +KKT step below which the iterate is stationary on its working set, on the +multiplier that releases a row, and on the symmetry of `H`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `MaxIterations` | `int` | `≤ 0` means `1000` | The number of active-set rounds. | +| `Tolerance` | `float64` | `≤ 0` means `1e-10` | The KKT step, multiplier-release and symmetry threshold. | + +**`DifferentialEvolutionOptions`**: tunes `MinimiseDifferentialEvolution`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Population` | `int` | `≤ 0` means `15·d`, at least `4` | The number of population members. A resolved value below 4 is an error, because the scheme draws three distinct others besides the target. | +| `F` | `float64` | `≤ 0` means `0.7` | The differential weight of the mutant vector. | +| `CR` | `float64` | `≤ 0` means `0.9` | The binomial crossover probability. | +| `Generations` | `int` | `≤ 0` means `1000` | The generation budget. Running it out returns the best point found, with no error. | +| `Seed` | `int64` | `0` is replaced by `42` | The xoshiro stream seed. Any other value, negatives included, seeds the stream directly. | + +**`CMAESOptions`**: tunes `MinimiseCMAES`. The tolerance ends the run when +either the largest principal axis of the search distribution, +`sigma·sqrt(max C_ii)`, has fallen to `Tolerance·max(1, ‖mean‖∞)`, or the +objective spread over one generation has fallen to `Tolerance·max(1, |best|)`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Sigma0` | `float64` | `≤ 0` means `0.3` | The initial step size, the tutorial's typical value for problems scaled to O(1). | +| `Generations` | `int` | `≤ 0` means `500` | The generation budget. | +| `Tolerance` | `float64` | `≤ 0` means `1e-12` | The collapse and flat-landscape threshold. | +| `Seed` | `int64` | `0` is replaced by `42` | The xoshiro stream seed, as in `MinimiseDifferentialEvolution`. | +| `AllowBudgetExit` | `bool` | `false` | Reports the best point with a nil error when the generation budget runs out instead of refusing it. | + +**`SimulatedAnnealingOptions`**: tunes `MinimiseSimulatedAnnealing`. The +schedule is the algorithm: the temperature runs from `Temperature0` down by +the factor `CoolingRate` at every proposal. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Steps` | `int` | `≤ 0` means `20000` | The number of proposals. | +| `Temperature0` | `float64` | `≤ 0` means `1` | The starting temperature. | +| `CoolingRate` | `float64` | `≤ 0` means `0.9995`, and a value above `1` is an error | The geometric factor per proposal. | +| `StepScale` | `float64` | `≤ 0` means `0.1` | The relative Gaussian proposal scale: coordinate `i` moves by `StepScale·max(1, |x_i|)` standard normals. | +| `Seed` | `int64` | `0` is replaced by `42` | The xoshiro stream seed, as in `MinimiseDifferentialEvolution`. | +| `Tolerance` | `float64` | `≤ 0` means `1e-6` | The convergence test on the frozen tail: the best value may not improve by more than `Tolerance·max(1, |best|)` over the final quarter of the schedule. | +| `AllowBudgetExit` | `bool` | `false` | Reports the best point when the schedule ended while the value was still improving, instead of refusing it. | + +**`BrentOptions`**: tunes `FindRootBrent`, in the vocabulary of +`RootSystemOptions`. The tolerance is a threshold on `|f|` at the returned +point, the scalar counterpart of the residual infinity norm. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Tolerance` | `float64` | `≤ 0` means `1e-10` | The accepted `|f|` at the returned point. | +| `MaxIterations` | `int` | `≤ 0` means `100` | The iteration budget; a run costs at most `MaxIterations + 2` evaluations of `f`. An exhausted budget is an error. | + +**`RootSystemOptions`**: tunes `FindRootSystem`. The tolerance is an +infinity-norm threshold on both the residual and the scaled step. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Tolerance` | `float64` | `≤ 0` means `1e-10` | The threshold the residual and the scaled step must both meet. | +| `MaxIterations` | `int` | `≤ 0` means `100` | The number of damped Newton rounds. | +| `UseBroyden` | `bool` | `false` | Builds the central-difference Jacobian once and carries its inverse between steps by the rank-one Broyden update, with a numerical rebuild whenever the update degrades. The default leaves the per-step Jacobian and the iteration's results exactly as they are. | +| `ParallelJacobian` | `bool` | `false` | Lets the central-difference Jacobian sweep its columns on several goroutines, the initial build and every Broyden rebuild included. Setting it is the caller's consent that the residual callback may run concurrently from more than one goroutine; the default keeps every evaluation on the caller's goroutine. The columns are independent, so the run is bit-identical either way. | + +#### Budget policy + +An iteration budget is a refusal, not an answer. `Minimise`, `MinimiseLBFGS`, +`LevenbergMarquardt`, `MinimiseCMAES`, `MinimiseSimulatedAnnealing`, +`MinimiseLinear`, `MinimiseLinearRows`, `FindRootBrent`, `FindRootNewton` and +`FindRootSystem` each return an error when the budget runs out before the +tolerance is met, naming the figure they reached and the tolerance they fell +short of, so a budget stop is never mistaken for a converged answer. Setting +`AllowBudgetExit` on the corresponding options reports the best point reached +instead, with a nil error; the default is `false` everywhere, and +`MinimiseConstrained` and `MinimiseNonlinearConstrained` set it only for their +inner solves, whose accuracy the outer loop's feasibility check judges. +`LevenbergMarquardtFit` needs no flag either: its budget stop is reported as +`FitBudget` on the last point. +`MinimiseDifferentialEvolution` needs no flag: its generation budget tunes the +search rather than deciding convergence, so running it out returns the best +point found. `FindRoot` has a fixed budget of 200 iterations and no options +struct at all. + +#### Errors + +Every error carries the library's `tensor: ` prefix. + +- A complex starting point, cost, right-hand side, bound or callback payload: refused by the entry points that read the value directly (`complex starting points are not supported`, `complex costs are not supported`, `complex constraint matrices are not supported`, and so on). +- A non-finite entry in a cost, right-hand side, bound, constraint coefficient or Hessian, and a NaN or crossed pair of walls: refused at entry, naming the row or the variable. +- A callback that returns a non-finite objective, residual, constraint value or gradient: an error naming the value it saw, never a run that quietly reads as converged. +- A callback whose array has the wrong shape or length (a gradient for `n` variables, a Jacobian of `nR×nP`, a residual of the same length as `x0`, an `r×n` constraint matrix): refused with the shape it received. +- An underdetermined least-squares problem (`nR < nP`), a residual whose length changes mid-fit, a damping that collapses, and a vanished or non-descent search direction: refused with the diagnosis. The `LevenbergMarquardtFit` surface reports the singular solve and the collapse as `FitStalled` on the best point reached instead of refusing them. +- A `Sigma` that is not a vector of positive variances or an exactly symmetric positive definite matrix: refused at entry, naming the offending element or pair. +- A bracket that does not change sign, or one that evaluates to a non-finite value: refused by `FindRoot` and `FindRootBrent`. +- An infeasible linear program: refused with the phase-1 infeasibility and the row that carries the worst of it; an unbounded one with the column that prices out as a profitable ray. +- A non-symmetric or non-positive-definite `H` in `MinimiseQP`: refused at entry with the failed pivot named; an `x0` that violates a row by more than 1e-9 is refused with the worst violation. +- An exhausted iteration budget in any solver whose `AllowBudgetExit` is `false`: refused, as described above. +- A failure to reach feasibility in 40 augmented-Lagrangian rounds: refused with the worst remaining row violation. Feasibility itself is judged against a fixed absolute threshold of `1e-10`, independent of the inner solver's tolerance. + +#### Workflow + +```go +package main + +import ( + "fmt" + "math" + + "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/optim" +) + +// Fit y = a·exp(−b·t) to measurements, then report the parameters and +// the residual sum of squares the fit reached. +func main() { + ts := []float64{0, 0.5, 1, 1.5, 2, 2.5, 3} + ys := []float64{2.50, 1.31, 0.69, 0.36, 0.19, 0.10, 0.05} + + // The residual holds one entry per observation: model(p, t) − y. + 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) - ys[i] + } + return tensor.FromFloats(out, len(out)) + } + + p0, _ := tensor.FromFloats([]float64{1, 1}, 2) // the starting guess + p, chi2, err := optim.LevenbergMarquardt(residual, p0, optim.LMOptions{}) + if err != nil { + fmt.Println("fit:", err) + return + } + fmt.Printf("a = %.4f, b = %.4f, chi2 = %.3e\n", p.FloatAt(0), p.FloatAt(1), chi2) + + // A one-dimensional objective inside a box. The walls are the + // contract: the start below is projected onto them, and the answer + // stops at the wall the gradient pushes against, which is the KKT + // point of the box problem. + x0, _ := tensor.FromFloats([]float64{-5}, 1) + bounded := optim.LBFGSOptions{Lower: []float64{0}, Upper: []float64{2}} + x, f, err := optim.MinimiseLBFGS( + func(v *tensor.Array) (float64, error) { + d := v.FloatAt(0) - 3 + return d * d, nil + }, + nil, x0, bounded) + if err != nil { + fmt.Println("minimisation:", err) + return + } + fmt.Printf("x = %.4f, f = %.3e\n", x.FloatAt(0), f) +} +``` + +### Package `io` + +Reading and writing the formats scientific data arrives in: comma-separated +text, FITS images and tables, HDF5, NetCDF classic and native-endian memory +maps. Every loader returns the library's own `Array`, so a file read is an +ordinary value the rest of the library operates on, and every writer takes one. + +Each format is covered by a documented subset, and what lies outside it is +refused with an error naming the feature rather than half-read or guessed at. +The subset, and every refusal, is written down under [Format +support](#format-support). + +#### CSV + +| Call | What it does | +|---|---| +| `LoadCSV(path, skipHeader)` | Reads a 2-D float64 array from a CSV file. Every row must have the same number of fields; a header row is data unless `skipHeader` is true. An empty file gives a zero-shaped array. | +| `LoadCSVReader(r, skipHeader)` | Reads the same from any `io.Reader`. A leading UTF-8 byte order mark is skipped. A field that is not a number is refused with its row and column. | +| `SaveCSV(path, a)` | Writes a 2-D array to `path` as comma-separated values. The close error is part of the write. | +| `SaveCSVWriter(w, a)` | Writes a 2-D array as CSV to `w`. Complex arrays are refused, as is an array with no columns, whose CSV form would read back as a zero-shaped array. | + +Values go out with the shortest text that round-trips the float64. + +#### FITS images + +| Call | What it does | +|---|---| +| `SaveFITS(path, a, headers)` | Writes a float64 or float32 array as a FITS primary image with the given header entries. Keywords are uppercased and must be 1 to 8 characters of `A-Z`, `0-9`, `-` and `_`; the reserved ones are refused. Values go out as FITS strings of at most 68 characters after quote escaping. | +| `LoadFITS(path)` | Reads a FITS primary image into a float64 (BITPIX -64) or float32 (BITPIX -32) array and returns every non-structural header entry with it. `NAXIS = 0` gives an empty array. | + +`SIMPLE`, `BITPIX`, `NAXIS`, `NAXISn`, `EXTEND`, `END`, `COMMENT`, `HISTORY`, +`BSCALE` and `BZERO` are reserved and refused as user keywords. `LoadFITS` +applies the `BSCALE`/`BZERO` affine map (physical = raw*scale + zero) whenever +the keywords differ from the identity; float32 values are computed in float64 +and rounded once. `SaveFITS` refuses those two because it writes physical +values, and a scaling card would make a conforming reader scale them a second +time. + +#### FITS tables + +| Call | What it does | +|---|---| +| `SaveFITSTable(path, ascii, cols, headers)` | Writes a zero-axis primary HDU followed by one table extension: a binary table by default, an ASCII table when `ascii` is set. The headers land in the extension's header beside the column cards. | +| `LoadFITSTable(path)` | Reads the first table extension (`XTENSION` `BINTABLE` or `TABLE`), skipping the primary HDU and any images before it. Numeric columns come back as arrays, character columns as string slices. | + +A binary column form is `D` (float64), `E` (float32), `K` (int64), `L` +(logical), `B` (unsigned byte), `I` (16-bit) or `J` (32-bit), each with a +repeat of one, or `nA` for a character string of n characters. The writer +stores `D`, `E`, `K`, `L` and `nA`; the reader decodes all of them and +truncates a character field at its trailing spaces. An integer column scaled +by `TSCALn`/`TZEROn` keeps the int dtype while the map stays integral and +promotes to float64 otherwise. + +The ASCII forms are `nA` (or `A`), `Iw`, `Fw.d`, `Ew.d` and `Dw.d`, where w is +the column width in characters. Cells are separated by one space and written +with `TBCOLn` column starts. A value that does not fit its declared width is +refused at write time. On read, a blank cell is skipped; the Fortran-specific +`Dw.d` exponent and the exponent-less `1.5+03` form are accepted. + +#### HDF5 + +| Call | What it does | +|---|---| +| `LoadHDF5(path)` | Reads every dataset of a file, in path order, as numeric arrays with their paths, shapes and attributes. | +| `SaveHDF5(path, datasets, groupAttrs, opts)` | Writes numeric datasets as an HDF5 file: the mirror of `LoadHDF5`. The paths build the group tree, so the file reads back with the same paths, shapes, dtypes and values. | +| `SaveHDF5Text(path, texts, opts)` | Writes fixed-length string datasets, the string side of the datatype message the reader refuses for data but accepts for attributes. These files are for other readers of the format. | + +The paths are absolute (`/scan/temperature`), each written once, and an +intermediate group is created as needed. A dataset's `Attrs` are written on +the dataset itself; the attributes of the root and of the groups come from +`groupAttrs`, keyed by group path with the root keyed `/`. Attribute values +are parsed back into typed attributes: a whole number becomes an int64 +attribute, a decimal a float64 one, a bracketed list an int64 or float64 +array, and anything else a fixed-length string. The written file is +deterministic: children are laid out in sorted name order. + +#### NetCDF classic + +| Call | What it does | +|---|---| +| `LoadNetCDF(path)` | Reads a NetCDF classic file (CDF-1 and CDF-2) into its dimensions, its variables and its global attributes. Each variable lands the core dtype its classic type code carries: `NC_BYTE` as int8 (signed in the classic model), `NC_CHAR` as uint8 raw bytes, `NC_SHORT` as int16, `NC_INT` as int32, and `NC_FLOAT` and `NC_DOUBLE` as float64. A type code beyond the classic six is refused by name. | +| `SaveNetCDF(path, dims, vars, attrs)` | Writes a NetCDF classic (CDF-1) file: dimensions, variables and text attributes. | + +Variables carry their own attributes; only the global attributes come back in +the map. Attribute keys are written in sorted order so the same inputs give +the same bytes. Names must follow the traditional grammar: the first +character alphanumeric or `_`, the rest alphanumeric or one of `_.@+-`. + +#### Memory mapping + +| Call | What it does | +|---|---| +| `MapFloats(path, offset, n)` | Maps n float64 values of `path`, starting at byte `offset`, into a read-only one-dimensional array. | +| `MapFloat32s(path, offset, n)` | Maps n float32 values, with the same contract. | +| `MapInts(path, offset, n)` | Maps n int64 values, with the same contract. | +| `SaveNativeFloats(path, values)` | Writes float64 values to `path` in the machine's native byte order, the format `MapFloats` reads back. | + +The mapped array is a live view of the file's pages: `release` unmaps it and +any use of the array afterwards is a use-after-free, so `release` comes +strictly last. The values must have been written in the machine's native byte +order (`binary.NativeEndian`); the offset must be a multiple of the element +size, which is 8 for `MapFloats` and `MapInts` and 4 for `MapFloat32s`. + +#### Options and results + +**`FITSTable`**: a parsed table extension. `Names`, `Units`, `Columns` and +`Text` run parallel to the file's column list: a numeric column holds its +values in `Columns` and nil in `Text`, a character column holds its strings in +`Text` and nil in `Columns`. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Kind` | `string` | from the file | `BINTABLE` or `TABLE`. | +| `Names` | `[]string` | from the file | The `TTYPEn` of each column. | +| `Units` | `[]string` | from the file | The `TUNITn` of each column, empty when the card is absent. | +| `Columns` | `[]*core.Array` | from the file | One array per column, nil for a character column. | +| `Text` | `[][]string` | from the file | One string slice per column, nil for a numeric column. | +| `Rows` | `int` | from the file | The row count, `NAXIS2`. | +| `Headers` | `map[string]string` | from the file | Every other header card of the extension. | + +**`FITSTableColumn`**: one column of a table to write. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Name` | `string` | required | The column name, written as `TTYPEn`. It must be 1 to 68 characters. | +| `Unit` | `string` | `""` | The unit, written as `TUNITn`; an empty unit writes no card. | +| `Form` | `string` | required | The TFORM descriptor, uppercased on write. See [FITS tables](#fits-tables) for the accepted forms. | +| `Data` | `*core.Array` | `nil` | The values of a numeric column, a vector of the dtype its form requires, one element per row. | +| `Text` | `[]string` | `nil` | The strings of a character column, one per row. A numeric column must carry `Data`, and a `Text` on it is ignored; a character column must carry `Text` and must not carry `Data`. | + +**`HDF5Dataset`**: one dataset of a file: its path from the root, its shape +(row-major, as the file stores it), its values and its attributes. The +attribute map also carries the attributes of the groups the dataset sits in, +the nearest one winning. On read, `Shape` and `Values` are filled. On write, +the shape is taken from `Values`; a non-nil `Shape` that disagrees with them +is refused. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Path` | `string` | required | The absolute path from the root, `/scan/temperature`. | +| `Shape` | `[]int` | `nil` | The shape to write; nil means the shape of `Values`. | +| `Values` | `*core.Array` | required | The values, int64, float32 or float64. A complex array is refused. | +| `Attrs` | `map[string]string` | `nil` | The attributes of the dataset, written on it. | + +**`HDF5TextDataset`**: one fixed-length string dataset for `SaveHDF5Text`: +`Text` holds the elements in row-major order, each padded to the longest +element in the file. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Path` | `string` | required | The absolute path from the root. | +| `Shape` | `[]int` | required | The shape of the dataset; the element count must equal `len(Text)`. | +| `Text` | `[]string` | required | The elements in row-major order. A NUL byte cannot be carried by a fixed-length element and is refused. | + +**`HDF5WriteOptions`**: tunes `SaveHDF5` and `SaveHDF5Text`. The zero value +writes the classic file layout with contiguous datasets. When several option +values are passed, the last one wins. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Latest` | `bool` | `false` | Writes superblock version 3 with version 2 object headers: groups become link messages and every structure the format checksums carries a lookup3 sum. Latest files hold contiguous datasets only, so combining `Latest` with a filter is refused. | +| `Gzip` | `int` | `0` | Applies the deflate filter at the given level: 0 disables it, -1 is the default level and 1 to 9 are the levels of the format. A level outside that set is refused. A filtered dataset is stored in chunks. | +| `Shuffle` | `bool` | `false` | Applies the shuffle filter before deflate, regrouping the bytes of each element so compression sees the high-order bytes together. Shuffle alone also forces chunks. | +| `ChunkBytes` | `int` | `0` | The target size of one chunk in bytes for filtered datasets; 0 selects 64 KiB. Datasets smaller than the target stay in one chunk. A negative target is refused, and a target so small that the dataset needs more than 4194304 chunks is refused with a hint to raise it. | + +**`NetCDFDim`**: one named dimension of a NetCDF classic file. A length of +zero is the record dimension (the unlimited one): only it may be zero, it must +lead the list, and its extent is the number of records, which a record +variable carries on its first axis. Reading a file and writing it back +preserves the record dimension. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Name` | `string` | required | The dimension name, under the traditional name grammar. A duplicate name is refused. | +| `Length` | `int` | `0` | The extent. A negative length is refused; zero declares the record dimension. | + +**`NetCDFVar`**: one named variable of a NetCDF classic file. Values are +row-major with the slowest dimension first, exactly as the file stores them, +and `Dims` names the dimensions in that same order. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Name` | `string` | required | The variable name, under the traditional name grammar. A duplicate name is refused. | +| `Dims` | `[]string` | `nil` | The dimension names, slowest first. Every one must be declared; a record dimension in a later position is refused. A variable with no dimensions is a scalar. | +| `Values` | `*core.Array` | required | The values: float64, float32 or int64. The element count must equal the extent of the dimensions; a record variable must hold a whole number of records. | +| `Attrs` | `map[string]string` | `nil` | The variable's attributes, written as text. | + +#### Format support + +| Format | Reads | Writes | The subset that is implemented | +|---|---|---|---| +| CSV | `LoadCSV`, `LoadCSVReader` | `SaveCSV`, `SaveCSVWriter` | 2-D arrays: the writer formats every numeric dtype the core carries (bool as 0 and 1, the integer classes as exact decimals, float16 through its exact widening). Complex arrays and arrays with no columns are refused on write. | +| FITS image | `LoadFITS` | `SaveFITS` | The primary HDU, BITPIX -64 and -32, any rank with positive axes. | +| FITS table | `LoadFITSTable` | `SaveFITSTable` | The `BINTABLE` and `TABLE` extension kinds, and nothing else. | +| HDF5 | `LoadHDF5` | `SaveHDF5`, `SaveHDF5Text` | See the two lists below. | +| NetCDF classic | `LoadNetCDF` | `SaveNetCDF` | The classic model: CDF-1 and CDF-2 read, CDF-1 written. | +| Native-endian map | `MapFloats`, `MapFloat32s`, `MapInts` | `SaveNativeFloats` | float64, float32 and int64 payloads in the machine's own byte order. | + +**HDF5 read.** Superblock generations 0 to 3. Generations 0 and 1 are the +classic layout; generations 2 and 3 carry a lookup3 checksum over the +superblock, which is verified, and name the root group by its object header +address. The address and length widths are 4 or 8 bytes. Refused by name: a +superblock version of 4 or more, a non-zero base address, the superblock +extension. + +Object headers of version 1 and version 2, version 2 headers being checksummed +and their continuation blocks with them; a header may chain through at most +512 continuation blocks. Groups are read from symbol tables (version 1 group +B-tree, local heap, symbol table nodes) and from version 1 link messages. +Refused by name: a dense group (the fractal-heap link storage), a link message +of another version, a group B-tree more than 32 levels deep. + +Datasets are read from contiguous, compact and chunked storage; a chunked +dataset in a superblock 2 or later is refused, because the reference library +indexes such chunks with the version 2 B-tree. The chunk B-tree is walked at +most 32 levels deep. Filter pipelines of version 1 are read, with the deflate +(1), shuffle (2) and fletcher32 (3) filters; the fletcher32 checksum is +verified, and any other filter identifier is refused by name. + +Datatypes: fixed-point values of 1, 2, 4 and 8 bytes land in the dtype of +their own width and sign: 1-byte signed and unsigned read as `Int8` and +`Uint8`, 2-byte as `Int16` and `Uint16`, 4-byte as `Int32` and `Uint32`, +8-byte signed as `Int`; unsigned 64-bit data is refused, because float64 +cannot hold every value of it. Floating-point values of 4 and 8 bytes read +as `Float32` and `Float`. A boolean dataset written as an HDF5 enumeration +(a one-byte unsigned base with member values 0 and 1) lands `Bool`; every +other enumeration and every bitfield is refused by class name. A big-endian +element is refused by name. Refused for data: string datasets, +variable-length data, and the remaining datatype classes. Dataspace message +versions 1 and 2 are read; a dataset with a dimension permutation is +refused. + +Attributes of an object header are read into a text map: a fixed-length string +verbatim, a variable-length string through the global heap, a numeric value +formatted as decimal, with several values as `[a, b, c]`. An attribute message +the reader cannot decode is left out of the map rather than failing the read. +An object reachable through several hard links is read once, under the first +path the traversal reaches it from, and a soft or external link is skipped, +since it names no object of this file. A hard-link cycle along the current +path and a walk deeper than 512 group levels are refused. + +A dataset of more than 2 GiB, and the datasets of a file summing past that +budget, are refused with a hint that larger datasets belong behind mmap. + +**HDF5 write.** Superblock version 0 (the default, classic layout) or 3 +(`Latest`); object headers version 1 or 2; groups as symbol tables or as link +messages; datasets contiguous, or chunked through a version 1 chunk B-tree +when a filter applies. Every address and length is eight bytes. The datatypes +written are the boolean enumeration (`Bool`), 1-, 2-, 4- and 8-byte +fixed-point at their native widths and signs (`Int8`, `Uint8`, `Int16`, +`Uint16`, `Int32`, `Uint32`, `Int64`), 4-byte and 8-byte floating-point +(`Float32`, `Float`), and fixed-length strings (`SaveHDF5Text`); `Float16` +and `Complex` are refused, because `LoadHDF5` decodes neither. Written files +stay byte-deterministic, and the files the previous releases wrote keep +their exact bytes. A filtered +dataset in a latest-version file is refused, as are the deflate and shuffle +filters for string datasets, which are written contiguously. Bounds: a rank of +at most 32, an object path of at most 4096 bytes, at most 4096 attributes on +one object. A group path in `groupAttrs` that names no group of the file is +refused. + +**FITS.** Only the primary HDU is read as an image: a file whose first card is +`XTENSION` is refused, and so is any `BITPIX` other than -64 and -32. The +table loader instead walks the HDUs and returns the first `BINTABLE` or +`TABLE` extension it finds, refusing any other extension kind by name rather +than guessing its size. A binary table column of another repeat count than one +(a `3E` vector) and every variable-length form (`P`, `Q`) are refused by name. +On the image side the header must open with `SIMPLE` and `SIMPLE = F` is +refused, `COMMENT`, `HISTORY` and blank cards carry no value and are skipped, +and a header without an `END` card is an error. + +**NetCDF.** The classic model only: CDF-1 and CDF-2 are read (CDF-2 addresses +its offsets with 64-bit words, its record count stays 32-bit), CDF-5 is +refused by version. The writer emits CDF-1 and refuses a file at or past +2 GiB, which a CDF-1 offset cannot address. Types read: `NC_BYTE` lands +`Int8`, `NC_SHORT` lands `Int16`, `NC_INT` lands `Int32` and `NC_CHAR` +lands `Uint8` (a CHAR variable carries raw bytes at the array level, not +text), while `NC_FLOAT` and `NC_DOUBLE` keep their float64 landing; a type +code beyond the classic set is refused by name. Types written: float64 as +`NC_DOUBLE`, float32 as `NC_FLOAT` and int64 as `NC_INT`, the widest +integer of the classic model: a value outside the int32 range refuses by +name rather than truncating, and a variable that landed a narrow dtype +from a file is refused by the writer until it is converted with `Astype`. +Attributes are +written as `NC_CHAR`; on read, a numeric attribute is rendered as decimal +text. + +A record dimension (length zero) may be declared first, at most one per file. +A variable whose first dimension is the record dimension is a record variable: +it must hold a whole number of records, every record variable must agree on +that number, and the file interleaves one slab of each per record, padded to +four bytes, the layout the classic model defines. On read, the record +dimension comes back with length zero and each record variable's first axis +carries the record count the file declares. A rank-0 variable comes back as a +one-element vector with no dimension names, since an array always carries at +least one dimension. + +#### Errors + +- A shape or dtype a writer cannot store (a complex array to CSV or HDF5, an + array that is not 2-D, an int64 value outside the int32 range NetCDF stores): + an error naming the shape, the dtype or the value. +- A refused format feature (an unsupported datatype class, storage layout, + filter, extension kind or NetCDF type): an error naming the feature, never a + partly read result. +- Malformed input (a truncated header, a card without an `END`, a row with the + wrong field count, a chunk outside the file): an error naming the byte + offset, the row or the field. +- A hostile declaration (a dimension or dataset extent larger than the file + can hold, a cycle in the group tree): an error from the bounds check that + runs before the allocation, not a panic and not a truncated read. +- A missing file or an unreadable one: the operating system's error, wrapped + with the name of the call. +- Every error carries the library's `tensor: ` prefix and the name of the + function that produced it. + +#### Workflow + +```go +package main + +import ( + "fmt" + "os" + "path/filepath" + + tensor "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/io" +) + +func main() { + dir, err := os.MkdirTemp("", "tensor-io") + if err != nil { + panic(err) + } + defer os.RemoveAll(dir) + + values, err := tensor.FromFloats([]float64{1.5, 2.25, 3.125, 4}, 2, 2) + if err != nil { + panic(err) + } + + // Write: the path builds the group tree, and the attributes of the root and + // of the group are keyed by path, the root as "/". + path := filepath.Join(dir, "scan.h5") + err = io.SaveHDF5(path, []io.HDF5Dataset{{ + Path: "/scan/temperature", + Values: values, + Attrs: map[string]string{"units": "degC"}, + }}, map[string]map[string]string{ + "/": {"title": "cruise"}, + "/scan": {"instrument": "thermistor"}, + }) + if err != nil { + panic(err) + } + + // Read: the same paths, shapes, dtypes and values, with the attributes of + // the enclosing groups merged into each dataset. + sets, err := io.LoadHDF5(path) + if err != nil { + panic(err) + } + for _, d := range sets { + fmt.Println(d.Path, d.Shape, d.Values.Dtype(), d.Attrs["units"], d.Attrs["title"]) + } +} +``` + +It prints `/scan/temperature [2 2] float degC cruise`: the dataset's own +`units` attribute beside the `title` it inherits from the root group. + +### Package `grad` + +Reverse-mode automatic differentiation over the shared array surface. A +computation written as ordinary Go calls over `Tensor` values records a +graph, and one call to `Backward` propagates the gradient from the +output back to every leaf that requires it: + +```mermaid +sequenceDiagram + participant Caller + participant Ops as the forward ops + participant Graph as the recorded graph + Caller->>Ops: build the pass from leaves + Ops-->>Graph: each call records its node and inputs + Caller->>Graph: loss.Backward() + Graph-->>Graph: reverse sweep, each node applies its adjoint + Graph-->>Caller: every leaf reads its Grad +``` + +The package deliberately stops at that graph. There is no forward-mode +differentiation, nothing beyond the second derivative, no graph +serialisation and no parameter registry: the caller owns the leaves, the +graph is a transient record of one forward pass, and the operation set is +closed, so a new primitive is a new method here and never a +user-registered op. + +#### Building a graph + +| Call | What it does | +|---|---| +| `FromFloat64s(vals, requiresGrad, shape...)` | A float64 leaf built from a copy of `vals`; the shape needs at least one dimension and must hold exactly `len(vals)` elements, and anything else is an error. | +| `FromArray(a, requiresGrad)` | Wraps an existing array as a tensor without copying it; `requiresGrad` marks it as a leaf whose gradient `Backward` fills. | +| `Tensor.Data()` | The underlying array. | +| `Tensor.Grad()` | The accumulated gradient, `nil` before the first `Backward`. | +| `Tensor.RequiresGrad()` | Whether the tensor is a trainable leaf. | +| `Tensor.ZeroGrad()` | Discards the accumulated gradient. | +| `Tensor.SetGrad(g)` | Replaces the gradient array outright: the hook for clipping, accumulation resets and test harnesses. | +| `Tensor.ReplaceWith(a)` | Swaps the underlying data array, the one mutable point the optimisers use for parameter updates. A graph built before the swap keeps differentiating against the operands its closures captured. | + +Only `Tensor` values that require grad carry a node; the result of an +operation requires grad when any of its inputs does, so `RequiresGrad` +is true for intermediate results as well as for leaves. `FromFloat64s` +and `FromArray` are the only constructors. + +#### Running the backward pass + +| Call | What it does | +|---|---| +| `Tensor.Backward()` | Seeds the output with ones, sweeps the recorded graph in reverse and accumulates the gradient into every leaf that requires it. It refuses an output that does not require grad, and a complex output, whose seed would not be a real scalar objective. | + +To differentiate the graph without touching the leaves, as the +second-order helpers do, the internal reverse pass computes into a local +map and commits nothing; the public `Backward` is the call that writes +the leaves' `Grad`. + +#### Arithmetic + +| Call | What it does | +|---|---| +| `Tensor.Add(u)` | `t + u`. | +| `Tensor.Sub(u)` | `t - u`. | +| `Tensor.Mul(u)` | Element-wise product `t·u`; the complex adjoint conjugates the other factor. | +| `Tensor.Div(u)` | Element-wise true division `t/u`. | +| `Tensor.Neg()` | `-t`. | +| `Tensor.Scale(f)` | Every element times the constant `f`; the backward multiplies by the same factor. | +| `Tensor.Pow(n)` | `xⁿ` for an integer `n ≥ 0`; a negative exponent is an error. The exponent-0 backward is zero everywhere, including at `x = 0`, where `n·xⁿ⁻¹` would be a NaN. | +| `Tensor.Abs()` | The absolute value of each element. The real subgradient is `sign(x)` with 0 at zero; the complex one is `z/(2r)` for the magnitude `r`, also zero at the origin. | +| `Tensor.Abs2()` | The squared magnitude as a real tensor. The complex backward is `dz = g·z`; a real operand is squared in float64 and rounded once, keeping its own width. | +| `Tensor.Sqrt()` | `√x`; the backward is `g/(2√x)` with the denominator floored at `1e-12`. A negative input is not refused: the NaN the forward produces flows into the gradient. | +| `Tensor.Exp()` | `eˣ`, complex included, where the holomorphic adjoint multiplies by the conjugated output. | +| `Tensor.Log()` | The natural logarithm. The domain is the caller's, so a non-positive input reaches the gradient without a report; shift the input inside its domain first. | +| `Tensor.Sigmoid()` | The logistic sigmoid; the backward is `g·σ·(1−σ)`. | +| `Tensor.Tanh()` | The hyperbolic tangent; the backward is `g·(1−tanh²)`. | +| `Tensor.Clip(lo, hi)` | Clamps into `[lo, hi]`; the gradient passes through only where `lo ≤ x ≤ hi`. | +| `Tensor.Floor()` | Rounds down with an all-zero gradient, its derivative being zero almost everywhere. | + +The dtype gate is uniform: every method accepts float, float32 and +complex input except where its own row says otherwise, and the exceptions +are exactly these. `Log`, `Sigmoid`, `Tanh`, `Clip`, `Floor`, `Sqrt`, +`MeanAxis`, `L2NormAxis` and `MatMulBatched` are float and float32 only +and refuse a complex tensor with an error naming the dtype they got. +`Imag` and `IRFFT` need a complex tensor, `RFFT` needs a real one, and +both of those name the dtype they got too. + +#### Matrix products + +| Call | What it does | +|---|---| +| `Tensor.MatMul(u)` | The differentiable matrix product over the shapes `MatMul2D` accepts: 2-D×2-D, 2-D×1-D and 1-D×2-D. The complex adjoint conjugates and transposes, `dA = g·Bᴴ` and `dB = Aᴴ·g`. | +| `Tensor.MatMulBatched(u)` | Stacked 3-D matrices multiplied batch-wise, `(N, M, K)·(N, K, P) → (N, M, P)`; a rank other than 3, a batch mismatch or an inner-dimension mismatch is an error. Float and float32 only. | + +#### Reductions + +| Call | What it does | +|---|---| +| `Tensor.Sum()` | All elements into a single-element tensor; a complex tensor sums in complex. | +| `Tensor.Mean()` | The mean into a single-element tensor, keeping the operand's dtype; an empty tensor is an error. | +| `Tensor.SumAxis(dim)` | Sums along `dim` and drops it; the backward broadcasts the incoming gradient back over the reduced dimension. | +| `Tensor.MeanAxis(dim)` | Averages along `dim` and drops it; the backward scales by `1/size(dim)` before the broadcast. Float and float32 only. | +| `Tensor.L2NormAxis(dim)` | The L2 norm along `dim`, dropped from the shape. The result is float64 whatever the operand's dtype, and the backward divides each element by its line norm, floored at `1e-12` so a zero line does not divide by zero. | + +#### Shape, layout and joining + +| Call | What it does | +|---|---| +| `Tensor.Reshape(shape...)` | A view-equivalent reshape; the backward reshapes the incoming gradient back. | +| `Tensor.Squeeze(dim)` | Removes a size-1 dimension; the gradient unsqueezes back. | +| `Tensor.Unsqueeze(dim)` | Inserts a size-1 dimension; the gradient squeezes back. | +| `Tensor.Transpose()` | Reverses the dimensions; the gradient transposes back. | +| `Tensor.TransposeAxes(dims...)` | Reorders the axes by a permutation; the backward applies the inverse permutation. | +| `Tensor.Slice(dim, start, stop)` | A range along one dimension; the backward writes the incoming gradient into the corresponding region of the original shape. | +| `Tensor.Concat(u, dim)` | Joins two tensors along an existing dimension, every other dimension agreeing; the backward routes each side its own span, narrowed to the side's dtype first. | +| `Tensor.BroadcastTo(shape...)` | Expands under broadcasting rules, size-1 dimensions replicating; the backward sums the gradient over every replicated dimension. The shape is validated first, so an impossible target reports the shape mismatch; an int tensor is then refused, because only float, float32 and complex tensors carry a graph. | + +#### Complex graphs + +Complex tensors differentiate under the Wirtinger convention, the +convention the optimiser ecosystem settled on: the loss must be real, +and 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))`, so every complex backward conjugates exactly where +the calculus puts it, and `Mul`, `Div`, `MatMul`, `Pow`, `Exp` and `Abs2` +follow that rule. + +`Backward` enforces the real-loss rule rather than approximating it: a +complex output is rejected with + +```text +tensor: Backward: the loss must be real-valued; reduce the complex result with Real, Imag, Abs or Abs2 first +``` + +| Call | What it does | +|---|---| +| `Tensor.Conj()` | The element-wise conjugate. It is anti-holomorphic, so its adjoint conjugates the incoming gradient, `dz = conj(g)`; on a real tensor it is a copy. | +| `Tensor.Real()` | The real part of each element as a float tensor; the complex backward halves the gradient, `∂Re z/∂z̄ = ½`. | +| `Tensor.Imag()` | The imaginary part of each element as a float tensor; the complex backward scales by `i/2`. A non-complex input is an error. | +| `Tensor.Abs()` | Of a complex tensor, the magnitude as a float tensor, with `dz = g·z/(2r)` for the magnitude `r`. | + +A real tensor inside a complex graph narrows the incoming complex +gradient by `2·Re`: for a real variable, `dL/dx = 2·Re(∂L/∂x̄)`. That +factor also cancels the `½` the `Real` and `Imag` backward paths +contribute, so a graph mixing the two dtypes composes exactly. + +#### Spectral transforms + +| Call | What it does | +|---|---| +| `Tensor.FFT()` | The forward DFT of a rank-1 tensor; the backward multiplies the incoming gradient by `Fᴴ`, which is the inverse transform scaled by `n`. | +| `Tensor.IFFT()` | The inverse transform of a rank-1 tensor; the backward runs the forward transform scaled by `1/n`. | +| `Tensor.FFT2()` | The 2-D forward transform; the backward is the 2-D inverse scaled by `H·W`. | +| `Tensor.IFFT2()` | The 2-D inverse transform; the backward is the 2-D forward scaled by `1/(H·W)`. | +| `Tensor.RFFT()` | The real-input half-spectrum transform of a rank-1 real tensor; the backward folds the half-spectrum gradient into `dx = 2·Re(F_halfᴴ·g)` with one padded forward FFT. A complex input is an error. | +| `Tensor.IRFFT(n)` | The inverse half-spectrum transform, `n/2+1` complex bins into a real signal of length `n`; the backward widens the real gradient to a full forward FFT and halves the self-mirrored DC bin and, for even `n`, the Nyquist bin. The input must be a rank-1 complex tensor. | + +Every rank check is an error, so `FFT2` on a vector and `FFT` on a matrix +are refused rather than silently reinterpreted. + +#### Second order + +| Call | What it does | +|---|---| +| `Hessian(f, x, opts)` | The dense Hessian of the scalar `f` at `x`, an `(n, n)` float64 array, at the cost of `2n` gradient evaluations. Column `j` is the central difference of the analytic gradient along coordinate `j`. | +| `HessianVectorProduct(f, x, v, opts)` | `H·v`, by a central difference along the direction itself with the step normalised so it never depends on `v`'s magnitude. Two gradient evaluations answer for any `n`. A zero `v` returns zeros shaped like `x`; a `v` of a different element count is an error. | + +Both call `f` with a tensor that requires grad and expect a single-element +real tensor. A complex point (and for the product, a complex direction) is +an error, as are a non-scalar result and an objective that does not depend +on `x` at all. Both differentiate without committing, so the accumulated +gradients of the tensors `f` closes over are left exactly as they were, on +the success and the error path alike. + +#### Newton-CG minimisation + +| Call | What it does | +|---|---| +| `MinimiseNewtonCG(f, x0, opts)` | A local minimum of the scalar `f` near `x0`, returned as the point, the objective's value there and an error. Each outer step solves `H·s = −∇f` with truncated conjugate gradients, then an Armijo backtracking line search secures descent. No dense Hessian is ever formed: each CG iteration buys one Hessian-vector product, two reverse passes, so the memory stays that of the point. The converged point is a fresh array the caller owns. | + +#### Hamiltonian Monte Carlo + +| Call | What it does | +|---|---| +| `SampleHMC(logDensity, q0, opts)` | Draws `opts.Samples` states from the unnormalised density whose logarithm `logDensity` computes, starting at the vector `q0`, and returns them as a `(Samples × dim)` array. Each trajectory integrates the Hamiltonian dynamics of a unit-mass particle with the leapfrog scheme, one gradient evaluation per step, and a Metropolis test on the energy change. A rejected trajectory repeats the current state, as Markov chain sampling does. | + +`logDensity` receives a leaf tensor that requires grad and must return a +scalar tensor connected to it. An error it raises at `q0` is fatal; an +error raised inside a proposal marks the state as outside the support and +rejects the trajectory, which is how a constrained density keeps the +chain out of forbidden regions. A nil density, a state that is not a +non-empty real vector, a non-positive `Step`, `Steps` or `Samples`, a +negative `BurnIn`, and a density that does not yield a gradient are +errors. + +#### Adjoint ODE sensitivities + +| Call | What it does | +|---|---| +| `AdjointODE(f, params, t0, t1, y0, lossGrad, opts)` | Differentiates the solution of `y' = f(t, y)` at `t1` with respect to the initial state and to the parameters `f` closes over. Returns `∂L/∂y0` and, parallel to `params`, each `∂L/∂θk` shaped like its parameter. `lossGrad` is `∂L/∂y(t1)`, the seed the loss itself contributes. | + +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 then integrated by the same +adaptive solver in reverse, so the cost is one extra solve whatever the +parameter count. Backward-in-time problems (`t1 < t0`) work, and a +zero-length span with `t0 == t1` short-circuits to `lossGrad` for the +state and zeros for the parameters. `opts` is the `integrate` package's +`ODEOptions` (`RelTol` 1e-6, `AbsTol` 1e-9, `MaxSteps` 100000 where the +field is not positive). + +#### Options and results + +**`Tensor`**: a differentiable n-dimensional array. It carries no +exported fields; its state is reached through `Data`, `Grad`, +`RequiresGrad`, `ReplaceWith`, `SetGrad` and `ZeroGrad`. + +**`HessianOptions`**: the stencil width the second-order helpers +perturb by. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Step` | `float64` | `0`; `Hessian` then takes `sqrt(2.22e-16)·max(1, abs(x_j))` per coordinate, about `1.49e-8·max(1, abs(x_j))`, and `HessianVectorProduct` takes an absolute `1e-5` along the normalised direction | The perturbation size. A positive value is used as given and applies in both directions, so the stencil stays central. | + +**`HMCOptions`**: the sampler's trajectory and chain settings. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `Step` | `float64` | none; must be positive, `0` is an error | The leapfrog step size. | +| `Steps` | `int` | none; must be at least 1, `0` is an error | Leapfrog steps per trajectory, so the trajectory length is `Step·Steps`. | +| `BurnIn` | `int` | `0` | Trajectories discarded before the first sample is kept; negative is an error. | +| `Thin` | `int` | `0` | Every `Thin`-th trajectory after the burn-in contributes one sample; zero or less normalises to `1`. | +| `Samples` | `int` | none; must be at least 1, `0` is an error | Kept states; the result is a `(Samples × dim)` array. | +| `Seed` | `int64` | `0` | Seeds the package's xoshiro generator, so a run is bit-reproducible. | + +**`NewtonCGOptions`**: the outer iteration budget, the stopping +tolerance and the inner solve budget. + +| Field | Type | Default | Effect | +|---|---|---|---| +| `MaxIterations` | `int` | `100`; zero or less takes the default | The outer Newton step budget. Exhausting it without reaching the tolerance is an error. | +| `Tolerance` | `float64` | `1e-8`; zero or less takes the default | Stops when the Euclidean norm of the gradient falls to it or below. | +| `MaxCGIterations` | `int` | `n`, the problem dimension; zero or less takes the default | The inner conjugate-gradient iterations per outer step. | + +The line search halves the step up to 40 times from 1.0 and stops at the +first point satisfying the Armijo condition with constant `1e-4`; no +point satisfying it is an error. `MaxCGIterations` defaults to the +dimension because the conjugate-gradient method terminates in at most `n` +steps in exact arithmetic, and the inner solve stops early when the +residual falls to a tenth of the current gradient norm or when it meets +non-positive curvature, in which case the first iteration falls back to +the steepest descent direction. + +#### Errors + +- an int tensor in the graph: every differentiable call returns + `tensor: autograd: needs a float, float32 or complex tensor, got int`; the float-only calls name float and float32 instead; +- `Imag` on a real tensor, `RFFT` on a complex tensor, `IRFFT` on a real one, and `MatMulBatched` on anything but float or float32: an error naming the dtype; +- `Backward` on an output that does not require grad, or on a complex output: an error naming the condition, the complex one naming `Real`, `Imag`, `Abs` and `Abs2` as the reducers; +- `Pow` with a negative exponent: `tensor: Pow: negative exponent has no general real gradient`, the exponent being an integer argument rather than a graph node; +- a shape or range the array core refuses: `Reshape`, `Slice`, `Concat`, `Squeeze`, `Unsqueeze`, `TransposeAxes`, `SumAxis`, `MeanAxis` and `L2NormAxis` return the core's own error, and `Mean` refuses a tensor with no elements; +- `Hessian` and `HessianVectorProduct`: a complex point, a complex direction, a `v` of the wrong length, a non-scalar `f`, and an `f` that does not depend on `x`; +- `MinimiseNewtonCG`: a nil `f`, a nil or empty starting point, a complex starting point, a non-scalar or non-finite objective, a CG direction that does not descend, a line search that finds no descent in 40 halvings, and an exhausted iteration budget (the error names the gradient norm it stopped at); +- `SampleHMC`: a nil density, a nil state, a state that is not a non-empty vector, a complex state, a non-positive `Step`, `Steps` or `Samples`, a negative `BurnIn`, a density returning a non-scalar, and a density that yields no gradient; +- `AdjointODE`: a nil `f`, a state or loss seed that is not a non-empty vector, a complex state, a seed of the wrong length, a parameter that does not require grad or has no elements, an `f` returning the wrong shape, a parameter disconnected from the dynamics, and a trajectory too short to interpolate (fewer than two recorded nodes). + +Every message carries the `tensor: ` prefix the library uses throughout. + +#### Workflow + +The flagship sequence: build the leaves, run the forward pass, call +`Backward` once, read the gradients. The imports are +`tensor "sourcedock.dev/petrbalvin/tensor"` for the array type and +`sourcedock.dev/petrbalvin/tensor/grad` for the graph. + +```go +// squaredErrorGradient builds Σ (w·x)² and returns ∂L/∂w for one +// forward pass and one reverse sweep. +func squaredErrorGradient(xs, ws []float64) (*tensor.Array, error) { + // Leaves: the data does not require grad, the parameter does. + x, err := grad.FromFloat64s(xs, false, len(xs)) + if err != nil { + return nil, err + } + w, err := grad.FromFloat64s(ws, true, len(ws)) + if err != nil { + return nil, err + } + + // Forward pass: every call records a node and its inputs. + prod, err := w.Mul(x) + if err != nil { + return nil, err + } + sq, err := prod.Pow(2) + if err != nil { + return nil, err + } + loss, err := sq.Sum() + if err != nil { + return nil, err + } + + // Reverse sweep, once, then the gradient is on the leaf. + if err := loss.Backward(); err != nil { + return nil, err + } + return w.Grad(), nil +} +``` + +The runnable form of this and the other flagship workflows, with the +numbers they produce, is in `grad/example_test.go` and rendered beside +the package on pkg.go.dev. + +### Package `plot` + +Deterministic SVG line charts for scientific figures: linear axes with +five ticks each, one legend line per series, titles and axis labels, +and nothing else. The output is deterministic by contract, so a figure +in a paper is regenerated and compared exactly like any other computed +number. The series colours are part of that contract: a fixed cycle of +seven samples of the Viridis perceptual-uniform map (Smith, van der +Walt and Firing, CC0), taken over the map's legible-on-white range, so +the same series always wears the same colour. The package is small by +intent; it draws the figures, it does not stage a cinema. + +| Member | What it does | +|---|---| +| `Chart` | one figure: title, axis labels, optional size, optional axis ranges, the series list | +| `Series`, `Point` | one named polyline and one data point in axis units | +| `Line(name, xs, ys)` | a `Series` from two rank-1 arrays of equal, non-zero length; any numeric dtype through the promotion ladder, non-finite values refused | +| `(*Chart).WriteSVG(path)` | renders the chart; the same chart always renders byte for byte the same file. Refuses a non-finite point in any series, the contract `Line` enforces on the caller's behalf; an axis range with a zero, inverted or non-finite span falls back to the data's own bounds | + +```go +xs, _ := tensor.FromFloats(linspace, 200) +ys, _ := tensor.FromFloats(spectrum, 200) +series, _ := tensor.Line("spectrum", xs, ys) +chart := tensor.Chart{ + Title: "Absorption spectrum", + XLabel: "wavelength [nm]", YLabel: "intensity", + Series: []tensor.Series{series}, +} +err := chart.WriteSVG("spectrum.svg") +``` + +The ranges fall back to the data extent when the zero value leaves them +unset, a chart needs at least two points across its series, and every +error is prefixed `tensor: ` like the rest of the library. + +### Package `spmd` + +An experiment: a new surface carrying a concept that is not fully +verified and remains the subject of research, expected to move as it +settles. Explicit SPMD worlds: one program runs on many ranks, over TCP +between machines (`Listen` at rank 0, `Join` everywhere else) or in +one process over channels (`Launch`), with the same collectives and +the same answers on both transports. The determinism contract is the +package's point: the order of a reduction is a function of the data, +never of the world size, the machine, the worker count or the order +frames arrive in, so a sharded reduction and the single-array +reduction carry the same bits, for one rank or for fifty. Any failure +or deadline fails the whole world, and no collective ever returns a +partial numeric result. The package is imported explicitly; the root +facade does not re-export it. + +| Member | What it does | +|---|---| +| `Launch(size, fn)` | runs the same function on `size` ranks of one process, one goroutine per rank; the failing ranks' errors come back joined in rank order | +| `Listen(addr, size, opts)` / `Join(addr, opts)` | builds the networked world: rank 0 listens and assigns ranks in dial order, the other ranks dial; the sharded reductions' results never depend on the assignment, while the same-shaped arrays' fold order is the rank order the program fixes | +| `Options` | `Timeout`, one collective's wait on the network (ten minutes by default), and `MaxMessage`, the frame ceiling a peer cannot talk the world into allocating past | +| `(*World).Rank` / `Size` / `Barrier` | this rank's place, the world's size, and a barrier | +| `Partition(globalN, size, rank)` | the rank's contiguous piece of a global axis, cut on the canonical fold partition's block boundaries; a `Span` names the global length and the piece's bounds; a negative `globalN` or a rank outside `[0, size)` is refused with an error naming the input, and valid arguments never fail | +| `(*World).Broadcast(a, root)` | delivers root's array to every rank, bits included | +| `(*World).Scatter(global, root)` | deals root's array out along the first dimension on the canonical boundaries; returns the rank's piece and its `Span` | +| `(*World).Gather(local, root)` / `AllGather(local)` | raises the pieces back into the whole on the root, or on every rank, joined in rank order | +| `(*World).ReduceShards(local, span, op, root)` / `AllReduceShards(local, span, op)` | reduces the shards of one global array: the answer is the single-array reduction's exact bits at any world size; `Sum`, `Min`, `Max`, `Prod`, and `Any`/`All` on `Bool` | +| `(*World).ReduceNormShards(local, span, p, root)` / `AllReduceNormShards(local, span, p)` | the Lp norm of the shards, the power sums folded through the canonical blocks and closed by the norm's own Sqrt and Pow; finite positive p | +| `(*World).ReduceDotShards(x, y, span, root)` / `AllReduceDotShards(x, y, span)` | the dot product of two equally sharded arrays, each rank folding its own blocks; the two shards carry the same dtype | +| `(*World).ExchangeHalos(local, halos)` | hands the edge rows of the rank's piece to its two neighbours and receives theirs, the ghost cells a domain-decomposed stencil needs; world edges answer nil, empty pieces join with empty edges | +| `(*World).ExchangeHalosOnGrid(local, halos, axis, grid)` | the same neighbour exchange on a process grid: `grid` lays the world's ranks out row-major and must cover the world exactly, the halos travel along `axis` between grid neighbours, and the edge slabs keep the piece's other dimensions | +| `(*World).Reduce(a, op, root)` / `AllReduce(a, op)` | folds the ranks' same-shaped arrays together elementwise in rank index order, the canonical order the program itself fixes; the fold runs chunk by chunk across the world instead of piling on one rank | +| `(*World).ReduceArgShards(local, span, op, root)` / `AllReduceArgShards(local, span, op)` | the global index of the extremum (`Min` or `Max`), ties by the earliest global index, NaNs skipped, the single-array ArgMax and ArgMin answer | +| `(*World).ReduceArgSortShards(local, span, root)` / `AllReduceArgSortShards(local, span)` | the global permutation that sorts the whole array, the single-array ArgSort's own answer with its tie and NaN placement | + +```go +err := spmd.Launch(4, func(w *spmd.World) error { + span, err := spmd.Partition(globalN, 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 + } + _ = got // the single-array Sum's exact bits, on every rank + return w.Barrier() +}) +``` + +The shards' piece must be the canonical partition's own cut, which +`Scatter` and `Partition` provide; any other cut is refused by name, +because the bit-identity contract lives on those boundaries. +`Partition` itself answers the `Span` or an error: valid arguments +never fail. The wire +form is a fixed little-endian header naming dtype, shape, sender and +destination, and raw element bits: NaN payloads and signed zeros +travel intact. + +## Notes + +All packages share the conventions [ARCHITECTURE.md](ARCHITECTURE.md) +documents: immutable arrays, loud errors prefixed `tensor: `, +deterministic parallel execution, and no third-party dependencies. + +Every array is the same type everywhere: `tensor.Array` is an alias +for the core array, so an array built by a root constructor is +accepted by a domain package without conversion, and a caller who +imports a domain package alone can still build inputs through +`linalg.ArrayFromFloatsSafe`. Importing the root package is what most +code does, and it costs nothing else. + +The exported surface is documented in godoc form in the source, and +`go doc` is the authority on signatures and types. The `-all` form is +the complete inventory of a package: every exported function, type, +method and option field with its signature, so the full list of what +the library offers is one command away and never a second copy that +can age: + +```sh +go doc -all sourcedock.dev/petrbalvin/tensor +go doc -all sourcedock.dev/petrbalvin/tensor/linalg +``` + diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md new file mode 100644 index 0000000..8f61222 --- /dev/null +++ b/docs/ARCHITECTURE.md @@ -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
FromFloats, Zeros, Grid"] + A2["io loaders
CSV, FITS, HDF5, NetCDF, mmap"] + C1["element-wise ops"] + C2["reductions"] + C3["domain kernels
linalg, signal, integrate,
stats, optim"] + C4["autograd graph
grad.Backward"] + S1["spmd collectives
Broadcast, Scatter, Gather,
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. diff --git a/docs/BENCHMARKING.md b/docs/BENCHMARKING.md new file mode 100644 index 0000000..39b3abd --- /dev/null +++ b/docs/BENCHMARKING.md @@ -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. diff --git a/docs/DEVELOPMENT.md b/docs/DEVELOPMENT.md new file mode 100644 index 0000000..a5c7d8c --- /dev/null +++ b/docs/DEVELOPMENT.md @@ -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. diff --git a/dtypes_facade_test.go b/dtypes_facade_test.go new file mode 100644 index 0000000..5f505aa --- /dev/null +++ b/dtypes_facade_test.go @@ -0,0 +1,1354 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package tensor + +import ( + "math" + "strings" + "testing" +) + +// The facade matrix pin for the narrow element types: every new dtype +// through the re-exported names only, against references computed in +// plain Go from the accessor-widened values. Every comparison is +// bit-exact. + +// dtNew lists the seven narrow dtypes; dtIntAnchor adds Int, the dtype +// whose existing treatment each domain package's census matches. +var dtNew = []Dtype{Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32} + +var dtNewWithInt = []Dtype{Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int} + +// dtStrings pins the rendering of all twelve dtypes. +var dtStrings = map[Dtype]string{ + Bool: "bool", Int8: "int8", Uint8: "uint8", Int16: "int16", + Uint16: "uint16", Int32: "int32", Uint32: "uint32", Int: "int", + Float16: "float16", Float32: "float32", Float: "float", Complex: "complex", +} + +// dtCast narrows a widened value into dt's value space the way the +// payload store does: Go conversion semantics, wrap-around included. +func dtCast(dt Dtype, v float64) float64 { + switch dt { + case Bool: + if v != 0 { + return 1 + } + return 0 + case Int8: + return float64(int8(int64(v))) + case Uint8: + return float64(uint8(int64(v))) + case Int16: + return float64(int16(int64(v))) + case Uint16: + return float64(uint16(int64(v))) + case Int32: + return float64(int32(int64(v))) + case Uint32: + return float64(uint32(int64(v))) + case Int: + return float64(int64(v)) + case Float32: + return float64(float32(v)) + case Float16: + return HalfToFloat64(HalfFromFloat64(v)) + default: + return v + } +} + +// dtMake builds an array of the given dtype through the facade +// constructors, casting each value into the dtype's space first. +func dtMake(t *testing.T, dt Dtype, vals []float64, shape ...int) *Array { + t.Helper() + switch dt { + case Bool: + bs := make([]bool, len(vals)) + for i, v := range vals { + bs[i] = v != 0 + } + a, err := FromBools(bs, shape...) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + return a + case Int8: + vs := make([]int8, len(vals)) + for i, v := range vals { + vs[i] = int8(int64(v)) + } + a, err := FromInt8s(vs, shape...) + if err != nil { + t.Fatalf("FromInt8s: %v", err) + } + return a + case Uint8: + vs := make([]uint8, len(vals)) + for i, v := range vals { + vs[i] = uint8(int64(v)) + } + a, err := FromUint8s(vs, shape...) + if err != nil { + t.Fatalf("FromUint8s: %v", err) + } + return a + case Int16: + vs := make([]int16, len(vals)) + for i, v := range vals { + vs[i] = int16(int64(v)) + } + a, err := FromInt16s(vs, shape...) + if err != nil { + t.Fatalf("FromInt16s: %v", err) + } + return a + case Uint16: + vs := make([]uint16, len(vals)) + for i, v := range vals { + vs[i] = uint16(int64(v)) + } + a, err := FromUint16s(vs, shape...) + if err != nil { + t.Fatalf("FromUint16s: %v", err) + } + return a + case Int32: + vs := make([]int32, len(vals)) + for i, v := range vals { + vs[i] = int32(int64(v)) + } + a, err := FromInt32s(vs, shape...) + if err != nil { + t.Fatalf("FromInt32s: %v", err) + } + return a + case Uint32: + vs := make([]uint32, len(vals)) + for i, v := range vals { + vs[i] = uint32(int64(v)) + } + a, err := FromUint32s(vs, shape...) + if err != nil { + t.Fatalf("FromUint32s: %v", err) + } + return a + case Int: + vs := make([]int64, len(vals)) + for i, v := range vals { + vs[i] = int64(v) + } + a, err := FromInts(vs, shape...) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + return a + case Float32: + vs := make([]float32, len(vals)) + for i, v := range vals { + vs[i] = float32(v) + } + a, err := FromFloat32s(vs, shape...) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + return a + case Float16: + vs := make([]float64, len(vals)) + copy(vs, vals) + a, err := FromFloat16s(vs, shape...) + if err != nil { + t.Fatalf("FromFloat16s: %v", err) + } + return a + case Complex: + vs := make([]complex128, len(vals)) + for i, v := range vals { + vs[i] = complex(v, 0) + } + a, err := FromComplexes(vs, shape...) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + return a + default: + fs := make([]float64, len(vals)) + copy(fs, vals) + a, err := FromFloats(fs, shape...) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a + } +} + +// dtWiden reads an array's elements through the widening accessor, the +// reference view every comparison in this file uses. The facade reader +// takes shape-matching indices, so the walk flattens row-major through +// full coordinates whatever the rank. +func dtWiden(t *testing.T, a *Array) []float64 { + t.Helper() + sh := a.Shape() + out := make([]float64, a.Len()) + coords := make([]int, len(sh)) + for i := range out { + v, err := FloatAt(a, coords...) + if err != nil { + t.Fatalf("FloatAt(%v): %v", coords, err) + } + out[i] = v + for d := len(sh) - 1; d >= 0; d-- { + coords[d]++ + if coords[d] < sh[d] { + break + } + coords[d] = 0 + } + } + return out +} + +// dtWant narrows reference values into dt's space, the operation's +// result dtype decided by the caller. +func dtWant(dt Dtype, vals []float64) []float64 { + out := make([]float64, len(vals)) + for i, v := range vals { + out[i] = dtCast(dt, v) + } + return out +} + +// dtAssertValues pins an array's widened values bit-exactly. +func dtAssertValues(t *testing.T, label string, got *Array, want []float64) { + t.Helper() + if got == nil { + t.Fatalf("%s: got a nil array", label) + } + if got.Len() != len(want) { + t.Fatalf("%s: got %d elements, want %d", label, got.Len(), len(want)) + } + vals := dtWiden(t, got) + for i, w := range want { + if vals[i] != w { + t.Fatalf("%s: element %d = %v, want %v (all: %v vs %v)", label, i, vals[i], w, vals, want) + } + } +} + +// dtRes pairs an array result with its error, the shape dtOf accepts +// straight from a multi-value call. +type dtRes struct { + a *Array + err error +} + +// dtOf captures an (array, error) pair for dtRun. +func dtOf(a *Array, err error) dtRes { return dtRes{a: a, err: err} } + +// dtRun unwraps a captured pair, failing the test on an error or a nil +// array with no error. +func dtRun(t *testing.T, label string, r dtRes) *Array { + t.Helper() + if r.err != nil { + t.Fatalf("%s: %v", label, r.err) + } + if r.a == nil { + t.Fatalf("%s: nil array, no error", label) + } + return r.a +} + +// TestDtypesFacadeConstruction pins the construction round-trips: each +// new dtype through its facade constructor, read back through the Raw* +// accessors and the widening readers, with the copying contract and the +// String renderings. +func TestDtypesFacadeConstruction(t *testing.T) { + cases := []struct { + dt Dtype + vals []float64 + }{ + {Bool, []float64{1, 0, 1, 1, 0}}, + {Int8, []float64{-128, -1, 0, 1, 127}}, + {Uint8, []float64{0, 1, 127, 128, 255}}, + {Int16, []float64{-32768, -1, 0, 1, 32767}}, + {Uint16, []float64{0, 1, 255, 256, 65535}}, + {Int32, []float64{-2147483648, -1, 0, 1, 2147483647}}, + {Uint32, []float64{0, 1, 65535, 65536, 4294967295}}, + } + for _, c := range cases { + name := dtStrings[c.dt] + t.Run(name, func(t *testing.T) { + a := dtMake(t, c.dt, c.vals, len(c.vals)) + if a.Dtype() != c.dt { + t.Fatalf("Dtype() = %s, want %s", a.Dtype(), c.dt) + } + if got := c.dt.String(); got != name { + t.Fatalf("String() = %q, want %q", got, name) + } + // Raw payload read-back, bit-exact against the source + // values cast into the dtype. + switch c.dt { + case Bool: + r := a.RawBools() + if len(r) != len(c.vals) { + t.Fatalf("RawBools length %d, want %d", len(r), len(c.vals)) + } + for i, v := range c.vals { + if r[i] != (v != 0) { + t.Fatalf("RawBools[%d] = %v, want %v", i, r[i], v != 0) + } + } + case Int8: + r := a.RawInt8s() + for i, v := range c.vals { + if r[i] != int8(int64(v)) { + t.Fatalf("RawInt8s[%d] = %d, want %d", i, r[i], int8(int64(v))) + } + } + case Uint8: + r := a.RawUint8s() + for i, v := range c.vals { + if r[i] != uint8(int64(v)) { + t.Fatalf("RawUint8s[%d] = %d, want %d", i, r[i], uint8(int64(v))) + } + } + case Int16: + r := a.RawInt16s() + for i, v := range c.vals { + if r[i] != int16(int64(v)) { + t.Fatalf("RawInt16s[%d] = %d, want %d", i, r[i], int16(int64(v))) + } + } + case Uint16: + r := a.RawUint16s() + for i, v := range c.vals { + if r[i] != uint16(int64(v)) { + t.Fatalf("RawUint16s[%d] = %d, want %d", i, r[i], uint16(int64(v))) + } + } + case Int32: + r := a.RawInt32s() + for i, v := range c.vals { + if r[i] != int32(int64(v)) { + t.Fatalf("RawInt32s[%d] = %d, want %d", i, r[i], int32(int64(v))) + } + } + case Uint32: + r := a.RawUint32s() + for i, v := range c.vals { + if r[i] != uint32(int64(v)) { + t.Fatalf("RawUint32s[%d] = %d, want %d", i, r[i], uint32(int64(v))) + } + } + } + // The widening readers, package-level through the facade. + for i, v := range c.vals { + want := dtCast(c.dt, v) + fv, err := FloatAt(a, i) + if err != nil || fv != want { + t.Fatalf("FloatAt(%d) = %v, %v; want %v", i, fv, err, want) + } + bv, err := BoolAt(a, i) + if err != nil || bv != (want != 0) { + t.Fatalf("BoolAt(%d) = %v, %v; want %v", i, bv, err, want != 0) + } + iv, err := IntAt(a, i) + if err != nil || iv != int64(want) { + t.Fatalf("IntAt(%d) = %d, %v; want %d", i, iv, err, int64(want)) + } + } + // The copying contract: mutating the source after + // construction must not reach the array. + switch c.dt { + case Int8: + src := []int8{7, 7, 7, 7, 7} + b, err := FromInt8s(src, len(src)) + if err != nil { + t.Fatal(err) + } + src[0] = 9 + if got := b.RawInt8s()[0]; got != 7 { + t.Fatalf("FromInt8s aliased its source: element 0 = %d, want 7", got) + } + case Uint32: + src := []uint32{7, 7, 7, 7, 7} + b, err := FromUint32s(src, len(src)) + if err != nil { + t.Fatal(err) + } + src[0] = 9 + if got := b.RawUint32s()[0]; got != 7 { + t.Fatalf("FromUint32s aliased its source: element 0 = %d, want 7", got) + } + } + // The *FromArray ownership family hands the payload + // through: the values read back are the source's own. + switch c.dt { + case Bool: + src := []bool{true, false, true} + b, err := BoolsFromArray(src, 3) + if err != nil { + t.Fatal(err) + } + if r := b.RawBools(); r[0] != true || r[1] != false || r[2] != true { + t.Fatalf("BoolsFromArray read back %v", r) + } + case Int8: + src := []int8{1, -2, 3} + b, err := Int8sFromArray(src, 3) + if err != nil { + t.Fatal(err) + } + if r := b.RawInt8s(); r[0] != 1 || r[1] != -2 || r[2] != 3 { + t.Fatalf("Int8sFromArray read back %v", r) + } + case Uint8: + src := []uint8{1, 2, 255} + b, err := Uint8sFromArray(src, 3) + if err != nil { + t.Fatal(err) + } + if r := b.RawUint8s(); r[2] != 255 { + t.Fatalf("Uint8sFromArray read back %v", r) + } + case Int16: + src := []int16{-300, 0, 300} + b, err := Int16sFromArray(src, 3) + if err != nil { + t.Fatal(err) + } + if r := b.RawInt16s(); r[0] != -300 || r[2] != 300 { + t.Fatalf("Int16sFromArray read back %v", r) + } + case Uint16: + src := []uint16{0, 40000, 65535} + b, err := Uint16sFromArray(src, 3) + if err != nil { + t.Fatal(err) + } + if r := b.RawUint16s(); r[1] != 40000 { + t.Fatalf("Uint16sFromArray read back %v", r) + } + case Int32: + src := []int32{-70000, 0, 70000} + b, err := Int32sFromArray(src, 3) + if err != nil { + t.Fatal(err) + } + if r := b.RawInt32s(); r[2] != 70000 { + t.Fatalf("Int32sFromArray read back %v", r) + } + case Uint32: + src := []uint32{0, 70000, 4294967295} + b, err := Uint32sFromArray(src, 3) + if err != nil { + t.Fatal(err) + } + if r := b.RawUint32s(); r[2] != 4294967295 { + t.Fatalf("Uint32sFromArray read back %v", r) + } + } + // Array rendering: dtype name, shape in parentheses, + // values, the diagnostic format the core pins ("int (2, 2) + // [1, 2, 3, 4]"). + sample := dtMake(t, c.dt, c.vals[:3], 3) + want := name + " (3) [" + for i := range 3 { + if i > 0 { + want += ", " + } + v := dtCast(c.dt, c.vals[i]) + switch c.dt { + case Bool: + if v != 0 { + want += "true" + } else { + want += "false" + } + default: + want += itoa64(int64(v)) + } + } + want += "]" + if got := sample.String(); got != want { + t.Fatalf("String() = %q, want %q", got, want) + } + // Zeros and Ones per dtype. + z, err := Zeros(c.dt, 2, 3) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + if z.Dtype() != c.dt || z.Shape()[0] != 2 || z.Shape()[1] != 3 { + t.Fatalf("Zeros: dtype %s shape %v", z.Dtype(), z.Shape()) + } + dtAssertValues(t, "Zeros", z, make([]float64, 6)) + o, err := Ones(c.dt, 4) + if err != nil { + t.Fatalf("Ones: %v", err) + } + if o.Dtype() != c.dt { + t.Fatalf("Ones dtype = %s, want %s", o.Dtype(), c.dt) + } + dtAssertValues(t, "Ones", o, []float64{1, 1, 1, 1}) + }) + } +} + +// itoa64 renders an int64 without pulling strconv into the pin's +// expected strings by hand. +func itoa64(v int64) string { + if v == math.MinInt64 { + return "-9223372036854775808" + } + neg := v < 0 + if neg { + v = -v + } + var buf []byte + for { + buf = append([]byte{byte('0' + v%10)}, buf...) + v /= 10 + if v == 0 { + break + } + } + if neg { + return "-" + string(buf) + } + return string(buf) +} + +// TestDtypesFacadeOperations pins the elementwise, logic, comparison, +// reduction and indexing families on every new dtype (Int anchored) +// against plain-Go references over the widened values, bit-exact. +func TestDtypesFacadeOperations(t *testing.T) { + av := []float64{3, 1, 4, 1, 5} + bv := []float64{2, 7, 1, 8, 2} + condv := []float64{1, 0, 1, 0, 1} + for _, dt := range dtNewWithInt { + name := dtStrings[dt] + t.Run(name, func(t *testing.T) { + a := dtMake(t, dt, av, len(av)) + b := dtMake(t, dt, bv, len(bv)) + cond := dtMake(t, Bool, condv, len(condv)) + // Elementwise with same-dtype operands: the result keeps + // the dtype and wraps exactly as the payload store wraps. + for _, op := range []struct { + name string + run func(x, y *Array) (*Array, error) + f func(x, y float64) float64 + }{ + {"Add", Add, func(x, y float64) float64 { return x + y }}, + {"Sub", Sub, func(x, y float64) float64 { return x - y }}, + {"Mul", Mul, func(x, y float64) float64 { return x * y }}, + } { + if dt == Bool { + // Bool carries no arithmetic: the loud core + // refusal, pinned in the refusal census. + continue + } + got := dtRun(t, op.name, dtOf(op.run(a, b))) + if got.Dtype() != dt { + t.Fatalf("%s dtype = %s, want %s", op.name, got.Dtype(), dt) + } + want := make([]float64, len(av)) + for i := range av { + aw, bw := dtCast(dt, av[i]), dtCast(dt, bv[i]) + want[i] = dtCast(dt, op.f(aw, bw)) + } + dtAssertValues(t, op.name, got, want) + } + // Where with a Bool condition; the result dtype is + // promote(dt, dt) = dt. + got := dtRun(t, "Where", dtOf(Where(cond, a, b))) + if got.Dtype() != dt { + t.Fatalf("Where dtype = %s, want %s", got.Dtype(), dt) + } + wantWhere := make([]float64, len(av)) + for i := range av { + if condv[i] != 0 { + wantWhere[i] = dtCast(dt, av[i]) + } else { + wantWhere[i] = dtCast(dt, bv[i]) + } + } + dtAssertValues(t, "Where", got, wantWhere) + // Comparisons answer bool masks whatever the operand dtype. + for _, op := range []struct { + name string + run func(x, y *Array) (*Array, error) + f func(x, y float64) bool + }{ + {"Lt", Lt, func(x, y float64) bool { return x < y }}, + {"Gt", Gt, func(x, y float64) bool { return x > y }}, + {"Eq", Eq, func(x, y float64) bool { return x == y }}, + } { + m := dtRun(t, op.name, dtOf(op.run(a, b))) + if m.Dtype() != Bool { + t.Fatalf("%s mask dtype = %s, want Bool", op.name, m.Dtype()) + } + wantMask := make([]float64, len(av)) + for i := range av { + if op.f(dtCast(dt, av[i]), dtCast(dt, bv[i])) { + wantMask[i] = 1 + } + } + dtAssertValues(t, op.name, m, wantMask) + } + // Reductions widen exactly: Sum answers an Int scalar for + // the integer class (bool counts its trues there too) and + // a float scalar otherwise; Mean is always float64; Min, + // Max and ArgMax read the widened values. + sumW := 0.0 + minW, maxW := math.Inf(1), math.Inf(-1) + argW, best := 0, math.Inf(-1) + for i := range av { + w := dtCast(dt, av[i]) + sumW += w + minW = min(minW, w) + maxW = max(maxW, w) + if w > best { + best, argW = w, i + } + } + s := Sum(a) + switch dt { + case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int: + if s.IsFloat() || s.Int() != int64(sumW) { + t.Fatalf("Sum = %s, want int %v", s, sumW) + } + default: + if !s.IsFloat() || s.Float() != sumW { + t.Fatalf("Sum = %s, want float %v", s, sumW) + } + } + mean, err := Mean(a) + if err != nil || mean != sumW/float64(len(av)) { + t.Fatalf("Mean = %v, %v; want %v", mean, err, sumW/float64(len(av))) + } + mn, err := Min(a) + if err != nil || mn.Float() != minW { + t.Fatalf("Min = %v, %v; want %v", mn, err, minW) + } + mx, err := Max(a) + if err != nil || mx.Float() != maxW { + t.Fatalf("Max = %v, %v; want %v", mx, err, maxW) + } + am, err := ArgMax(a) + if err != nil || am != argW { + t.Fatalf("ArgMax = %d, %v; want %d", am, err, argW) + } + // Selecting and reshaping entries keep dtype and values. + cat := dtRun(t, "Concat", dtOf(Concat(a, b, 0))) + if cat.Dtype() != dt { + t.Fatalf("Concat dtype = %s, want %s", cat.Dtype(), dt) + } + dtAssertValues(t, "Concat", cat, append(dtWant(dt, av), dtWant(dt, bv)...)) + st := dtRun(t, "Stack", dtOf(Stack(a, b, 0))) + if st.Dtype() != dt || st.Shape()[0] != 2 { + t.Fatalf("Stack dtype %s shape %v", st.Dtype(), st.Shape()) + } + dtAssertValues(t, "Stack", st, append(dtWant(dt, av), dtWant(dt, bv)...)) + idx := dtMake(t, Int, []float64{4, 0, 2}, 3) + g1 := dtRun(t, "Gather", dtOf(Gather(a, 0, idx))) + if g1.Dtype() != dt { + t.Fatalf("Gather dtype = %s, want %s", g1.Dtype(), dt) + } + dtAssertValues(t, "Gather", g1, []float64{dtCast(dt, av[4]), dtCast(dt, av[0]), dtCast(dt, av[2])}) + tk := dtRun(t, "Take", dtOf(Take(a, idx))) + if tk.Dtype() != dt { + t.Fatalf("Take dtype = %s, want %s", tk.Dtype(), dt) + } + dtAssertValues(t, "Take", tk, []float64{dtCast(dt, av[4]), dtCast(dt, av[0]), dtCast(dt, av[2])}) + // Scatter(self, dim, index, src) writes src values into + // self at the indexed positions, one src element per index. + sc := dtRun(t, "Scatter", dtOf(Scatter(a, 0, idx, dtMake(t, dt, bv[:3], 3)))) + if sc.Dtype() != dt { + t.Fatalf("Scatter dtype = %s, want %s", sc.Dtype(), dt) + } + wantSc := dtWant(dt, av) + wantSc[4] = dtCast(dt, bv[0]) + wantSc[0] = dtCast(dt, bv[1]) + wantSc[2] = dtCast(dt, bv[2]) + dtAssertValues(t, "Scatter", sc, wantSc) + rv := Reverse(a) + if rv.Dtype() != dt { + t.Fatalf("Reverse dtype = %s, want %s", rv.Dtype(), dt) + } + dtAssertValues(t, "Reverse", rv, []float64{ + dtCast(dt, av[4]), dtCast(dt, av[3]), dtCast(dt, av[2]), + dtCast(dt, av[1]), dtCast(dt, av[0]), + }) + mat := dtMake(t, dt, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + tr := Transpose(mat) + if tr.Dtype() != dt || tr.Shape()[0] != 3 || tr.Shape()[1] != 2 { + t.Fatalf("Transpose dtype %s shape %v", tr.Dtype(), tr.Shape()) + } + dtAssertValues(t, "Transpose", tr, []float64{ + dtCast(dt, 1), dtCast(dt, 4), dtCast(dt, 2), + dtCast(dt, 5), dtCast(dt, 3), dtCast(dt, 6), + }) + cp := Copy(mat) + if cp.Dtype() != dt { + t.Fatalf("Copy dtype = %s", cp.Dtype()) + } + dtAssertValues(t, "Copy", cp, dtWiden(t, mat)) + sl := dtRun(t, "Slice", dtOf(Slice(mat, 0, 1, 2))) + if sl.Dtype() != dt { + t.Fatalf("Slice dtype = %s", sl.Dtype()) + } + dtAssertValues(t, "Slice", sl, []float64{dtCast(dt, 4), dtCast(dt, 5), dtCast(dt, 6)}) + // Zero-aware entries over a sparse sample. + sparse := dtMake(t, dt, []float64{0, 2, 0, 3, 5}, 5) + aw := dtRun(t, "Argwhere", dtOf(Argwhere(sparse))) + if aw.Dtype() != Int { + t.Fatalf("Argwhere dtype = %s, want Int", aw.Dtype()) + } + // Row-major coordinates of the non-zero widened elements. + wantCoords := [][]float64{} + for i, v := range []float64{0, 2, 0, 3, 5} { + if dtCast(dt, v) != 0 { + wantCoords = append(wantCoords, []float64{float64(i), 0}) + } + } + // Argwhere answers (nnz, ndim): one coordinate column per + // dimension, rank 1 here so a single column per row. + if aw.Shape()[1] != 1 { + t.Fatalf("Argwhere shape %v, want (nnz, 1)", aw.Shape()) + } + for r, wc := range wantCoords { + if v, err := FloatAt(aw, r, 0); err != nil || v != wc[0] { + t.Fatalf("Argwhere[%d] = %v, %v; want %v", r, v, err, wc[0]) + } + } + nz, err := Nonzero(sparse) + if err != nil { + t.Fatalf("Nonzero: %v", err) + } + if len(nz) != 1 || len(nz[0]) != len(wantCoords) { + t.Fatalf("Nonzero = %v, want %d entries on one axis", nz, len(wantCoords)) + } + for i, wc := range wantCoords { + if nz[0][i] != int(wc[0]) { + t.Fatalf("Nonzero[0][%d] = %d, want %d", i, nz[0][i], int(wc[0])) + } + } + cnt, err := CountNonzero(sparse) + if err != nil || cnt != len(wantCoords) { + t.Fatalf("CountNonzero = %d, %v; want %d", cnt, err, len(wantCoords)) + } + // BroadcastTo replicates a size-1 axis, dtype kept. + col := dtMake(t, dt, []float64{2, 3, 5}, 3, 1) + bc := dtRun(t, "BroadcastTo", dtOf(BroadcastTo(col, 3, 2))) + if bc.Dtype() != dt { + t.Fatalf("BroadcastTo dtype = %s", bc.Dtype()) + } + dtAssertValues(t, "BroadcastTo", bc, []float64{ + dtCast(dt, 2), dtCast(dt, 2), dtCast(dt, 3), + dtCast(dt, 3), dtCast(dt, 5), dtCast(dt, 5), + }) + }) + } + // Bool-operand surfaces through the facade: comparisons answer bool + // masks over Bool operands and Where consumes Bool conditions, + // both pinned in the per-dtype loop above; the Bool arithmetic + // refusals are pinned in the refusal census. The core logic + // entries And/Or/Xor/Not (internal/core/mask.go) ride + // facade_generated.go rows, pinned here on Bool operands for + // values and in the refusal census on Int operands by name. + ba := dtMake(t, Bool, []float64{1, 1, 0, 0}, 4) + bb := dtMake(t, Bool, []float64{1, 0, 1, 0}, 4) + m, err := Eq(ba, bb) + if err != nil { + t.Fatalf("Eq on Bool operands: %v", err) + } + if m.Dtype() != Bool { + t.Fatalf("Eq(bool, bool) mask dtype = %s, want Bool", m.Dtype()) + } + dtAssertValues(t, "Eq bool mask", m, []float64{1, 0, 0, 1}) + // And/Or/Xor/Not on Bool operands, through the facade names. + for _, op := range []struct { + name string + run func(x, y *Array) (*Array, error) + want []float64 + }{ + {"And", And, []float64{1, 0, 0, 0}}, + {"Or", Or, []float64{1, 1, 1, 0}}, + {"Xor", Xor, []float64{0, 1, 1, 0}}, + } { + got := dtRun(t, op.name, dtOf(op.run(ba, bb))) + if got.Dtype() != Bool { + t.Fatalf("%s(bool, bool) dtype = %s, want Bool", op.name, got.Dtype()) + } + dtAssertValues(t, op.name+" bool", got, op.want) + } + nb := dtRun(t, "Not", dtOf(Not(ba))) + if nb.Dtype() != Bool { + t.Fatalf("Not(bool) dtype = %s, want Bool", nb.Dtype()) + } + dtAssertValues(t, "Not bool", nb, []float64{0, 0, 1, 1}) + // The int8 wrap-around pin, explicitly: the payload arithmetic + // wraps exactly as Go's int8 arithmetic wraps. + i8a := dtMake(t, Int8, []float64{127, -128, 16}, 3) + i8b := dtMake(t, Int8, []float64{1, 1, 8}, 3) + sum8 := dtRun(t, "Add wrap", dtOf(Add(i8a, i8b))) + dtAssertValues(t, "Add int8 wrap", sum8, []float64{-128, -127, 24}) + dif8 := dtRun(t, "Sub wrap", dtOf(Sub(i8a, i8b))) + dtAssertValues(t, "Sub int8 wrap", dif8, []float64{126, 127, 8}) + prd8 := dtRun(t, "Mul wrap", dtOf(Mul(i8a, i8b))) + dtAssertValues(t, "Mul int8 wrap", prd8, []float64{127, -128, -128}) + u8a := dtMake(t, Uint8, []float64{255, 200}, 2) + u8b := dtMake(t, Uint8, []float64{1, 100}, 2) + sumU8 := dtRun(t, "Add uint8 wrap", dtOf(Add(u8a, u8b))) + dtAssertValues(t, "Add uint8 wrap", sumU8, []float64{0, 44}) +} + +// TestDtypesFacadeRefusals pins the refusal census: the core narrow-dtype +// deferrals read through the facade, the Bool-arithmetic wording, the +// preserved domain refusals and the Astype range and conversion texts. +func TestDtypesFacadeRefusals(t *testing.T) { + mk := func(dt Dtype, vals []float64, shape ...int) *Array { return dtMake(t, dt, vals, shape...) } + // identityOp is the callback shape the GMRES census row needs. + identityOp := func(v *Array) (*Array, error) { return v, nil } + harmonic := func(q *Array) (*Array, error) { return MulF(q, -1), nil } + cases := []struct { + name string + call func() error + want []string + }{ + {"MatMul2D int8", func() error { + _, err := MatMul2D(mk(Int8, []float64{1, 2, 3, 4}, 2, 2), mk(Int8, []float64{1, 0, 0, 1}, 2, 2)) + return err + }, []string{"MatMul", "int8", "convert with Astype"}}, + {"Einsum uint16", func() error { + _, err := Einsum("ij,jk->ik", mk(Uint16, []float64{1, 2, 3, 4}, 2, 2), mk(Uint16, []float64{1, 0, 0, 1}, 2, 2)) + return err + }, []string{"Einsum", "uint16", "convert with Astype"}}, + {"Sort int16", func() error { + _, err := Sort(mk(Int16, []float64{3, 1, 2}, 3)) + return err + }, []string{"Sort", "int16", "convert with Astype"}}, + {"ArgSort uint8", func() error { + _, err := ArgSort(mk(Uint8, []float64{3, 1, 2}, 3)) + return err + }, []string{"ArgSort", "uint8", "convert with Astype"}}, + {"Kron int32", func() error { + _, err := Kron(mk(Int32, []float64{1, 2}, 1, 2), mk(Int32, []float64{3, 4}, 1, 2)) + return err + }, []string{"Kron", "int32", "convert with Astype"}}, + {"Prod uint32", func() error { + _, err := Prod(mk(Uint32, []float64{2, 3}, 2), 0, false) + return err + }, []string{"Prod", "uint32", "convert with Astype"}}, + {"CumSum bool", func() error { + _, err := CumSum(mk(Bool, []float64{1, 0}, 2), 0) + return err + }, []string{"CumSum", "bool", "convert with Astype"}}, + {"CumProd int8", func() error { + _, err := CumProd(mk(Int8, []float64{2, 3}, 2), 0) + return err + }, []string{"CumProd", "int8", "convert with Astype"}}, + {"Norm uint16", func() error { + _, err := Norm(mk(Uint16, []float64{3, 4}, 2), 2, 0, false) + return err + }, []string{"Norm", "uint16", "convert with Astype"}}, + {"ClipI int16", func() error { + _, err := ClipI(mk(Int16, []float64{5, 50}, 2), 0, 10) + return err + }, []string{"ClipI", "int16", "convert with Astype"}}, + {"ClipF bool", func() error { + _, err := ClipF(mk(Bool, []float64{1, 0}, 2), 0, 1) + return err + }, []string{"ClipF", "bool", "convert with Astype"}}, + {"Sqrt int32", func() error { + _, err := Sqrt(mk(Int32, []float64{4, 9}, 2)) + return err + }, []string{"Sqrt", "int32", "convert with Astype"}}, + {"SparseFrom uint8", func() error { + _, err := SparseFrom(mk(Uint8, []float64{1, 0, 0, 1}, 2, 2)) + return err + }, []string{"SparseFrom", "uint8", "convert with Astype"}}, + {"NewSparseCOO int8 values", func() error { + ind := mk(Int, []float64{0, 0, 1, 1}, 2, 2) + _, err := NewSparseCOO(ind, mk(Int8, []float64{1, 1}, 2), []int{2, 2}) + return err + }, []string{"NewSparseCOO", "int8", "convert with Astype"}}, + {"SpMul narrow dense", func() error { + ind := mk(Int, []float64{0, 0, 1, 1}, 2, 2) + s, serr := NewSparseCOO(ind, mk(Float, []float64{2, 3}, 2), []int{2, 2}) + if serr != nil { + return serr + } + _, err := SpMul(s, mk(Int16, []float64{4, 5, 6, 7}, 2, 2)) + return err + }, []string{"SpMul", "int16", "convert with Astype"}}, + {"SpMatMul narrow operand", func() error { + ind := mk(Int, []float64{0, 0, 1, 1}, 2, 2) + s, serr := NewSparseCOO(ind, mk(Float, []float64{2, 3}, 2), []int{2, 2}) + if serr != nil { + return serr + } + _, err := SpMatMul(s, mk(Uint8, []float64{4, 5, 6, 7}, 2, 2)) + return err + }, []string{"SpMatMul", "uint8", "convert with Astype"}}, + {"Add bool arithmetic", func() error { + _, err := Add(mk(Bool, []float64{1, 0}, 2), mk(Bool, []float64{1, 1}, 2)) + return err + }, []string{"Add", "bool arrays have no arithmetic"}}, + {"Sub bool arithmetic", func() error { + _, err := Sub(mk(Bool, []float64{1, 0}, 2), mk(Bool, []float64{1, 1}, 2)) + return err + }, []string{"Sub", "bool arrays have no arithmetic"}}, + {"Mul bool arithmetic", func() error { + _, err := Mul(mk(Bool, []float64{1, 0}, 2), mk(Bool, []float64{1, 1}, 2)) + return err + }, []string{"Mul", "bool arrays have no arithmetic"}}, + {"Div bool arithmetic", func() error { + _, err := Div(mk(Bool, []float64{1, 0}, 2), mk(Bool, []float64{1, 1}, 2)) + return err + }, []string{"Div", "bool arrays have no arithmetic"}}, + // The domain refusals, existing wording preserved. + {"GMRES narrow b", func() error { + _, err := GMRES(identityOp, mk(Int8, []float64{1, 2}, 2), 2, 5, 1e-10) + return err + }, []string{"GMRES", "b must be a real dtype", "int8"}}, + {"DetComplex narrow", func() error { + _, err := DetComplex(mk(Int8, []float64{1, 0, 0, 1}, 2, 2)) + return err + }, []string{"DetComplex", "complex"}}, + {"SolvePoissonPeriodic narrow", func() error { + _, err := SolvePoissonPeriodic(mk(Uint16, []float64{ + 1, 2, 3, 4, 4, 3, 2, 1, 1, 2, 3, 4, 4, 3, 2, 1, + }, 4, 4), 2*math.Pi, 2*math.Pi) + return err + }, []string{"SolvePoissonPeriodic", "float64 array", "uint16"}}, + {"NewTriangleMesh2D narrow connectivity", func() error { + verts := mk(Float, []float64{0, 0, 1, 0, 0, 1}, 3, 2) + _, err := NewTriangleMesh2D(verts, mk(Int32, []float64{0, 1, 2}, 1, 3)) + return err + }, []string{"NewTriangleMesh2D", "must hold integers", "int32"}}, + {"IntegrateVerlet int state", func() error { + _, _, err := IntegrateVerlet(harmonic, 0, 1, mk(Int, []float64{1}, 1), mk(Float, []float64{0}, 1), 5) + return err + }, []string{"IntegrateVerlet", "int states", "float or float32"}}, + {"IntegrateVerlet narrow state", func() error { + _, _, err := IntegrateVerlet(harmonic, 0, 1, mk(Int8, []float64{1}, 1), mk(Float, []float64{0}, 1), 5) + return err + }, []string{"IntegrateVerlet", "int states", "float or float32"}}, + {"Quantile narrow rides the Sort deferral", func() error { + _, err := Quantile(mk(Int8, []float64{3, 1, 2}, 3), []float64{0.5}) + return err + }, []string{"Sort", "int8", "convert with Astype"}}, + {"SpSolveComplexCG narrow b", func() error { + ind := mk(Int, []float64{0, 0, 1, 1}, 2, 2) + s, serr := NewSparseCOO(ind, mk(Complex, []float64{2, 3}, 2), []int{2, 2}) + if serr != nil { + return serr + } + _, err := SpSolveComplexCG(s, mk(Uint16, []float64{1, 1}, 2), 1e-10, 10) + return err + }, []string{"complex128", "uint16"}}, + // Astype through the facade: range errors, zero-test bool + // targets and the complex refusals. + {"Astype int to int8 range", func() error { + _, err := Astype(mk(Int, []float64{200}, 1), Int8) + return err + }, []string{"Astype", "value 200", "index 0", "does not fit int8"}}, + {"Astype float to uint8 fractional", func() error { + _, err := Astype(mk(Float, []float64{1.5}, 1), Uint8) + return err + }, []string{"Astype", "value 1.5", "does not fit uint8"}}, + {"Astype NaN to int16", func() error { + _, err := Astype(mk(Float, []float64{math.NaN()}, 1), Int16) + return err + }, []string{"Astype", "does not fit int16"}}, + {"Astype uint32 to int8 range", func() error { + _, err := Astype(mk(Uint32, []float64{4294967295}, 1), Int8) + return err + }, []string{"Astype", "does not fit int8"}}, + {"Astype complex to int8", func() error { + _, err := Astype(mk(Complex, []float64{1, 2}, 2), Int8) + return err + }, []string{"Astype", "cannot narrow complex to int8"}}, + {"Astype complex to uint16", func() error { + _, err := Astype(mk(Complex, []float64{1, 2}, 2), Uint16) + return err + }, []string{"Astype", "cannot narrow complex to uint16"}}, + {"Astype complex to float16", func() error { + _, err := Astype(mk(Complex, []float64{1, 2}, 2), Float16) + return err + }, []string{"Astype", "cannot narrow complex to float16"}}, + // The transcendental family's narrow-dtype refusals through the + // facade rows: the refusal names the dtype and the conversion. + {"Exp int8", func() error { + _, err := Exp(mk(Int8, []float64{1, 2}, 2)) + return err + }, []string{"Exp", "int8", "convert with Astype"}}, + {"Log uint16", func() error { + _, err := Log(mk(Uint16, []float64{1, 2}, 2)) + return err + }, []string{"Log", "uint16", "convert with Astype"}}, + {"Floor int32", func() error { + _, err := Floor(mk(Int32, []float64{1, 2}, 2)) + return err + }, []string{"Floor", "int32", "convert with Astype"}}, + {"Ceil bool", func() error { + _, err := Ceil(mk(Bool, []float64{1, 0}, 2)) + return err + }, []string{"Ceil", "bool", "convert with Astype"}}, + // The logic rows refuse non-bool operands by name. + {"And int operands", func() error { + _, err := And(mk(Int, []float64{1, 0}, 2), mk(Int, []float64{1, 1}, 2)) + return err + }, []string{"And", "operands must be bool arrays", "int"}}, + {"Or int operands", func() error { + _, err := Or(mk(Int, []float64{1, 0}, 2), mk(Int, []float64{1, 1}, 2)) + return err + }, []string{"Or", "operands must be bool arrays", "int"}}, + {"Xor int8 operands", func() error { + _, err := Xor(mk(Int8, []float64{1, 0}, 2), mk(Int8, []float64{1, 1}, 2)) + return err + }, []string{"Xor", "operands must be bool arrays", "int8"}}, + {"Not int operand", func() error { + _, err := Not(mk(Int, []float64{1, 0}, 2)) + return err + }, []string{"Not", "operands must be bool arrays", "int"}}, + // Interpolate2D refuses a narrow grid at the entry. + {"Interpolate2D int8 grid", func() error { + _, err := Interpolate2D(mk(Int8, []float64{1, 2, 3, 4}, 2, 2), + mk(Float, []float64{0.5}, 1), mk(Float, []float64{0.5}, 1), 0, 0, 1, 1) + return err + }, []string{"Interpolate2D", "int8", "convert with Astype"}}, + // Pow on a mixed integer-class pair refuses a negative + // exponent wherever promote lands; PowI refuses the narrow + // widths with the conversion wording. + {"Pow int with int8 negative exponent", func() error { + _, err := Pow(mk(Int, []float64{2, 3}, 2), mk(Int8, []float64{-3, 4}, 2)) + return err + }, []string{"Pow", "negative exponent"}}, + {"Pow uint32 with int16 negative exponent", func() error { + _, err := Pow(mk(Uint32, []float64{2, 3}, 2), mk(Int16, []float64{-1, 2}, 2)) + return err + }, []string{"Pow", "negative exponent"}}, + {"PowI int8", func() error { + _, err := PowI(mk(Int8, []float64{2, 3}, 2), 2) + return err + }, []string{"PowI", "int8", "convert with Astype"}}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + err := c.call() + if err == nil { + t.Fatalf("accepted; want a refusal carrying %v", c.want) + } + msg := err.Error() + for _, w := range c.want { + if !strings.Contains(msg, w) { + t.Fatalf("error %q does not contain %q", msg, w) + } + } + }) + } + // Abs answers nil on the narrow widths: the no-error entry's + // refusal shape, no panic and no silent zero array. + if got := Abs(mk(Int8, []float64{-1, 2}, 2)); got != nil { + t.Fatalf("Abs(int8) = %v, want nil", got) + } + if got := Abs(mk(Bool, []float64{1, 0}, 2)); got != nil { + t.Fatalf("Abs(bool) = %v, want nil", got) + } + // Astype bool targets: the zero-test, every source reaching bool + // without an error, NaN included. + b1 := dtRun(t, "Astype float to bool", dtOf(Astype(mk(Float, []float64{0, 2.5, -3, math.NaN()}, 4), Bool))) + for i, want := range []bool{false, true, true, true} { + got, err := BoolAt(b1, i) + if err != nil || got != want { + t.Fatalf("BoolAt(%d) = %v, %v; want %v", i, got, err, want) + } + } + b2 := dtRun(t, "Astype int to bool", dtOf(Astype(mk(Int, []float64{-5, 0, 7}, 3), Bool))) + for i, want := range []bool{true, false, true} { + got, err := BoolAt(b2, i) + if err != nil || got != want { + t.Fatalf("BoolAt(%d) = %v, %v; want %v", i, got, err, want) + } + } + b3 := dtRun(t, "Astype complex to bool", dtOf(Astype(mk(Complex, []float64{0, 3}, 2), Bool))) + for i, want := range []bool{false, true} { + got, err := BoolAt(b3, i) + if err != nil || got != want { + t.Fatalf("BoolAt(%d) = %v, %v; want %v", i, got, err, want) + } + } + b4 := dtRun(t, "Astype uint32 to bool", dtOf(Astype(mk(Uint32, []float64{0, 4294967295}, 2), Bool))) + for i, want := range []bool{false, true} { + got, err := BoolAt(b4, i) + if err != nil || got != want { + t.Fatalf("BoolAt(%d) = %v, %v; want %v", i, got, err, want) + } + } + // Astype widenings stay exact through the facade. + w1 := dtRun(t, "Astype int8 to int", dtOf(Astype(mk(Int8, []float64{-1, 127}, 2), Int))) + dtAssertValues(t, "int8 to int", w1, []float64{-1, 127}) + w2 := dtRun(t, "Astype uint32 to int", dtOf(Astype(mk(Uint32, []float64{4294967295}, 1), Int))) + dtAssertValues(t, "uint32 to int", w2, []float64{4294967295}) + w3 := dtRun(t, "Astype uint8 to complex", dtOf(Astype(mk(Uint8, []float64{255}, 1), Complex))) + cv, err := ComplexAt(w3, 0) + if err != nil || cv != complex(255, 0) { + t.Fatalf("uint8 to complex = %v, %v; want (255+0i)", cv, err) + } + // Same-dtype Astype copies unchanged. + same := dtRun(t, "Astype int16 to int16", dtOf(Astype(mk(Int16, []float64{-300, 300}, 2), Int16))) + if same.Dtype() != Int16 { + t.Fatalf("same-dtype Astype dtype = %s", same.Dtype()) + } + dtAssertValues(t, "same-dtype Astype", same, []float64{-300, 300}) +} + +// dtNoPanic runs an (array, error) call, converting a panic into a +// test failure instead of taking the shared test binary down with it. +// The mixed-pair pins below were written while the core elementwise +// kernel panicked on the Int-promoting pairs; the pin is the +// contract, and the guard turns any returning panic into a failure +// that names the call rather than a dead test binary. +func dtNoPanic(t *testing.T, label string, call func() (*Array, error)) (a *Array, err error) { + t.Helper() + defer func() { + if r := recover(); r != nil { + t.Fatalf("%s: panicked: %v", label, r) + } + }() + return call() +} + +// TestDtypesFacadeStringPin pins the diagnostic rendering of all +// twelve dtypes through the facade constants, the seven narrow names +// alongside the five the ladder already carried. +func TestDtypesFacadeStringPin(t *testing.T) { + all := []Dtype{Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int, Float16, Float32, Float, Complex} + for _, d := range all { + if got := d.String(); got != dtStrings[d] { + t.Fatalf("Dtype(%d).String() = %q, want %q", d, got, dtStrings[d]) + } + } + // Zeros and Ones answer the requested dtype for every one of the + // twelve, not only the seven narrow ones. + for _, d := range all { + z, err := Zeros(d, 3) + if err != nil || z.Dtype() != d || z.Len() != 3 { + t.Fatalf("Zeros(%s, 3) = %s len %d, %v", d, z.Dtype(), z.Len(), err) + } + o, err := Ones(d, 2) + if err != nil || o.Dtype() != d { + t.Fatalf("Ones(%s, 2) = %s, %v", d, o.Dtype(), err) + } + } +} + +// TestDtypesFacadePromote pins promotion through public operations +// only: mixed-dtype calls whose result dtype and values are the +// facade's answer to the promote question. +func TestDtypesFacadePromote(t *testing.T) { + cases := []struct { + name string + a, b Dtype + want Dtype + av, bv []float64 + wantVal []float64 + }{ + {"int8+uint8", Int8, Uint8, Int16, []float64{100}, []float64{200}, []float64{300}}, + {"int16+uint16", Int16, Uint16, Int32, []float64{-30000}, []float64{60000}, []float64{30000}}, + {"int32+uint32", Int32, Uint32, Int, []float64{2147483647}, []float64{1}, []float64{2147483648}}, + {"uint8+int16", Uint8, Int16, Int16, []float64{250}, []float64{-300}, []float64{-50}}, + {"uint16+int32", Uint16, Int32, Int32, []float64{65535}, []float64{100000}, []float64{165535}}, + {"uint32+int", Uint32, Int, Int, []float64{4294967295}, []float64{1}, []float64{4294967296}}, + {"bool+int8", Bool, Int8, Int8, []float64{1}, []float64{7}, []float64{8}}, + {"bool+uint16", Bool, Uint16, Uint16, []float64{1}, []float64{700}, []float64{701}}, + {"int8+float16", Int8, Float16, Float16, []float64{3}, []float64{0.5}, []float64{3.5}}, + {"uint32+float16", Uint32, Float16, Float16, []float64{2}, []float64{0.5}, []float64{2.5}}, + {"float16+float32", Float16, Float32, Float32, []float64{0.5}, []float64{0.25}, []float64{0.75}}, + {"float32+float", Float32, Float, Float, []float64{0.5}, []float64{0.25}, []float64{0.75}}, + {"int8+complex", Int8, Complex, Complex, []float64{3}, []float64{2}, []float64{5}}, + {"bool+float", Bool, Float, Float, []float64{1}, []float64{2.5}, []float64{3.5}}, + {"int+float16", Int, Float16, Float16, []float64{4}, []float64{0.5}, []float64{4.5}}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + a := dtMake(t, c.a, c.av, len(c.av)) + b := dtMake(t, c.b, c.bv, len(c.bv)) + sum := dtRun(t, "Add", dtOf(dtNoPanic(t, "Add "+c.name, func() (*Array, error) { return Add(a, b) }))) + if sum.Dtype() != c.want { + t.Fatalf("Add(%s, %s) dtype = %s, want %s", c.a, c.b, sum.Dtype(), c.want) + } + if c.want == Complex { + got, err := ComplexAt(sum, 0) + if err != nil || real(got) != c.wantVal[0] || imag(got) != 0 { + t.Fatalf("Add value = %v, %v; want %v", got, err, c.wantVal[0]) + } + } else { + dtAssertValues(t, "Add", sum, c.wantVal) + } + // Sub and Mul walk the same promote rule. + prod := dtRun(t, "Mul", dtOf(dtNoPanic(t, "Mul "+c.name, func() (*Array, error) { return Mul(a, b) }))) + if prod.Dtype() != c.want { + t.Fatalf("Mul(%s, %s) dtype = %s, want %s", c.a, c.b, prod.Dtype(), c.want) + } + if c.want != Complex { + // The product computes in the promoted dtype and wraps + // there exactly as the payload store wraps: uint8 with + // int16 promotes to int16, whose value space the + // product -75000 leaves by wrapping to -9464. + pv := make([]float64, len(c.av)) + for i := range c.av { + pv[i] = dtCast(c.want, c.av[i]*c.bv[i]) + } + dtAssertValues(t, "Mul", prod, pv) + } + }) + } + // Where promotes across its value operands the same way. + cond := dtMake(t, Bool, []float64{1, 0}, 2) + x := dtMake(t, Int8, []float64{7, 7}, 2) + y := dtMake(t, Uint16, []float64{700, 700}, 2) + w := dtRun(t, "Where mixed", dtOf(Where(cond, x, y))) + if w.Dtype() != Int32 { + t.Fatalf("Where(int8, uint16) dtype = %s, want int32 (the mixed operands promote through the containment table)", w.Dtype()) + } + dtAssertValues(t, "Where mixed", w, []float64{7, 700}) + // Sum's scalar kind is part of the promote contract: the integer + // class answers an Int scalar whatever the narrow width. + for _, dt := range dtNewWithInt { + if dt == Bool { + continue + } + a := dtMake(t, dt, []float64{1, 2, 3}, 3) + s := Sum(a) + if s.IsFloat() || s.Int() != 6 { + t.Fatalf("Sum(%s) = %s, want int 6", dt, s) + } + } + bs := Sum(dtMake(t, Bool, []float64{1, 0, 1, 1}, 4)) + if bs.IsFloat() || bs.Int() != 3 { + t.Fatalf("Sum(bool) = %s, want int 3 (the count of trues)", bs) + } +} + +// TestDtypesFacadeRowsPresent calls every facade row the dtypes wave +// added, once each, through the facade name: the seven dtype +// constants, the fourteen narrow constructors, BoolAt and the four +// logic rows. A row deleted from facade_generated.go breaks this +// test's compilation, so an unreferenced row cannot vanish from the +// facade while the suite stays green. +func TestDtypesFacadeRowsPresent(t *testing.T) { + // The seven dtype constants. + consts := map[Dtype]string{ + Bool: "bool", Int8: "int8", Uint8: "uint8", Int16: "int16", + Uint16: "uint16", Int32: "int32", Uint32: "uint32", + } + for dt, name := range consts { + if got := dt.String(); got != name { + t.Fatalf("Dtype constant %s rendered %q", name, got) + } + } + // The fourteen constructors, each called once. + must := func(a *Array, err error) *Array { + if err != nil { + t.Fatalf("constructor: %v", err) + } + return a + } + ctors := map[string]*Array{ + "FromBools": must(FromBools([]bool{true, false}, 2)), + "FromInt8s": must(FromInt8s([]int8{1, -2}, 2)), + "FromUint8s": must(FromUint8s([]uint8{1, 2}, 2)), + "FromInt16s": must(FromInt16s([]int16{1, -2}, 2)), + "FromUint16s": must(FromUint16s([]uint16{1, 2}, 2)), + "FromInt32s": must(FromInt32s([]int32{1, -2}, 2)), + "FromUint32s": must(FromUint32s([]uint32{1, 2}, 2)), + "BoolsFromArray": must(BoolsFromArray([]bool{true, false}, 2)), + "Int8sFromArray": must(Int8sFromArray([]int8{1, -2}, 2)), + "Uint8sFromArray": must(Uint8sFromArray([]uint8{1, 2}, 2)), + "Int16sFromArray": must(Int16sFromArray([]int16{1, -2}, 2)), + "Uint16sFromArray": must(Uint16sFromArray([]uint16{1, 2}, 2)), + "Int32sFromArray": must(Int32sFromArray([]int32{1, -2}, 2)), + "Uint32sFromArray": must(Uint32sFromArray([]uint32{1, 2}, 2)), + } + if len(ctors) != 14 { + t.Fatalf("constructor census = %d rows, want 14", len(ctors)) + } + for name, a := range ctors { + if a == nil || a.Len() != 2 { + t.Fatalf("%s answered %v, want a two-element array", name, a) + } + } + // BoolAt, once. + if v, err := BoolAt(ctors["FromBools"], 0); err != nil || !v { + t.Fatalf("BoolAt = %v, %v; want true, nil", v, err) + } + // The four logic rows, once each. + ba := must(FromBools([]bool{true, false}, 2)) + if a, err := And(ba, ba); err != nil || a.Dtype() != Bool { + t.Fatalf("And = %v, %v", a, err) + } + if a, err := Or(ba, ba); err != nil || a.Dtype() != Bool { + t.Fatalf("Or = %v, %v", a, err) + } + if a, err := Xor(ba, ba); err != nil || a.Dtype() != Bool { + t.Fatalf("Xor = %v, %v", a, err) + } + if a, err := Not(ba); err != nil || a.Dtype() != Bool { + t.Fatalf("Not = %v, %v", a, err) + } +} + +// TestDtypesFacadeNarrowValuePins pins the narrow-width value answers +// through the facade: Unique's dedupe and dtype keeping, the extrema +// in uint32's comparison space, Astype at the target dtype's +// boundary, Elements[int64]'s exact widening of the integer class, and +// Abs's nil refusal on a narrow width. +func TestDtypesFacadeNarrowValuePins(t *testing.T) { + mk := func(dt Dtype, vals []float64, shape ...int) *Array { return dtMake(t, dt, vals, shape...) } + // Unique on int8: sorted, deduped, dtype kept. + ui := dtRun(t, "Unique int8", dtOf(Unique(mk(Int8, []float64{3, 1, 3, 1, 2}, 5)))) + if ui.Dtype() != Int8 { + t.Fatalf("Unique(int8) dtype = %s, want int8", ui.Dtype()) + } + dtAssertValues(t, "Unique int8", ui, []float64{1, 2, 3}) + // Unique on bool: false orders before true, dtype kept. + ub := dtRun(t, "Unique bool", dtOf(Unique(mk(Bool, []float64{1, 0, 1}, 3)))) + if ub.Dtype() != Bool { + t.Fatalf("Unique(bool) dtype = %s, want bool", ub.Dtype()) + } + dtAssertValues(t, "Unique bool", ub, []float64{0, 1}) + // The extrema in uint32's comparison space: no float32 transit may + // round 2^32-1 onto 2^32. + mx, err := Max(mk(Uint32, []float64{4294967295, 1}, 2)) + if err != nil || mx.Float() != 4294967295 { + t.Fatalf("Max(uint32) = %v, %v; want 4294967295", mx, err) + } + mn, err := Min(mk(Uint32, []float64{0, 4294967295}, 2)) + if err != nil || mn.Float() != 0 { + t.Fatalf("Min(uint32) = %v, %v; want 0", mn, err) + } + // Astype at the target boundary: 127 fits int8 exactly. + ba := dtRun(t, "Astype boundary", dtOf(Astype(mk(Int, []float64{127}, 1), Int8))) + if ba.Dtype() != Int8 || ba.RawInt8s()[0] != 127 { + t.Fatalf("Astype(Int{127}, Int8) = %s %v, want int8 [127]", ba.Dtype(), ba.RawInt8s()) + } + // Elements[int64] widens the integer class exactly, the IntAt rule. + ei, err := mk(Int8, []float64{-128, -1, 0, 1, 127}, 5).Elements[int64]() + if err != nil { + t.Fatalf("Elements[int64] on int8: %v", err) + } + for i, w := range []int64{-128, -1, 0, 1, 127} { + if ei[i] != w { + t.Fatalf("Elements[int64] int8[%d] = %d, want %d", i, ei[i], w) + } + } + // Abs answers nil on a narrow width: the no-error entry's refusal + // shape, also pinned in the refusal census above. + if got := Abs(mk(Uint16, []float64{1, 2}, 2)); got != nil { + t.Fatalf("Abs(uint16) = %v, want nil", got) + } +} diff --git a/example_test.go b/example_test.go new file mode 100644 index 0000000..f6e00ef --- /dev/null +++ b/example_test.go @@ -0,0 +1,308 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/examples/deconv/main.go b/examples/deconv/main.go new file mode 100644 index 0000000..b686ce1 --- /dev/null +++ b/examples/deconv/main.go @@ -0,0 +1,165 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/examples/fft/main.go b/examples/fft/main.go new file mode 100644 index 0000000..3a6db53 --- /dev/null +++ b/examples/fft/main.go @@ -0,0 +1,62 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/examples/fits/main.go b/examples/fits/main.go new file mode 100644 index 0000000..01dd394 --- /dev/null +++ b/examples/fits/main.go @@ -0,0 +1,98 @@ +// Copyright (c) 2026 Petr Balvín (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)") +} diff --git a/examples/helmholtz/main.go b/examples/helmholtz/main.go new file mode 100644 index 0000000..fbda728 --- /dev/null +++ b/examples/helmholtz/main.go @@ -0,0 +1,242 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/examples/hmc/main.go b/examples/hmc/main.go new file mode 100644 index 0000000..757feef --- /dev/null +++ b/examples/hmc/main.go @@ -0,0 +1,102 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/examples/netcdf/main.go b/examples/netcdf/main.go new file mode 100644 index 0000000..a24a607 --- /dev/null +++ b/examples/netcdf/main.go @@ -0,0 +1,108 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/examples/ode-fit/main.go b/examples/ode-fit/main.go new file mode 100644 index 0000000..bd186f1 --- /dev/null +++ b/examples/ode-fit/main.go @@ -0,0 +1,135 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Command ode-fit fits the parameters of a damped oscillator to +// endpoint measurements by the adjoint method: AdjointODE hands back +// dL/dθ for every parameter at the cost of one extra solve, and plain +// gradient descent walks the parameters to the truth. This is the +// data-assimilation loop no other Go library expresses. +// +// Usage: go run ./examples/ode-fit +package main + +import ( + "fmt" + "log" + "math" + + "sourcedock.dev/petrbalvin/tensor" +) + +// coreDynamics is y' = (v, −c·v − ω²·y) over plain arrays, the shape +// IntegrateODE wants; it generates the data. +func coreDynamics(c, omega float64) func(float64, *tensor.Array) (*tensor.Array, error) { + return func(t float64, y *tensor.Array) (*tensor.Array, error) { + return tensor.FromFloats([]float64{ + y.FloatAt(1), + -c*y.FloatAt(1) - omega*omega*y.FloatAt(0), + }, 2) + } +} + +// graphDynamics is the same equation with (c, ω) as differentiable +// leaves, the shape AdjointODE wants. +func graphDynamics(c, omega *tensor.Tensor) func(float64, *tensor.Tensor) (*tensor.Tensor, error) { + return func(t float64, y *tensor.Tensor) (*tensor.Tensor, error) { + yv, err := y.Slice(0, 0, 1) + if err != nil { + return nil, err + } + vv, err := y.Slice(0, 1, 2) + if err != nil { + return nil, err + } + w2, err := omega.Pow(2) + if err != nil { + return nil, err + } + damping, err := c.Mul(vv) + if err != nil { + return nil, err + } + restoring, err := w2.Mul(yv) + if err != nil { + return nil, err + } + acc, err := damping.Add(restoring) + if err != nil { + return nil, err + } + neg, err := acc.Scale(-1) + if err != nil { + return nil, err + } + return vv.Concat(neg, 0) + } +} + +func main() { + const ( + trueC = 0.8 + trueOmega = 3.0 + ) + y0, err := tensor.FromFloats([]float64{1, 0}, 2) + if err != nil { + log.Fatal(err) + } + + // Measurements of y(t) at three times from the true system. + times := []float64{0.4, 0.8, 1.2, 1.6, 2.0, 2.4} + data := make([]float64, len(times)) + trueF := coreDynamics(trueC, trueOmega) + for i, T := range times { + end, err := tensor.IntegrateODE(trueF, 0, T, y0, tensor.ODEOptions{}) + if err != nil { + log.Fatal(err) + } + data[i] = end.FloatAt(0) + } + + // The fit: gradient descent on L = Σ (y(Tᵢ; θ) − dataᵢ)² with the + // gradient from one adjoint pass per data point. + cVal, wVal := 0.3, 1.8 + const rate = 0.04 + for iter := 1; iter <= 400; iter++ { + gc, gw := 0.0, 0.0 + loss := 0.0 + forwardF := coreDynamics(cVal, wVal) + cArr, _ := tensor.FromFloats([]float64{cVal}, 1) + wArr, _ := tensor.FromFloats([]float64{wVal}, 1) + c := tensor.FromArray(cArr, true) + omega := tensor.FromArray(wArr, true) + adjF := graphDynamics(c, omega) + for i, T := range times { + end, err := tensor.IntegrateODE(forwardF, 0, T, y0, tensor.ODEOptions{}) + if err != nil { + log.Fatal(err) + } + res := end.FloatAt(0) - data[i] + loss += res * res + // dL/dy(T) = 2·res on the position component only; the + // velocity component carries no loss. + seed, err := tensor.FromFloats([]float64{2 * res, 0}, 2) + if err != nil { + log.Fatal(err) + } + _, paramGrads, err := tensor.AdjointODE(adjF, []*tensor.Tensor{c, omega}, + 0, T, y0, seed, tensor.ODEOptions{}) + if err != nil { + log.Fatal(err) + } + gc += paramGrads[0].FloatAt(0) + gw += paramGrads[1].FloatAt(0) + } + if iter%100 == 0 { + fmt.Printf("iter %3d c = %.4f omega = %.4f loss = %.3e\n", iter, cVal, wVal, loss) + } + cVal -= rate * gc + wVal -= rate * gw + } + fmt.Printf("fitted c = %.4f (true %.4f), omega = %.4f (true %.4f)\n", + cVal, trueC, wVal, trueOmega) + if math.Abs(cVal-trueC) > 0.05 || math.Abs(wVal-trueOmega) > 0.05 { + log.Fatal("the fit did not converge to the truth") + } +} diff --git a/examples/pde/main.go b/examples/pde/main.go new file mode 100644 index 0000000..0e8ba8d --- /dev/null +++ b/examples/pde/main.go @@ -0,0 +1,103 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/examples/pendulum/main.go b/examples/pendulum/main.go new file mode 100644 index 0000000..0ebc212 --- /dev/null +++ b/examples/pendulum/main.go @@ -0,0 +1,72 @@ +// Copyright (c) 2026 Petr Balvín (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) +} diff --git a/examples/qmc/main.go b/examples/qmc/main.go new file mode 100644 index 0000000..cc94937 --- /dev/null +++ b/examples/qmc/main.go @@ -0,0 +1,77 @@ +// Copyright (c) 2026 Petr Balvín (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") +} diff --git a/examples/regression/main.go b/examples/regression/main.go new file mode 100644 index 0000000..0c130b5 --- /dev/null +++ b/examples/regression/main.go @@ -0,0 +1,118 @@ +// Copyright (c) 2026 Petr Balvín (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) +} diff --git a/examples/spectral/main.go b/examples/spectral/main.go new file mode 100644 index 0000000..b9c4d81 --- /dev/null +++ b/examples/spectral/main.go @@ -0,0 +1,136 @@ +// Copyright (c) 2026 Petr Balvín (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] } diff --git a/examples/wavelets/main.go b/examples/wavelets/main.go new file mode 100644 index 0000000..c2d8b3d --- /dev/null +++ b/examples/wavelets/main.go @@ -0,0 +1,139 @@ +// Copyright (c) 2026 Petr Balvín (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]) + } +} diff --git a/facade_generated.go b/facade_generated.go new file mode 100644 index 0000000..148f920 --- /dev/null +++ b/facade_generated.go @@ -0,0 +1,785 @@ +// Copyright (c) 2026 Petr Balvín (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 diff --git a/facade_helpers_test.go b/facade_helpers_test.go new file mode 100644 index 0000000..4f243a5 --- /dev/null +++ b/facade_helpers_test.go @@ -0,0 +1,24 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/facade_test.go b/facade_test.go new file mode 100644 index 0000000..c50adbd --- /dev/null +++ b/facade_test.go @@ -0,0 +1,80 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..120703e --- /dev/null +++ b/go.mod @@ -0,0 +1,3 @@ +module sourcedock.dev/petrbalvin/tensor + +go 1.27.1 diff --git a/grad/adjoint.go b/grad/adjoint.go new file mode 100644 index 0000000..4ab893b --- /dev/null +++ b/grad/adjoint.go @@ -0,0 +1,313 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/adjoint_test.go b/grad/adjoint_test.go new file mode 100644 index 0000000..49b5fb8 --- /dev/null +++ b/grad/adjoint_test.go @@ -0,0 +1,349 @@ +// Copyright (c) 2026 Petr Balvín (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)) + } +} diff --git a/grad/backwardcoverage_test.go b/grad/backwardcoverage_test.go new file mode 100644 index 0000000..7209854 --- /dev/null +++ b/grad/backwardcoverage_test.go @@ -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) + } +} diff --git a/grad/batched.go b/grad/batched.go new file mode 100644 index 0000000..32cab20 --- /dev/null +++ b/grad/batched.go @@ -0,0 +1,165 @@ +// Copyright (c) 2026 Petr Balvín (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)) + } + } +} diff --git a/grad/batched_test.go b/grad/batched_test.go new file mode 100644 index 0000000..0b5ae51 --- /dev/null +++ b/grad/batched_test.go @@ -0,0 +1,188 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/bench_perf_test.go b/grad/bench_perf_test.go new file mode 100644 index 0000000..a44a345 --- /dev/null +++ b/grad/bench_perf_test.go @@ -0,0 +1,292 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/grad/bench_test.go b/grad/bench_test.go new file mode 100644 index 0000000..4907930 --- /dev/null +++ b/grad/bench_test.go @@ -0,0 +1,112 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/grad/broadcast_backward_pins_test.go b/grad/broadcast_backward_pins_test.go new file mode 100644 index 0000000..b0ddb64 --- /dev/null +++ b/grad/broadcast_backward_pins_test.go @@ -0,0 +1,145 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/grad/complex.go b/grad/complex.go new file mode 100644 index 0000000..25bddd7 --- /dev/null +++ b/grad/complex.go @@ -0,0 +1,503 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/complex_test.go b/grad/complex_test.go new file mode 100644 index 0000000..6ad5796 --- /dev/null +++ b/grad/complex_test.go @@ -0,0 +1,539 @@ +// Copyright (c) 2026 Petr Balvín (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) +} diff --git a/grad/concat.go b/grad/concat.go new file mode 100644 index 0000000..b3e9769 --- /dev/null +++ b/grad/concat.go @@ -0,0 +1,128 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/concat_test.go b/grad/concat_test.go new file mode 100644 index 0000000..f1c3850 --- /dev/null +++ b/grad/concat_test.go @@ -0,0 +1,167 @@ +// Copyright (c) 2026 Petr Balvín (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]) + } + } +} diff --git a/grad/data_ops_test.go b/grad/data_ops_test.go new file mode 100644 index 0000000..871b737 --- /dev/null +++ b/grad/data_ops_test.go @@ -0,0 +1,200 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/grad/doc.go b/grad/doc.go new file mode 100644 index 0000000..5151235 --- /dev/null +++ b/grad/doc.go @@ -0,0 +1,82 @@ +// Copyright (c) 2026 Petr Balvín (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 diff --git a/grad/example_test.go b/grad/example_test.go new file mode 100644 index 0000000..286fad4 --- /dev/null +++ b/grad/example_test.go @@ -0,0 +1,218 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/fuzz_test.go b/grad/fuzz_test.go new file mode 100644 index 0000000..39c55d0 --- /dev/null +++ b/grad/fuzz_test.go @@ -0,0 +1,187 @@ +// Copyright (c) 2026 Petr Balvín (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 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)) + } + } + }) +} diff --git a/grad/grad_guard_pins_test.go b/grad/grad_guard_pins_test.go new file mode 100644 index 0000000..80471ee --- /dev/null +++ b/grad/grad_guard_pins_test.go @@ -0,0 +1,875 @@ +// Copyright (c) 2026 Petr Balvín (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) + }) +} diff --git a/grad/gradient_hygiene_pins_test.go b/grad/gradient_hygiene_pins_test.go new file mode 100644 index 0000000..83cee65 --- /dev/null +++ b/grad/gradient_hygiene_pins_test.go @@ -0,0 +1,230 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/grad/helpers_test.go b/grad/helpers_test.go new file mode 100644 index 0000000..61a0c11 --- /dev/null +++ b/grad/helpers_test.go @@ -0,0 +1,47 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/hessian.go b/grad/hessian.go new file mode 100644 index 0000000..d757147 --- /dev/null +++ b/grad/hessian.go @@ -0,0 +1,178 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/hessian_nil_guard_test.go b/grad/hessian_nil_guard_test.go new file mode 100644 index 0000000..1f259f9 --- /dev/null +++ b/grad/hessian_nil_guard_test.go @@ -0,0 +1,40 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/grad/hessian_test.go b/grad/hessian_test.go new file mode 100644 index 0000000..be1b4d6 --- /dev/null +++ b/grad/hessian_test.go @@ -0,0 +1,193 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/grad/hmc.go b/grad/hmc.go new file mode 100644 index 0000000..2b61ea9 --- /dev/null +++ b/grad/hmc.go @@ -0,0 +1,194 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/hmc_test.go b/grad/hmc_test.go new file mode 100644 index 0000000..57a0a67 --- /dev/null +++ b/grad/hmc_test.go @@ -0,0 +1,217 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/grad/hvp_shape_pin_test.go b/grad/hvp_shape_pin_test.go new file mode 100644 index 0000000..99de132 --- /dev/null +++ b/grad/hvp_shape_pin_test.go @@ -0,0 +1,38 @@ +// Copyright (c) 2026 Petr Balvín (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)) + } + } +} diff --git a/grad/matmul_adjoint_pins_test.go b/grad/matmul_adjoint_pins_test.go new file mode 100644 index 0000000..1178c16 --- /dev/null +++ b/grad/matmul_adjoint_pins_test.go @@ -0,0 +1,158 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/newtoncg.go b/grad/newtoncg.go new file mode 100644 index 0000000..e1353a4 --- /dev/null +++ b/grad/newtoncg.go @@ -0,0 +1,248 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/newtoncg_test.go b/grad/newtoncg_test.go new file mode 100644 index 0000000..f93271f --- /dev/null +++ b/grad/newtoncg_test.go @@ -0,0 +1,197 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/grad/permute.go b/grad/permute.go new file mode 100644 index 0000000..aba50cd --- /dev/null +++ b/grad/permute.go @@ -0,0 +1,42 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/permute_test.go b/grad/permute_test.go new file mode 100644 index 0000000..740e047 --- /dev/null +++ b/grad/permute_test.go @@ -0,0 +1,100 @@ +// Copyright (c) 2026 Petr Balvín (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)) + } + } +} diff --git a/grad/pool.go b/grad/pool.go new file mode 100644 index 0000000..cec9ede --- /dev/null +++ b/grad/pool.go @@ -0,0 +1,352 @@ +// Copyright (c) 2026 Petr Balvín (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) +} diff --git a/grad/pool_bench_test.go b/grad/pool_bench_test.go new file mode 100644 index 0000000..b1d99e4 --- /dev/null +++ b/grad/pool_bench_test.go @@ -0,0 +1,200 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/grad/pooled_sweep_pin_test.go b/grad/pooled_sweep_pin_test.go new file mode 100644 index 0000000..5084f64 --- /dev/null +++ b/grad/pooled_sweep_pin_test.go @@ -0,0 +1,164 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/grad/shape.go b/grad/shape.go new file mode 100644 index 0000000..928c6a9 --- /dev/null +++ b/grad/shape.go @@ -0,0 +1,80 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/grad/spectral.go b/grad/spectral.go new file mode 100644 index 0000000..9b3ce26 --- /dev/null +++ b/grad/spectral.go @@ -0,0 +1,246 @@ +// Copyright (c) 2026 Petr Balvín (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 (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) + } + } +} diff --git a/grad/sweep_test.go b/grad/sweep_test.go new file mode 100644 index 0000000..08b6fd4 --- /dev/null +++ b/grad/sweep_test.go @@ -0,0 +1,286 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/grad/tensor.go b/grad/tensor.go new file mode 100644 index 0000000..d0cdeec --- /dev/null +++ b/grad/tensor.go @@ -0,0 +1,2303 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package grad + +import ( + "fmt" + "math" + "slices" + "strconv" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Automatic differentiation. Tensor wraps an immutable Array +// with reverse-mode autograd: every differentiable operation records a +// node in the computation graph, and Backward propagates gradients from +// the output to every leaf that requires them. This is the training +// substrate on top of the Array surface. +// +// Contract: +// - float, float32 and complex tensors are differentiable; int +// tensors error loudly. Backward needs a real-valued loss: a +// complex loss is an error, and complex leaves receive the +// conjugate-Wirtinger gradient dL/dz̄, the coefficient g of +// dL = 2·Re(g·dz). +// - Gradients carry the same dtype as the data they update; a +// mixed-dtype graph narrows each leaf gradient to the leaf's dtype. +// - Backward accumulates into existing leaf gradients; call ZeroGrad +// before the next backward. +// - The graph is rebuilt on every forward pass; tensors are cheap +// values and never alias the arrays they wrap. + +// Tensor is a differentiable n-dimensional array. +type Tensor struct { + data *core.Array + grad *core.Array + reqGrad bool + node *gradNode +} + +// gradNode is one recorded operation in the reverse-mode graph. The +// operands are held inline, so the tape pays no second allocation for a +// one- or two-element input list, and the backward closures read the +// operand arrays they were handed at forward time: Tensor.data is +// mutable through ReplaceWith, so nothing may be read off an operand at +// backward time. A node carries no sweep state: everything a reverse +// pass needs lives in locals, so two passes over one graph never share +// a cell and never see each other's stamps. +type gradNode struct { + op string + grad func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error + in [2]*Tensor + out *Tensor + arity int8 +} + +// opResult is a differentiable result and its graph node in one +// allocation: a node lives exactly as long as the tensor carrying it, +// so the tape pays one allocation per operation instead of two. +type opResult struct { + t Tensor + n gradNode +} + +// FromArray wraps a as a tensor; requiresGrad marks it as a leaf whose +// gradient Backward fills. +func FromArray(a *core.Array, requiresGrad bool) *Tensor { + return &Tensor{data: a, reqGrad: requiresGrad} +} + +// FromFloat64s builds a float tensor from vals and a shape, mirroring +// FromFloats. +func FromFloat64s(vals []float64, requiresGrad bool, shape ...int) (*Tensor, error) { + a, err := core.FromFloats(vals, shape...) + if err != nil { + return nil, err + } + return FromArray(a, requiresGrad), nil +} + +// Data returns the underlying array. +func (t *Tensor) Data() *core.Array { return t.data } + +// Grad returns the accumulated gradient, or nil before Backward. +func (t *Tensor) Grad() *core.Array { return t.grad } + +// RequiresGrad reports whether the tensor is a trainable leaf. +func (t *Tensor) RequiresGrad() bool { return t.reqGrad } + +// ZeroGrad discards the accumulated gradient. +func (t *Tensor) ZeroGrad() { t.grad = nil } + +// Backward propagates gradients from t back to every leaf that requires +// them, accumulating into the leaves' existing gradients. It seeds the +// output with ones, so a scalar loss is the usual caller. The seed is +// real: a complex output means the "loss" is not a scalar objective, +// and Backward rejects it with an error pointing at Real, Abs or Abs2. +// The reverse sweep recycles its intermediate gradient buffers; the +// arrays committed here escape that recycling, because from Grad on +// they belong to the caller. +func (t *Tensor) Backward() error { + grads, err := t.reverseGradsPooled() + if err != nil { + return err + } + // Commit the accumulated gradients into the leaves, merging with any + // gradient left over from an earlier Backward. The sweep returned + // leaf entries only. + for leaf, s := range grads { + if s.arr == nil || !leaf.reqGrad { + continue + } + if leaf.grad == nil { + leaf.grad = s.arr + continue + } + acc, err := core.Add(leaf.grad, s.arr) + if err != nil { + return err + } + releaseGrad(s.arr, s.sh) + leaf.grad = acc + } + return nil +} + +// reverseGrads runs the reverse pass over t's graph and returns the +// gradient every reached tensor receives from this pass alone. Nothing +// is committed: the pass neither reads nor writes any tensor's +// accumulated Grad, so a caller that differentiates the graph for its +// own purposes (the second-order helpers, whose objective closure may +// touch the caller's trainable tensors) leaves the graph's gradients +// exactly as it found them. +// +// The sweep walks the graph once to fix the topological order and +// again in reverse, folding each node's contributions into its inputs. +// Both orders are the ones a recursive children-first walk produces: +// the fold order decides the rounding of an accumulation, so it is +// part of the result and not an implementation detail. Everything the +// pass needs is a local: two passes over one graph, concurrent +// included, cannot see each other's state. +func (t *Tensor) reverseGrads() (map[*Tensor]*core.Array, error) { + if !t.reqGrad { + return nil, errf("Backward: tensor does not require grad") + } + if isComplexArr(t.data) { + return nil, errf("Backward: the loss must be real-valued; reduce the complex result with Real, Imag, Abs or Abs2 first") + } + ones, err := core.Ones(t.data.Dtype(), t.data.Shape()...) + if err != nil { + return nil, err + } + if t.node == nil { + // A leaf tensor's gradient w.r.t. itself is ones, accumulated + // like any other backward pass. + return map[*Tensor]*core.Array{t: ones}, nil + } + type frame struct { + node *gradNode + next int + } + // Topological order of the graph, children first: an explicit + // post-order walk that marks a node when it is first reached and + // appends it once its operands are done, the order the recursive + // walk produced. + seen := make(map[*Tensor]bool) + order := make([]*gradNode, 0, 64) + stack := make([]frame, 0, 64) + seen[t] = true + stack = append(stack, frame{node: t.node}) + for len(stack) > 0 { + f := &stack[len(stack)-1] + if f.next < int(f.node.arity) { + in := f.node.in[f.next] + f.next++ + if in.node != nil && !seen[in] { + seen[in] = true + stack = append(stack, frame{node: in.node}) + } + continue + } + order = append(order, f.node) + stack = stack[:len(stack)-1] + } + grads := map[*Tensor]*core.Array{t: ones} + var slots [2]gradSlot + for _, node := range slices.Backward(order) { + gOut, ok := grads[node.out] + if !ok { + continue + } + for i := range slots { + slots[i] = gradSlot{} + } + // A nil arena: this path hands the map to callers who read it + // after the sweep, so its buffers are plain allocations the + // collector reclaims, never pooled ones. + if err := node.grad(gradSlot{arr: gOut}, &slots, nil); err != nil { + return nil, err + } + for j := range int(node.arity) { + in := node.in[j] + gj := slots[j].arr + if !in.reqGrad || gj == nil { + continue + } + if gj.Dtype() != in.data.Dtype() { + sl, nerr := narrowGradient(gradSlot{arr: gj}, in.data.Dtype()) + if nerr != nil { + return nil, nerr + } + gj = sl.arr + if gj == nil { + continue + } + } + // Accumulate into the incoming-gradient map: a tensor used + // twice (like z in z·z) receives both contributions. + if prev, ok := grads[in]; ok { + var aerr error + gj, aerr = core.Add(prev, gj) + if aerr != nil { + return nil, aerr + } + } + grads[in] = gj + } + } + return grads, nil +} + +// reverseGradsPooled is Backward's own reverse pass: the same walk, the +// same fold order and the same per-element arithmetic as reverseGrads, +// with the gradient buffers borrowed from the sweep's arena. The +// returned map holds leaf entries only; every intermediate gradient is +// released the moment its producing node has consumed it. That release +// is safe because the closures maintain one invariant: a closure +// returns freshly built buffers in its slots, never the incoming +// gradient itself, never one buffer in two slots, and never a view of +// another live array, so at the producing node's fold nothing but the +// map entry itself can still reach the buffer. A leaf entry is never +// released here: Backward commits it to the leaf, where it escapes the +// pool into caller ownership. An error path commits nothing to any +// tensor: the caller's gradients are exactly as they were. Buffers +// already released to the arena return through its flush, where the +// pool zeroes every one of them at the next borrow; buffers still +// referenced when the error surfaces go to the collector unreleased. +func (t *Tensor) reverseGradsPooled() (map[*Tensor]gradSlot, error) { + if !t.reqGrad { + return nil, errf("Backward: tensor does not require grad") + } + if isComplexArr(t.data) { + return nil, errf("Backward: the loss must be real-valued; reduce the complex result with Real, Imag, Abs or Abs2 first") + } + ar := borrowArena() + seed := gradSlot{sh: t.data.Shape()} + switch dt := t.data.Dtype(); dt { + case core.Float, core.Float32, core.Complex: + seed.arr = ar.borrowGrad(dt, seed.sh) + if dt == core.Complex { + fillGradSlotC(seed, complex(1, 0)) + } else { + fillGradSlot(seed, 1) + } + default: + ones, err := core.Ones(t.data.Dtype(), seed.sh...) + if err != nil { + ar.flush() + return nil, err + } + seed.arr = ones + } + if t.node == nil { + ar.flush() + return map[*Tensor]gradSlot{t: seed}, nil + } + w := borrowTapeWork() + fail := func(err error) (map[*Tensor]gradSlot, error) { + releaseTapeWork(w) + ar.flush() + return nil, err + } + // The same explicit children-first post-order walk reverseGrads + // runs, on the borrowed traversal scratch. + w.seen[t] = true + w.stack = append(w.stack, tapeFrame{node: t.node}) + for len(w.stack) > 0 { + f := &w.stack[len(w.stack)-1] + if f.next < int(f.node.arity) { + in := f.node.in[f.next] + f.next++ + if in.node != nil && !w.seen[in] { + w.seen[in] = true + w.stack = append(w.stack, tapeFrame{node: in.node}) + } + continue + } + w.order = append(w.order, f.node) + w.stack = w.stack[:len(w.stack)-1] + } + grads := map[*Tensor]gradSlot{t: seed} + var slots [2]gradSlot + for _, node := range slices.Backward(w.order) { + gOut, ok := grads[node.out] + if !ok { + continue + } + for i := range slots { + slots[i] = gradSlot{} + } + if err := node.grad(gOut, &slots, ar); err != nil { + return fail(err) + } + // The producing node has consumed its output gradient, and no + // closure returned the buffer itself, so this entry was the + // last live reference. + delete(grads, node.out) + ar.releaseGrad(gOut.arr, gOut.sh) + for j := range int(node.arity) { + in := node.in[j] + s := slots[j] + if s.arr == nil { + continue + } + if !in.reqGrad { + ar.releaseGrad(s.arr, s.sh) + continue + } + gj := s + if gj.arr.Dtype() != in.data.Dtype() { + nj, nerr := narrowGradient(gj, in.data.Dtype()) + if nerr != nil { + return fail(nerr) + } + if nj.arr == gj.arr { + gj = nj + } else { + ar.releaseGrad(gj.arr, gj.sh) + gj = nj + } + } + // Accumulate into the incoming-gradient map: a tensor used + // twice (like z in z·z) receives both contributions, folded + // in the same order the walk has always produced. + if prev, ok := grads[in]; ok { + acc, aerr := addGradSlots(ar, prev, gj) + if aerr != nil { + return fail(aerr) + } + ar.releaseGrad(prev.arr, prev.sh) + ar.releaseGrad(gj.arr, gj.sh) + grads[in] = acc + continue + } + grads[in] = gj + } + } + releaseTapeWork(w) + // Intermediates were released at their producing nodes; the sweep + // below keeps the leaf-only return invariant explicit rather than + // trusted. + for k, s := range grads { + if k.node != nil { + ar.releaseGrad(s.arr, s.sh) + delete(grads, k) + } + } + ar.flush() + return grads, nil +} + +// gradShape returns the shape the incoming gradient carries, or, on +// the legacy sweep path whose slots travel without a shape, a fresh +// copy of a's own. For a shape-preserving op the two are equal, which +// is what Add's and Mul's use of g.sh already rests on. +func gradShape(g gradSlot, a *core.Array) []int { + if g.sh != nil { + return g.sh + } + return a.Shape() +} + +// elemSweepMin is the element count below which the new gradient +// helpers stay on the calling goroutine: below the element-wise spawn +// floor a parallel split cannot pay for its own closure, and the +// helpers' buffers at that size are the tape's small tensors. +const elemSweepMin = 1024 + +// copyGradSlot returns an independent buffer holding g's values, +// shaped sh. The copy is what keeps every gradient entry exclusively +// owned, which in turn is what lets the sweep recycle an entry the +// moment its producing node is done: a passthrough of the incoming +// buffer would leave a second live reference and break that proof. The +// same-dtype dense path writes the copy element for element, the same +// values Astype's clone produced; anything else keeps Astype itself. +func copyGradSlot(ar *gradArena, g gradSlot, sh []int) (gradSlot, error) { + if g.arr == nil { + return gradSlot{}, nil + } + dt := g.arr.Dtype() + if g.arr.Strided() || len(sh) == 0 || len(sh) > 6 { + c, err := core.Astype(g.arr, dt) + if err != nil { + return gradSlot{}, err + } + return gradSlot{arr: c, sh: sh}, nil + } + n := g.arr.Len() + var out *core.Array + switch dt { + case core.Float, core.Float32, core.Complex: + out = ar.borrowGrad(dt, sh) + default: + c, err := core.Astype(g.arr, dt) + if err != nil { + return gradSlot{}, err + } + return gradSlot{arr: c, sh: sh}, nil + } + switch dt { + case core.Float: + copy(out.RawFloats()[:n], g.arr.RawFloats()[:n]) + case core.Float32: + copy(out.RawFloat32s()[:n], g.arr.RawFloat32s()[:n]) + default: + copy(out.RawComplexes()[:n], g.arr.RawComplexes()[:n]) + } + return gradSlot{arr: out, sh: sh}, nil +} + +// negGradSlot returns an independent buffer holding g negated. Each +// dtype multiplies by the negated one in the expression MulI writes, +// so the bits are MulI(g, -1)'s own. +func negGradSlot(ar *gradArena, g gradSlot, sh []int) (gradSlot, error) { + if g.arr == nil { + return gradSlot{}, nil + } + dt := g.arr.Dtype() + if g.arr.Strided() || len(sh) == 0 || len(sh) > 6 { + c := core.MulI(g.arr, -1) + return gradSlot{arr: c, sh: sh}, nil + } + n := g.arr.Len() + var out *core.Array + switch dt { + case core.Float, core.Float32, core.Complex: + out = ar.borrowGrad(dt, sh) + default: + c := core.MulI(g.arr, -1) + return gradSlot{arr: c, sh: sh}, nil + } + switch dt { + case core.Float: + gs, os := g.arr.RawFloats()[:n], out.RawFloats()[:n] + for i := range os { + os[i] = gs[i] * -1 + } + case core.Float32: + gs, os := g.arr.RawFloat32s()[:n], out.RawFloat32s()[:n] + for i := range os { + os[i] = float32(float64(gs[i]) * -1) + } + default: + gs, os := g.arr.RawComplexes()[:n], out.RawComplexes()[:n] + for i := range os { + os[i] = gs[i] * complex(-1, 0) + } + } + return gradSlot{arr: out, sh: sh}, nil +} + +// addGradSlots folds two contributions to one tensor's gradient into a +// fresh buffer. Both carry the tensor's gradient dtype and shape, so +// the dense path writes core.Add's own per-element sum into a borrowed +// buffer; anything else keeps core.Add unchanged. +func addGradSlots(ar *gradArena, p, q gradSlot) (gradSlot, error) { + if p.arr == nil { + return q, nil + } + if q.arr == nil { + return p, nil + } + dt := p.arr.Dtype() + if dt != q.arr.Dtype() || p.arr.Strided() || q.arr.Strided() || + len(p.sh) == 0 || len(p.sh) > 6 { + out, err := core.Add(p.arr, q.arr) + if err != nil { + return gradSlot{}, err + } + return gradSlot{arr: out, sh: p.sh}, nil + } + n := p.arr.Len() + out := ar.borrowGrad(dt, p.sh) + switch dt { + case core.Float: + os, as, bs := out.RawFloats()[:n], p.arr.RawFloats()[:n], q.arr.RawFloats()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = as[i] + bs[i] + } + break + } + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = as[i] + bs[i] + } + }) + case core.Float32: + os, as, bs := out.RawFloat32s()[:n], p.arr.RawFloat32s()[:n], q.arr.RawFloat32s()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = float32(float64(as[i]) + float64(bs[i])) + } + break + } + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = float32(float64(as[i]) + float64(bs[i])) + } + }) + default: + os, as, bs := out.RawComplexes()[:n], p.arr.RawComplexes()[:n], q.arr.RawComplexes()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = as[i] + bs[i] + } + break + } + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = as[i] + bs[i] + } + }) + } + return gradSlot{arr: out, sh: p.sh}, nil +} + +// mulGradSlot writes g⊙b into a borrowed buffer shaped sh: the same +// per-element product core.Mul computes, dtype for dtype. The second +// return is false when the operands leave the dense same-dtype paths, +// and the caller keeps core.Mul. +func mulGradSlot(ar *gradArena, g gradSlot, b *core.Array, sh []int) (gradSlot, bool) { + if g.arr == nil || b == nil || g.arr.Strided() || b.Strided() || + len(sh) == 0 || len(sh) > 6 || g.arr.Len() != b.Len() { + return gradSlot{}, false + } + dt := g.arr.Dtype() + if dt != b.Dtype() { + return gradSlot{}, false + } + n := g.arr.Len() + switch dt { + case core.Float: + out := ar.borrowGrad(dt, sh) + gs, bs, os := g.arr.RawFloats()[:n], b.RawFloats()[:n], out.RawFloats()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = gs[i] * bs[i] + } + } else { + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = gs[i] * bs[i] + } + }) + } + return gradSlot{arr: out, sh: sh}, true + case core.Float32: + out := ar.borrowGrad(dt, sh) + gs, bs, os := g.arr.RawFloat32s()[:n], b.RawFloat32s()[:n], out.RawFloat32s()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = float32(float64(gs[i]) * float64(bs[i])) + } + } else { + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = float32(float64(gs[i]) * float64(bs[i])) + } + }) + } + return gradSlot{arr: out, sh: sh}, true + case core.Complex: + out := ar.borrowGrad(dt, sh) + gs, bs, os := g.arr.RawComplexes()[:n], b.RawComplexes()[:n], out.RawComplexes()[:n] + for i := range os { + os[i] = gs[i] * bs[i] + } + return gradSlot{arr: out, sh: sh}, true + } + return gradSlot{}, false +} + +// mulConjGradSlot writes g·conj(b) into a borrowed buffer shaped sh, +// the complex product rule's adjoint. The conjugate is formed per +// element exactly as conjArray forms it, and the product is the same +// one core.Mul computes on the conjugated operand. +func mulConjGradSlot(ar *gradArena, g gradSlot, b *core.Array, sh []int) (gradSlot, bool) { + if g.arr == nil || b == nil || g.arr.Strided() || b.Strided() || + len(sh) == 0 || len(sh) > 6 || g.arr.Len() != b.Len() || + g.arr.Dtype() != core.Complex || b.Dtype() != core.Complex { + return gradSlot{}, false + } + n := g.arr.Len() + out := ar.borrowGrad(core.Complex, sh) + gs, bs, os := g.arr.RawComplexes()[:n], b.RawComplexes()[:n], out.RawComplexes()[:n] + for i := range os { + z := bs[i] + os[i] = gs[i] * complex(real(z), -imag(z)) + } + return gradSlot{arr: out, sh: sh}, true +} + +// divGradSlot writes g/b into a borrowed buffer shaped sh with the +// per-element quotient core.Div computes. The second return is false +// outside the dense same-dtype real and complex paths. +func divGradSlot(ar *gradArena, g gradSlot, b *core.Array, sh []int) (gradSlot, bool) { + if g.arr == nil || b == nil || g.arr.Strided() || b.Strided() || + len(sh) == 0 || len(sh) > 6 || g.arr.Len() != b.Len() { + return gradSlot{}, false + } + dt := g.arr.Dtype() + if dt != b.Dtype() { + return gradSlot{}, false + } + n := g.arr.Len() + switch dt { + case core.Float: + out := ar.borrowGrad(dt, sh) + gs, bs, os := g.arr.RawFloats()[:n], b.RawFloats()[:n], out.RawFloats()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = gs[i] / bs[i] + } + } else { + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = gs[i] / bs[i] + } + }) + } + return gradSlot{arr: out, sh: sh}, true + case core.Float32: + out := ar.borrowGrad(dt, sh) + gs, bs, os := g.arr.RawFloat32s()[:n], b.RawFloat32s()[:n], out.RawFloat32s()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = float32(float64(gs[i]) / float64(bs[i])) + } + } else { + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = float32(float64(gs[i]) / float64(bs[i])) + } + }) + } + return gradSlot{arr: out, sh: sh}, true + case core.Complex: + out := ar.borrowGrad(dt, sh) + gs, bs, os := g.arr.RawComplexes()[:n], b.RawComplexes()[:n], out.RawComplexes()[:n] + for i := range os { + os[i] = gs[i] / bs[i] + } + return gradSlot{arr: out, sh: sh}, true + } + return gradSlot{}, false +} + +// divGradReal writes the real Div backward, ga = g/b and gb = −g·a/b², +// into two borrowed buffers shaped sh. gb's sub-expressions run in the +// order the staged core chain produced them in (neg = g·(−1), +// num = neg·a, bb = b·b, gb = num/bb), so every element rounds exactly +// as the chain rounded it. The second return is false outside the +// dense float64 and float32 paths. +func divGradReal(ar *gradArena, g gradSlot, a, b *core.Array, sh []int) (gradSlot, gradSlot, bool) { + if g.arr == nil || a == nil || b == nil || g.arr.Strided() || a.Strided() || b.Strided() || + len(sh) == 0 || len(sh) > 6 || g.arr.Len() != a.Len() || g.arr.Len() != b.Len() { + return gradSlot{}, gradSlot{}, false + } + dt := g.arr.Dtype() + if dt != a.Dtype() || dt != b.Dtype() || (dt != core.Float && dt != core.Float32) { + return gradSlot{}, gradSlot{}, false + } + n := g.arr.Len() + ga := ar.borrowGrad(dt, sh) + gb := ar.borrowGrad(dt, sh) + if dt == core.Float { + gs, as, bs := g.arr.RawFloats()[:n], a.RawFloats()[:n], b.RawFloats()[:n] + da, db := ga.RawFloats()[:n], gb.RawFloats()[:n] + if n < elemSweepMin { + for i := range n { + da[i] = gs[i] / bs[i] + neg := gs[i] * -1 + num := neg * as[i] + bb := bs[i] * bs[i] + db[i] = num / bb + } + return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true + } + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + da[i] = gs[i] / bs[i] + neg := gs[i] * -1 + num := neg * as[i] + bb := bs[i] * bs[i] + db[i] = num / bb + } + }) + return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true + } + gs, as, bs := g.arr.RawFloat32s()[:n], a.RawFloat32s()[:n], b.RawFloat32s()[:n] + da, db := ga.RawFloat32s()[:n], gb.RawFloat32s()[:n] + if n < elemSweepMin { + for i := range n { + da[i] = float32(float64(gs[i]) / float64(bs[i])) + neg := float32(float64(gs[i]) * -1) + num := float32(float64(neg) * float64(as[i])) + bb := float32(float64(bs[i]) * float64(bs[i])) + db[i] = float32(float64(num) / float64(bb)) + } + return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true + } + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + da[i] = float32(float64(gs[i]) / float64(bs[i])) + neg := float32(float64(gs[i]) * -1) + num := float32(float64(neg) * float64(as[i])) + bb := float32(float64(bs[i]) * float64(bs[i])) + db[i] = float32(float64(num) / float64(bb)) + } + }) + return gradSlot{arr: ga, sh: sh}, gradSlot{arr: gb, sh: sh}, true +} + +// transposeGradSlot writes the transpose of the rank-2 dense g into a +// borrowed buffer shaped sh: pure data motion, the same elements +// core.Transpose copies, only into recycled storage. The caller shapes +// g as (sh[1], sh[0]), the transpose node's own output shape, so out +// element (i, j) is g element (j, i). The second return is false +// elsewhere and the caller keeps core.Transpose. +func transposeGradSlot(ar *gradArena, g gradSlot, sh []int) (gradSlot, bool) { + if g.arr == nil || g.arr.Strided() || g.arr.NDim() != 2 || len(sh) != 2 { + return gradSlot{}, false + } + dt := g.arr.Dtype() + switch dt { + case core.Float, core.Float32, core.Complex: + default: + return gradSlot{}, false + } + m, n := sh[0], sh[1] + if g.arr.Len() != m*n { + return gradSlot{}, false + } + out := ar.borrowGrad(dt, sh) + switch dt { + case core.Float: + gs, os := g.arr.RawFloats(), out.RawFloats() + for i := range m { + for j := range n { + os[i*n+j] = gs[j*m+i] + } + } + case core.Float32: + gs, os := g.arr.RawFloat32s(), out.RawFloat32s() + for i := range m { + for j := range n { + os[i*n+j] = gs[j*m+i] + } + } + default: + gs, os := g.arr.RawComplexes(), out.RawComplexes() + for i := range m { + for j := range n { + os[i*n+j] = gs[j*m+i] + } + } + } + return gradSlot{arr: out, sh: sh}, true +} + +// zeros allocates a zeroed array, ignoring the constructor's error: +// every shape here comes from existing arrays, so it cannot be invalid. +func zeros(dt core.Dtype, shape []int) *core.Array { + a, _ := core.Zeros(dt, shape...) + return a +} + +// checkFloat rejects non-float tensors for differentiation. +func (t *Tensor) checkFloat(name string) error { + if t.data.Dtype() != core.Float && t.data.Dtype() != core.Float32 { + return errf("autograd: %s needs a float or float32 tensor, got %s", name, t.data.Dtype()) + } + return nil +} + +// unaryResult wraps out as the result of a one-input op on t and, when +// the graph reaches t, records the backward node. The grad closure must +// capture the operand arrays it needs: Tensor.data is mutable through +// ReplaceWith, so nothing may be read off t at backward time. The +// closure writes its answers into dst and never returns the incoming +// gradient buffer itself: the sweep recycles buffers on the strength of +// that invariant. Tensor and node share one allocation. +func (t *Tensor) unaryResult(op string, out *core.Array, + grad func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error) *Tensor { + if !t.reqGrad { + // No graph reaches this result, so it needs no node. + return &Tensor{data: out} + } + res := &opResult{t: Tensor{data: out, reqGrad: true}} + res.n = gradNode{op: op, grad: grad, in: [2]*Tensor{t}, out: &res.t, arity: 1} + res.t.node = &res.n + return &res.t +} + +// binaryResult is unaryResult for a two-input op over t and u. +func binaryResult(op string, t, u *Tensor, out *core.Array, + grad func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error) *Tensor { + if !t.reqGrad && !u.reqGrad { + return &Tensor{data: out} + } + res := &opResult{t: Tensor{data: out, reqGrad: true}} + res.n = gradNode{op: op, grad: grad, in: [2]*Tensor{t, u}, out: &res.t, arity: 2} + res.t.node = &res.n + return &res.t +} + +// binary is the shared plumbing for element-wise differentiable ops. +// The four arithmetic ops all differentiate complex input, so the gate +// here is checkDiff. +func (t *Tensor) binary(name string, u *Tensor, + forward func(a, b *core.Array) (*core.Array, error), + backward func(g gradSlot, a, b *core.Array, dst *[2]gradSlot, ar *gradArena) error, +) (*Tensor, error) { + if err := t.checkDiff(name); err != nil { + return nil, err + } + if err := u.checkDiff(name); err != nil { + return nil, err + } + out, err := forward(t.data, u.data) + if err != nil { + return nil, err + } + at, au := t.data, u.data + return binaryResult(name, t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + return backward(g, at, au, dst, ar) + }), nil +} + +// Add returns t + u. +func (t *Tensor) Add(u *Tensor) (*Tensor, error) { + return t.binary("Add", u, core.Add, func(g gradSlot, _, _ *core.Array, dst *[2]gradSlot, ar *gradArena) error { + // The gradient is the seed for both inputs, and the two leaves + // must not share one buffer: a shared instance would make a + // write through one leaf's Grad corrupt the other's, and the + // commits into leaf.grad would alias. Both answers are + // independent copies with identical bits, which also keeps each + // gradient entry exclusively owned for the sweep's recycling. + sh := g.sh + c0, err := copyGradSlot(ar, g, sh) + if err != nil { + return err + } + c1, err := copyGradSlot(ar, g, sh) + if err != nil { + return err + } + dst[0], dst[1] = c0, c1 + return nil + }) +} + +// Sub returns t - u. +func (t *Tensor) Sub(u *Tensor) (*Tensor, error) { + return t.binary("Sub", u, core.Sub, func(g gradSlot, _, _ *core.Array, dst *[2]gradSlot, ar *gradArena) error { + sh := g.sh + c0, err := copyGradSlot(ar, g, sh) + if err != nil { + return err + } + c1, err := negGradSlot(ar, g, sh) + if err != nil { + return err + } + dst[0], dst[1] = c0, c1 + return nil + }) +} + +// Mul returns the element-wise product t·u. The complex adjoint +// conjugates the other factor: dz = g·w̄, the Wirtinger rule for a +// holomorphic product. +func (t *Tensor) Mul(u *Tensor) (*Tensor, error) { + return t.binary("Mul", u, core.Mul, func(g gradSlot, a, b *core.Array, dst *[2]gradSlot, ar *gradArena) error { + sh := g.sh + if eitherComplex(a, b) { + ga, ok := mulConjGradSlot(ar, g, b, sh) + if !ok { + var err error + if ga.arr, err = core.Mul(g.arr, conjArray(b)); err != nil { + return err + } + ga.sh = sh + } + gb, ok := mulConjGradSlot(ar, g, a, sh) + if !ok { + var err error + if gb.arr, err = core.Mul(g.arr, conjArray(a)); err != nil { + return err + } + gb.sh = sh + } + dst[0], dst[1] = ga, gb + return nil + } + ga, ok := mulGradSlot(ar, g, b, sh) + if !ok { + var err error + if ga.arr, err = core.Mul(g.arr, b); err != nil { + return err + } + ga.sh = sh + } + gb, ok := mulGradSlot(ar, g, a, sh) + if !ok { + var err error + if gb.arr, err = core.Mul(g.arr, a); err != nil { + return err + } + gb.sh = sh + } + dst[0], dst[1] = ga, gb + return nil + }) +} + +// Div returns the element-wise true division t/u. The complex adjoint +// is da = g/conj(b), db = −g·conj(a)/conj(b)². +func (t *Tensor) Div(u *Tensor) (*Tensor, error) { + return t.binary("Div", u, core.Div, func(g gradSlot, a, b *core.Array, dst *[2]gradSlot, ar *gradArena) error { + sh := g.sh + if eitherComplex(a, b) { + cb := conjArray(b) + ga, err := core.Div(g.arr, cb) + if err != nil { + return err + } + cbsq, err := core.Mul(cb, cb) + if err != nil { + return err + } + ca := conjArray(a) + num, err := core.Mul(g.arr, ca) + if err != nil { + return err + } + gb, err := core.Div(core.MulI(num, -1), cbsq) + if err != nil { + return err + } + dst[0] = gradSlot{arr: ga, sh: sh} + dst[1] = gradSlot{arr: gb, sh: sh} + return nil + } + if ga, gb, ok := divGradReal(ar, g, a, b, sh); ok { + dst[0], dst[1] = ga, gb + return nil + } + ga, err := core.Div(g.arr, b) + if err != nil { + return err + } + bb, err := core.Mul(b, b) + if err != nil { + return err + } + neg := core.MulI(g.arr, -1) + num, err := core.Mul(neg, a) + if err != nil { + return err + } + gb, err := core.Div(num, bb) + if err != nil { + return err + } + dst[0] = gradSlot{arr: ga, sh: sh} + dst[1] = gradSlot{arr: gb, sh: sh} + return nil + }) +} + +// MatMul is the differentiable matrix product (2-D×2-D, 2-D×1-D and +// 1-D×2-D, mirroring MatMul2D). It rides the same plumbing as the +// element-wise binary ops; only the forward and backward differ. The +// complex adjoint conjugates and transposes: dA = g·Bᴴ, dB = Aᴴ·g.arr. +func (t *Tensor) MatMul(u *Tensor) (*Tensor, error) { + if err := t.checkDiff("MatMul"); err != nil { + return nil, err + } + if err := u.checkDiff("MatMul"); err != nil { + return nil, err + } + out, err := core.MatMul2D(t.data, u.data) + if err != nil { + return nil, err + } + at, au := t.data, u.data + return binaryResult("MatMul", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + var ga, gb gradSlot + var err error + if eitherComplex(at, au) { + ga, gb, err = matmulGradComplex(ar, g, at, au) + } else { + ga, gb, err = matmulGrad(ar, g, at, au) + } + if err != nil { + return err + } + dst[0], dst[1] = ga, gb + return nil + }), nil +} + +// matmulGradComplex is the Wirtinger backward of MatMul2D: every +// operand enters the adjoint conjugated, so the plain transposes of +// the real backward become conjugate transposes. The transposes and +// reshapes are closure-local scratch, released back to the arena once +// the products are formed. +func matmulGradComplex(ar *gradArena, g gradSlot, a, b *core.Array) (gradSlot, gradSlot, error) { + ca, cb := conjArray(a), conjArray(b) + shA, shB := a.Shape(), b.Shape() + switch { + case a.NDim() == 2 && b.NDim() == 2: + bh, ok := transposeGradSlot(ar, gradSlot{arr: cb}, []int{shB[1], shB[0]}) + if !ok { + bh = gradSlot{arr: core.Transpose(cb), sh: []int{shB[1], shB[0]}} + } + da, err := core.MatMul2D(g.arr, bh.arr) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + ah, ok := transposeGradSlot(ar, gradSlot{arr: ca}, []int{shA[1], shA[0]}) + if !ok { + ah = gradSlot{arr: core.Transpose(ca), sh: []int{shA[1], shA[0]}} + } + db, err := core.MatMul2D(ah.arr, g.arr) + ar.releaseGrad(bh.arr, bh.sh) + ar.releaseGrad(ah.arr, ah.sh) + // conjArray answers a real operand unchanged, so the conjugate + // is released only when it is a buffer this call built. + if ca != a { + ar.releaseGrad(ca, shA) + } + if cb != b { + ar.releaseGrad(cb, shB) + } + if err != nil { + return gradSlot{}, gradSlot{}, err + } + return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil + case a.NDim() == 2 && b.NDim() == 1: + // da[i, j] = g[i]·b̄[j], the outer product against the + // conjugated vector. + gr, err := core.Reshape(g.arr, shA[0], 1) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + br, err := core.Reshape(cb, 1, b.Len()) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + da, err := core.MatMul2D(gr, br) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + ah, ok := transposeGradSlot(ar, gradSlot{arr: ca}, []int{shA[1], shA[0]}) + if !ok { + ah = gradSlot{arr: core.Transpose(ca), sh: []int{shA[1], shA[0]}} + } + db, err := core.MatMul2D(ah.arr, g.arr) + ar.releaseGrad(ah.arr, ah.sh) + if ca != a { + ar.releaseGrad(ca, shA) + } + if cb != b { + ar.releaseGrad(cb, shB) + } + if err != nil { + return gradSlot{}, gradSlot{}, err + } + return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil + case a.NDim() == 1 && b.NDim() == 2: + bh, ok := transposeGradSlot(ar, gradSlot{arr: cb}, []int{shB[1], shB[0]}) + if !ok { + bh = gradSlot{arr: core.Transpose(cb), sh: []int{shB[1], shB[0]}} + } + da, err := core.MatMul2D(g.arr, bh.arr) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + ar2, err := core.Reshape(ca, a.Len(), 1) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + gr, err := core.Reshape(g.arr, 1, g.arr.Len()) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + db, err := core.MatMul2D(ar2, gr) + ar.releaseGrad(bh.arr, bh.sh) + if ca != a { + ar.releaseGrad(ca, shA) + } + if cb != b { + ar.releaseGrad(cb, shB) + } + if err != nil { + return gradSlot{}, gradSlot{}, err + } + return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil + } + return gradSlot{}, gradSlot{}, errf("autograd MatMul: unsupported shapes %s and %s", prettyShape(a.Shape()), prettyShape(b.Shape())) +} + +// matmulGrad is the backward of MatMul2D for every supported shape +// combination. The transposes of the operands are closure-local +// scratch: a dense rank-2 operand transposes into an arena buffer, and +// whichever route produced it, the buffer returns to the arena once +// the two products are formed. +func matmulGrad(ar *gradArena, g gradSlot, a, b *core.Array) (gradSlot, gradSlot, error) { + shA, shB := a.Shape(), b.Shape() + switch { + case a.NDim() == 2 && b.NDim() == 2: + bt, ok := transposeGradSlot(ar, gradSlot{arr: b}, []int{shB[1], shB[0]}) + if !ok { + bt = gradSlot{arr: core.Transpose(b), sh: []int{shB[1], shB[0]}} + } + da, err := core.MatMul2D(g.arr, bt.arr) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + at, ok := transposeGradSlot(ar, gradSlot{arr: a}, []int{shA[1], shA[0]}) + if !ok { + at = gradSlot{arr: core.Transpose(a), sh: []int{shA[1], shA[0]}} + } + db, err := core.MatMul2D(at.arr, g.arr) + ar.releaseGrad(bt.arr, bt.sh) + ar.releaseGrad(at.arr, at.sh) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil + case a.NDim() == 2 && b.NDim() == 1: + // da[i, j] = g[i]·b[j], the outer product as a 1×k by n×1 + // matrix product. + gr, err := core.Reshape(g.arr, shA[0], 1) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + br, err := core.Reshape(b, 1, b.Len()) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + da, err := core.MatMul2D(gr, br) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + at, ok := transposeGradSlot(ar, gradSlot{arr: a}, []int{shA[1], shA[0]}) + if !ok { + at = gradSlot{arr: core.Transpose(a), sh: []int{shA[1], shA[0]}} + } + db, err := core.MatMul2D(at.arr, g.arr) + ar.releaseGrad(at.arr, at.sh) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil + case a.NDim() == 1 && b.NDim() == 2: + bt, ok := transposeGradSlot(ar, gradSlot{arr: b}, []int{shB[1], shB[0]}) + if !ok { + bt = gradSlot{arr: core.Transpose(b), sh: []int{shB[1], shB[0]}} + } + da, err := core.MatMul2D(g.arr, bt.arr) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + ar_, err := core.Reshape(a, a.Len(), 1) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + gr, err := core.Reshape(g.arr, 1, g.arr.Len()) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + db, err := core.MatMul2D(ar_, gr) + ar.releaseGrad(bt.arr, bt.sh) + if err != nil { + return gradSlot{}, gradSlot{}, err + } + return gradSlot{arr: da, sh: shA}, gradSlot{arr: db, sh: shB}, nil + } + return gradSlot{}, gradSlot{}, errf("autograd MatMul: unsupported shapes %s and %s", prettyShape(a.Shape()), prettyShape(b.Shape())) +} + +// Sum reduces the tensor to a single-element tensor holding the sum of +// all elements. A complex tensor sums in complex. +func (t *Tensor) Sum() (*Tensor, error) { + if err := t.checkDiff("Sum"); err != nil { + return nil, err + } + a := t.data + if isComplexArr(a) { + out := scalarComplex(core.Sum(a).Complex()) + return t.unaryResult("Sum", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := a.Shape() + s := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} + fillGradSlotC(s, g.arr.ComplexAt(0)) + dst[0] = s + return nil + }), nil + } + out := scalarLike(t.data, core.Sum(t.data).Float()) + return t.unaryResult("Sum", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := a.Shape() + s := gradSlot{arr: ar.borrowGrad(a.Dtype(), sh), sh: sh} + fillGradSlot(s, g.arr.FloatAt(0)) + dst[0] = s + return nil + }), nil +} + +// Mean reduces the tensor to a single-element tensor holding the mean +// of all elements. A complex tensor means in complex (the core Mean +// is real-only, so the complex path divides the sum directly). +func (t *Tensor) Mean() (*Tensor, error) { + if err := t.checkDiff("Mean"); err != nil { + return nil, err + } + a := t.data + if isComplexArr(a) { + if a.Len() == 0 { + // The real branch's core.Mean refuses an empty reduction; + // this branch divides by the element count, so an empty + // tensor would answer 0/0 = NaN instead. + return nil, errf("Mean: an empty array has no mean") + } + s := core.Sum(a).Complex() + n := complex(float64(a.Len()), 0) + out := scalarComplex(s / n) + return t.unaryResult("Mean", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := a.Shape() + da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} + fillGradSlotC(da, g.arr.ComplexAt(0)/n) + dst[0] = da + return nil + }), nil + } + m, err := core.Mean(t.data) + if err != nil { + return nil, err + } + out := scalarLike(t.data, m) + n := float64(a.Len()) + return t.unaryResult("Mean", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := a.Shape() + da := gradSlot{arr: ar.borrowGrad(a.Dtype(), sh), sh: sh} + fillGradSlot(da, g.arr.FloatAt(0)/n) + dst[0] = da + return nil + }), nil +} + +// Exp returns e raised to each element. Complex input differentiates +// too: exp is holomorphic, so the Wirtinger adjoint multiplies by the +// conjugated output. +func (t *Tensor) Exp() (*Tensor, error) { + if err := t.checkDiff("Exp"); err != nil { + return nil, err + } + if isComplexArr(t.data) { + a := t.data + out := zeros(core.Complex, a.Shape()) + cs := out.RawComplexes() + if a.Strided() { + for i := range cs { + z := a.ComplexAt(i) + cs[i] = complex(math.Exp(real(z))*math.Cos(imag(z)), math.Exp(real(z))*math.Sin(imag(z))) + } + } else { + as := a.RawComplexes() + for i := range cs { + z := as[i] + cs[i] = complex(math.Exp(real(z))*math.Cos(imag(z)), math.Exp(real(z))*math.Sin(imag(z))) + } + } + return t.unaryResult("Exp", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} + ds := da.arr.RawComplexes()[:da.arr.Len()] + if g.arr.Strided() { + for i := range ds { + ds[i] = g.arr.ComplexAt(i) * conj(cs[i]) + } + dst[0] = da + return nil + } + gs := g.arr.RawComplexes() + for i := range ds { + ds[i] = gs[i] * conj(cs[i]) + } + dst[0] = da + return nil + }), nil + } + out, err := core.Exp(t.data) + if err != nil { + return nil, err + } + a := t.data + return t.unaryResult("Exp", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + da, ok := mulGradSlot(ar, g, out, sh) + if !ok { + arr, nerr := core.Mul(g.arr, out) + if nerr != nil { + return nerr + } + da = gradSlot{arr: arr, sh: sh} + } + dst[0] = da + return nil + }), nil +} + +// Log returns the natural logarithm of each element. The domain is the +// caller's: a non-positive input is not refused, and the NaN or ±Inf +// the forward produces flows into the gradient without a report, so +// shift the input inside its domain before it reaches the graph. +func (t *Tensor) Log() (*Tensor, error) { + if err := t.checkFloat("Log"); err != nil { + return nil, err + } + out, err := core.Log(t.data) + if err != nil { + return nil, err + } + a := t.data + return t.unaryResult("Log", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + da, ok := divGradSlot(ar, g, a, sh) + if !ok { + arr, nerr := core.Div(g.arr, a) + if nerr != nil { + return nerr + } + da = gradSlot{arr: arr, sh: sh} + } + dst[0] = da + return nil + }), nil +} + +// Sigmoid returns the logistic sigmoid of each element. +func (t *Tensor) Sigmoid() (*Tensor, error) { + if err := t.checkFloat("Sigmoid"); err != nil { + return nil, err + } + out, err := core.Sigmoid(t.data) + if err != nil { + return nil, err + } + return t.unaryResult("Sigmoid", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + // dz = g·σ·(1−σ) in one pass; the fallback keeps the staged + // core chain for anything the raw paths cannot see through. + sh := gradShape(g, out) + if g.arr.Dtype() != out.Dtype() || g.arr.Strided() || out.Strided() { + comp, err := core.Sub(fillLike(out, 1), out) + if err != nil { + return err + } + so, err := core.Mul(out, comp) + if err != nil { + return err + } + da, err := core.Mul(g.arr, so) + if err != nil { + return err + } + dst[0] = gradSlot{arr: da, sh: sh} + return nil + } + da := gradSlot{arr: ar.borrowGrad(out.Dtype(), sh), sh: sh} + if out.Dtype() == core.Float32 { + gs, os := g.arr.RawFloat32s(), out.RawFloat32s() + ds := da.arr.RawFloat32s()[:da.arr.Len()] + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + cp := float32(1 - float64(os[i])) + so := float32(float64(os[i]) * float64(cp)) + ds[i] = float32(float64(gs[i]) * float64(so)) + } + }) + dst[0] = da + return nil + } + gs, os := g.arr.RawFloats(), out.RawFloats() + ds := da.arr.RawFloats()[:da.arr.Len()] + if len(ds) < elemSweepMin { + for i := range ds { + ds[i] = gs[i] * (os[i] * (1 - os[i])) + } + } else { + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + ds[i] = gs[i] * (os[i] * (1 - os[i])) + } + }) + } + dst[0] = da + return nil + }), nil +} + +// Tanh returns the hyperbolic tangent of each element. +func (t *Tensor) Tanh() (*Tensor, error) { + if err := t.checkFloat("Tanh"); err != nil { + return nil, err + } + out, err := core.Tanh(t.data) + if err != nil { + return nil, err + } + return t.unaryResult("Tanh", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + // dz = g·(1−tanh²) in one pass; the float32 branch rounds at the + // same three points the Mul/Sub chain it replaces did, so the + // payload is unchanged element for element. + sh := gradShape(g, out) + if g.arr.Dtype() != out.Dtype() || g.arr.Strided() || out.Strided() { + sq, err := core.Mul(out, out) + if err != nil { + return err + } + comp, err := core.Sub(fillLike(out, 1), sq) + if err != nil { + return err + } + da, err := core.Mul(g.arr, comp) + if err != nil { + return err + } + dst[0] = gradSlot{arr: da, sh: sh} + return nil + } + da := gradSlot{arr: ar.borrowGrad(out.Dtype(), sh), sh: sh} + if out.Dtype() == core.Float32 { + gs, os := g.arr.RawFloat32s(), out.RawFloat32s() + ds := da.arr.RawFloat32s()[:da.arr.Len()] + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + sq := float32(float64(os[i]) * float64(os[i])) + cp := float32(1 - float64(sq)) + ds[i] = float32(float64(gs[i]) * float64(cp)) + } + }) + dst[0] = da + return nil + } + gs, os := g.arr.RawFloats(), out.RawFloats() + ds := da.arr.RawFloats()[:da.arr.Len()] + if len(ds) < elemSweepMin { + for i := range ds { + ds[i] = gs[i] * (1 - os[i]*os[i]) + } + } else { + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + ds[i] = gs[i] * (1 - os[i]*os[i]) + } + }) + } + dst[0] = da + return nil + }), nil +} + +// Neg returns -t. +func (t *Tensor) Neg() (*Tensor, error) { + if err := t.checkDiff("Neg"); err != nil { + return nil, err + } + a := t.data + return t.unaryResult("Neg", core.MulI(t.data, -1), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + da, err := negGradSlot(ar, g, sh) + if err != nil { + return err + } + dst[0] = da + return nil + }), nil +} + +// Transpose reverses the dimensions; the gradient transposes back. +func (t *Tensor) Transpose() (*Tensor, error) { + if err := t.checkDiff("Transpose"); err != nil { + return nil, err + } + a := t.data + return t.unaryResult("Transpose", core.Transpose(t.data), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := a.Shape() + da, ok := transposeGradSlot(ar, g, sh) + if !ok { + da = gradSlot{arr: core.Transpose(g.arr), sh: sh} + } + dst[0] = da + return nil + }), nil +} + +// Squeeze removes size-1 dimensions; the gradient unsqueezes back. +func (t *Tensor) Squeeze(dim int) (*Tensor, error) { + if err := t.checkDiff("Squeeze"); err != nil { + return nil, err + } + out, err := core.Squeeze(t.data, dim) + if err != nil { + return nil, err + } + orig := t.data.Shape() + return t.unaryResult("Squeeze", 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 +} + +// Unsqueeze inserts a size-1 dimension; the gradient squeezes back. +func (t *Tensor) Unsqueeze(dim int) (*Tensor, error) { + if err := t.checkDiff("Unsqueeze"); err != nil { + return nil, err + } + out, err := core.Unsqueeze(t.data, dim) + if err != nil { + return nil, err + } + orig := t.data.Shape() + return t.unaryResult("Unsqueeze", 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 +} + +// Clip clamps each element into [lo, hi]; the gradient passes through +// only where the value was inside the range. +func (t *Tensor) Clip(lo, hi float64) (*Tensor, error) { + if err := t.checkFloat("Clip"); err != nil { + return nil, err + } + out, err := core.ClipF(t.data, lo, hi) + if err != nil { + return nil, err + } + a := t.data + return t.unaryResult("Clip", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + // d/dx clip(x) = 1 where lo ≤ x ≤ hi, 0 elsewhere. The + // comparisons answer bool masks, so their composition is the + // logic conjunction; And keeps the mask a bool payload. + sh := gradShape(g, a) + ge, err := core.GeF(a, lo) + if err != nil { + return err + } + le, err := core.LeF(a, hi) + if err != nil { + return err + } + mask, err := core.And(ge, le) + if err != nil { + return err + } + // Mul reads the bool mask through its exact 0/1 widening: g·1 + // is g bit for bit and g·0 is zero, the values the int mask's + // product used to write. + da, err := core.Mul(g.arr, mask) + if err != nil { + return err + } + dst[0] = gradSlot{arr: da, sh: sh} + return nil + }), nil +} + +// scalarLike builds a 1-element array of a's dtype holding v. +func scalarLike(a *core.Array, v float64) *core.Array { + out := zeros(a.Dtype(), []int{1}) + out.SetFloatAt(0, v) + return out +} + +// errf formats an error message with the package prefix. +func errf(format string, args ...any) error { + return fmt.Errorf("tensor: "+format, args...) +} + +// fillLike returns an array shaped like a with every element set to v. +// One pass over a fresh payload, unlike the OnesLike-then-MulF walk it +// replaces: ones·v multiplies to exactly v per element (1·v is the +// IEEE identity), so the bits are unchanged. +func fillLike(a *core.Array, v float64) *core.Array { + return fillConst(a.Dtype(), a.Shape(), v) +} + +// fillConst returns a fresh array of dt and shape with every element +// set to v, written straight into the raw payload. An int dtype +// promotes to float, mirroring the scalar-multiply semantics the +// accessor path had. The sweep stays on the calling goroutine: it is a +// plain store stream over one buffer, so a split buys no bandwidth and +// measures slower than the serial store. +func fillConst(dt core.Dtype, shape []int, v float64) *core.Array { + if dt == core.Int { + dt = core.Float + } + out := zeros(dt, shape) + switch dt { + case core.Float: + fs := out.RawFloats() + for i := range fs { + fs[i] = v + } + case core.Float32: + fs := out.RawFloat32s() + fv := float32(v) + for i := range fs { + fs[i] = fv + } + default: + cs := out.RawComplexes() + z := complex(v, 0) + for i := range cs { + cs[i] = z + } + } + return out +} + +// elemPass runs fn over n elements under the elementwise spawn policy: +// ranges too small to amortise a worker stay whole on the calling +// goroutine, large ones split into disjoint chunks. The split cannot +// move a bit: every element is computed from its own operands alone. +func elemPass(n int, fn func(s, e int)) { engine.ParallelMin(n, 1024, fn) } + +// mathElemMinPerWorker is the per-worker chunk floor for the sweeps +// whose per-element cost is a math.Pow call: tens of cycles each, so a +// spawned worker amortises its own start-up at a far smaller chunk +// than the arithmetic maps' floor, which holds every sweep below +// roughly 32k elements on the calling goroutine. A divide is cheap +// enough per element that the arithmetic floor stays the right one for +// it: measured at 20k elements, moving the square-root backward to +// this floor turned a 63µs serial sweep into a 94µs parallel one, +// while the two power sweeps halved. +const mathElemMinPerWorker = 256 + +// mathElemPass is elemPass for the costly per-element sweeps: same +// split, same arithmetic, only the spawn floor differs. +func mathElemPass(n int, fn func(s, e int)) { engine.ParallelMin(n, mathElemMinPerWorker, fn) } + +// l2LinesPerWorker is the per-worker line floor for the L2 norm +// backward: one line costs two passes over the reduced axis, so the +// element-wise floor of about a thousand elements per worker is that +// many line slots, and never below one line. +func l2LinesPerWorker(size int) int { + if size < 1 { + return 1 + } + return max(1, 1024/(2*size)) +} + +// flatFloats copies a's elements into a fresh float64 slice, reading +// the contiguous payload directly when possible and falling back to +// the widening accessor for views and other dtypes. +func flatFloats(a *core.Array) []float64 { + out := make([]float64, a.Len()) + if !a.Strided() && a.Dtype() == core.Float { + copy(out, a.RawFloats()) + return out + } + for i := range out { + out[i] = a.FloatAt(i) + } + return out +} + +// prettyShape renders a shape for diagnostics as "(2, 3)". +func prettyShape(shape []int) string { + parts := make([]string, len(shape)) + for i, d := range shape { + parts[i] = strconv.Itoa(d) + } + return "(" + strings.Join(parts, ", ") + ")" +} + +// ReplaceWith swaps the underlying data array, the one mutable point +// the optimisers rely on for parameter updates. +func (t *Tensor) ReplaceWith(a *core.Array) { t.data = a } + +// SumAxis sums the elements along dim, dropping it; the backward +// broadcasts the incoming gradient back over the reduced dimension. +func (t *Tensor) SumAxis(dim int) (*Tensor, error) { + if err := t.checkDiff("SumAxis"); err != nil { + return nil, err + } + out, err := core.SumAxis(t.data, dim) + if err != nil { + return nil, err + } + shape := append([]int{}, t.data.Shape()...) + return t.unaryResult("SumAxis", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + expanded, err := expandToShape(g.arr, dim, shape) + if err != nil { + return err + } + dst[0] = gradSlot{arr: expanded, sh: shape} + return nil + }), nil +} + +// MeanAxis averages along dim, dropping it; the backward is the sum +// backward scaled by 1/size(dim). +func (t *Tensor) MeanAxis(dim int) (*Tensor, error) { + if err := t.checkFloat("MeanAxis"); err != nil { + return nil, err + } + out, err := core.MeanAxis(t.data, dim) + if err != nil { + return nil, err + } + shape := append([]int{}, t.data.Shape()...) + n := float64(shape[dim]) + return t.unaryResult("MeanAxis", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + // Scale the incoming gradient before the broadcast: the division + // hits every source element once either way, and the broadcast + // copies the quotient exactly. The elements are independent, so + // the split cannot move a bit, and a divide is cheap enough for + // the arithmetic maps' spawn floor. + gsh := g.sh + scaled := gradSlot{arr: ar.borrowGrad(g.arr.Dtype(), gsh), sh: gsh} + if !g.arr.Strided() && g.arr.Dtype() == core.Float32 { + gs := g.arr.RawFloat32s() + ss := scaled.arr.RawFloat32s()[:scaled.arr.Len()] + elemPass(len(ss), func(s, e int) { + for i := s; i < e; i++ { + ss[i] = float32(float64(gs[i]) / n) + } + }) + } else if !g.arr.Strided() && g.arr.Dtype() == core.Float { + gs := g.arr.RawFloats() + ss := scaled.arr.RawFloats()[:scaled.arr.Len()] + elemPass(len(ss), func(s, e int) { + for i := s; i < e; i++ { + ss[i] = gs[i] / n + } + }) + } else { + for i := range g.arr.Len() { + scaled.arr.SetFloatAt(i, g.arr.FloatAt(i)/n) + } + } + expanded, err := expandToShape(scaled.arr, dim, shape) + ar.releaseGrad(scaled.arr, gsh) + if err != nil { + return err + } + dst[0] = gradSlot{arr: expanded, sh: shape} + return nil + }), nil +} + +// L2NormAxis computes the L2 norm along dim, dropping it. The backward +// divides each input element by the line norm: dx = g·x/‖x‖. +func (t *Tensor) L2NormAxis(dim int) (*Tensor, error) { + if err := t.checkFloat("L2NormAxis"); err != nil { + return nil, err + } + out, err := core.Norm(t.data, 2, dim, false) + if err != nil { + return nil, err + } + x := t.data + size := x.Shape()[dim] + stride := 1 + for k := dim + 1; k < x.NDim(); k++ { + stride *= x.Shape()[k] + } + // An empty dimension (or an empty trailing axis) empties the whole + // tensor; the gradient of an empty tensor is empty, and the line + // count below must not divide by a zero stride. + perLine := size * stride + lines := 0 + if perLine > 0 { + lines = x.Len() / perLine + } + const tiny = 1e-12 + xShape := x.Shape() + return t.unaryResult("L2NormAxis", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + dx := gradSlot{arr: ar.borrowGrad(core.Float, xShape), sh: xShape} + // The gradient arrives as float64 (the norm's dtype). The raw + // sweeps below widen a float32 input exactly where FloatAt did, + // and the fallback widens a strided operand once into a linear + // buffer instead of calling the accessor per element. + fast64 := !x.Strided() && !g.arr.Strided() && + x.Dtype() == core.Float && g.arr.Dtype() == core.Float + fast32 := !x.Strided() && !g.arr.Strided() && + x.Dtype() == core.Float32 && g.arr.Dtype() == core.Float + linesPerWorker := l2LinesPerWorker(size) + switch { + case fast64: + gs, xs, ds := g.arr.RawFloats(), x.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] + engine.ParallelMin(lines*stride, linesPerWorker, func(s, e int) { + for key := s; key < e; key++ { + line := key / stride + off := key % stride + var sq float64 + for k := range size { + v := xs[line*size*stride+k*stride+off] + sq += v * v + } + denom := math.Sqrt(sq) + if denom < tiny { + denom = tiny + } + gv := gs[key] + for k := range size { + pos := line*size*stride + k*stride + off + ds[pos] = gv * xs[pos] / denom + } + } + }) + case fast32: + gs, xs32, ds := g.arr.RawFloats(), x.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] + engine.ParallelMin(lines*stride, linesPerWorker, func(s, e int) { + for key := s; key < e; key++ { + line := key / stride + off := key % stride + var sq float64 + for k := range size { + v := float64(xs32[line*size*stride+k*stride+off]) + sq += v * v + } + denom := math.Sqrt(sq) + if denom < tiny { + denom = tiny + } + gv := gs[key] + for k := range size { + pos := line*size*stride + k*stride + off + ds[pos] = gv * float64(xs32[pos]) / denom + } + } + }) + default: + xs, gs, ds := flatFloats(x), flatFloats(g.arr), dx.arr.RawFloats()[:dx.arr.Len()] + engine.ParallelMin(lines*stride, linesPerWorker, func(s, e int) { + for key := s; key < e; key++ { + line := key / stride + off := key % stride + var sq float64 + for k := range size { + v := xs[line*size*stride+k*stride+off] + sq += v * v + } + denom := math.Sqrt(sq) + if denom < tiny { + denom = tiny + } + gv := gs[key] + for k := range size { + pos := line*size*stride + k*stride + off + ds[pos] = gv * xs[pos] / denom + } + } + }) + } + dst[0] = dx + return nil + }), nil +} + +// expandToShape reinserts a dropped axis at dim (size 1) and broadcasts +// g to shape, the inverse of an axis reduction. The reshape and the +// broadcast validate their shapes, so a mismatch travels back to the +// backward pass as an error instead of a nil array. +func expandToShape(g *core.Array, dim int, shape []int) (*core.Array, error) { + gShape := g.Shape() + newShape := make([]int, 0, len(shape)) + newShape = append(newShape, gShape[:dim]...) + newShape = append(newShape, 1) + newShape = append(newShape, gShape[dim:]...) + r, err := core.Reshape(g, newShape...) + if err != nil { + return nil, err + } + return core.BroadcastTo(r, shape...) +} + +// BroadcastTo expands the tensor to shape under broadcasting rules +// (size-1 dimensions replicate). The backward sums the incoming +// gradient over every replicated dimension, collapsing it back to the +// original shape. Only float, float32 and complex tensors carry a +// graph, so an int tensor is refused; the shape is validated first, so +// an impossible target keeps reporting the shape mismatch. +func (t *Tensor) BroadcastTo(shape ...int) (*Tensor, error) { + out, err := core.BroadcastTo(t.data, shape...) + if err != nil { + return nil, err + } + if err := t.checkDiff("BroadcastTo"); err != nil { + return nil, err + } + orig := t.data.Shape() + return t.unaryResult("BroadcastTo", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + red, err := sumToShape(g.arr, orig) + if err != nil { + return err + } + dst[0] = gradSlot{arr: red, sh: orig} + return nil + }), nil +} + +// sumToShape reduces g back to target by summing across every expanded +// dimension: leading extra axes are collapsed one by one, then any dim +// whose size grew (target == 1) is summed at its own position. A rank-1 +// target of size 1 has no summable axis left, so the whole gradient +// collapses to a single sum. +func sumToShape(g *core.Array, target []int) (*core.Array, error) { + for g.NDim() > len(target) { + g2, err := core.SumAxis(g, 0) + if err != nil { + return nil, err + } + g = g2 + } + for i := range len(target) { + if g.Shape()[i] == target[i] { + continue + } + if g.NDim() == 1 { + // The only dimension left: under the broadcasting rules a + // mismatch here means target[0] == 1, so the full sum is + // the correct reduction. + if isComplexArr(g) { + return scalarComplex(core.Sum(g).Complex()), nil + } + return scalarLike(g, core.Sum(g).Float()), nil + } + g2, err := core.SumAxis(g, i) + if err != nil { + return nil, err + } + ns := make([]int, 0, g.NDim()) + ns = append(ns, g.Shape()[:i]...) + ns = append(ns, 1) + ns = append(ns, g.Shape()[i+1:]...) + if g, err = core.Reshape(g2, ns...); err != nil { + return nil, err + } + } + return g, nil +} + +// Pow raises each element to an integer exponent n ≥ 0 on the graph; +// backward multiplies by n·x^(n−1). The exponent-0 backward is the +// zero gradient everywhere, including at x = 0, where n·x^(n−1) would +// evaluate 0·∞ and come out NaN. A complex tensor conjugates the +// derivative, per the Wirtinger convention. +func (t *Tensor) Pow(n int64) (*Tensor, error) { + if n < 0 { + return nil, errf("Pow: negative exponent has no general real gradient") + } + if err := t.checkDiff("Pow"); err != nil { + return nil, err + } + if isComplexArr(t.data) { + a := t.data + out := powComplexForward(a, n) + return t.unaryResult("Pow", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + dst[0] = gradSlot{arr: powComplexGrad(ar, g, a, n, sh), sh: sh} + return nil + }), nil + } + out := zeros(t.Data().Dtype(), t.Data().Shape()) + a := t.Data() + exp := float64(n) + switch { + case a.Strided(): + // The destination is fresh and dense, so its payload takes the + // power directly while the strided operand keeps the accessor + // read that rebases the index. + if out.Dtype() == core.Float32 { + os := out.RawFloat32s() + engine.Parallel(len(os), func(s, e int) { + for i := s; i < e; i++ { + os[i] = float32(math.Pow(a.FloatAt(i), exp)) + } + }) + } else { + os := out.RawFloats() + engine.Parallel(len(os), func(s, e int) { + for i := s; i < e; i++ { + os[i] = math.Pow(a.FloatAt(i), exp) + } + }) + } + case a.Dtype() == core.Float32: + as, os := a.RawFloat32s(), out.RawFloat32s() + mathElemPass(len(os), func(s, e int) { + for i := s; i < e; i++ { + os[i] = float32(math.Pow(float64(as[i]), exp)) + } + }) + default: + as, os := a.RawFloats(), out.RawFloats() + mathElemPass(len(os), func(s, e int) { + for i := s; i < e; i++ { + os[i] = math.Pow(as[i], exp) + } + }) + } + return t.unaryResult("Pow", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + dx := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh} + if exp == 0 { + // d/dx x⁰ = 0, including at x = 0. + dst[0] = dx + return nil + } + // The backward keeps the gradient sweep on raw payloads; dx + // stays float64 whatever the operand width, as the accessor + // path's zeros(core.Float) always did. + if !g.arr.Strided() && g.arr.Dtype() == core.Float32 && !a.Strided() && a.Dtype() == core.Float32 { + gs, as, ds := g.arr.RawFloat32s(), a.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] + mathElemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + ds[i] = float64(gs[i]) * exp * math.Pow(float64(as[i]), exp-1) + } + }) + dst[0] = dx + return nil + } + if !g.arr.Strided() && g.arr.Dtype() == core.Float && !a.Strided() && a.Dtype() == core.Float { + gs, as, ds := g.arr.RawFloats(), a.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] + mathElemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + ds[i] = gs[i] * exp * math.Pow(as[i], exp-1) + } + }) + dst[0] = dx + return nil + } + // Mixed widths or a strided operand: the accessors convert and + // rebase, the destination is dx's own dense float64 payload. + ds := dx.arr.RawFloats()[:dx.arr.Len()] + engine.Parallel(a.Len(), func(s, e int) { + for i := s; i < e; i++ { + x := a.FloatAt(i) + ds[i] = g.arr.FloatAt(i) * exp * math.Pow(x, exp-1) + } + }) + dst[0] = dx + return nil + }), nil +} + +// Abs returns |x|; the real subgradient is sign(x) (0 at zero), the +// complex one z/(2|z|). +func (t *Tensor) Abs() (*Tensor, error) { + if err := t.checkDiff("Abs"); err != nil { + return nil, err + } + if isComplexArr(t.data) { + return t.absComplex() + } + a := t.data + return t.unaryResult("Abs", core.Abs(a), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + dx := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh} + // dx = g·sign(x) on raw payloads; the gradient is float64 here + // whatever the operand width, so the dtype branch sits outside + // the loop. + if !g.arr.Strided() && !a.Strided() && g.arr.Dtype() == core.Float && a.Dtype() == core.Float { + gs, as, ds := g.arr.RawFloats(), a.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + dv := 0.0 + if v := as[i]; v > 0 { + dv = 1 + } else if v < 0 { + dv = -1 + } + ds[i] = gs[i] * dv + } + }) + dst[0] = dx + return nil + } + if !g.arr.Strided() && !a.Strided() && g.arr.Dtype() == core.Float32 && a.Dtype() == core.Float32 { + gs, as, ds := g.arr.RawFloat32s(), a.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + dv := 0.0 + if v := as[i]; v > 0 { + dv = 1 + } else if v < 0 { + dv = -1 + } + ds[i] = float64(gs[i]) * dv + } + }) + dst[0] = dx + return nil + } + ds := dx.arr.RawFloats()[:dx.arr.Len()] + engine.Parallel(a.Len(), func(s, e int) { + for i := s; i < e; i++ { + v := a.FloatAt(i) + dv := 0.0 + if v > 0 { + dv = 1 + } else if v < 0 { + dv = -1 + } + ds[i] = g.arr.FloatAt(i) * dv + } + }) + dst[0] = dx + return nil + }), nil +} + +// Sqrt returns √x; backward is 1/(2√x). The domain is the caller's: a +// negative input is not refused, and the NaN the forward produces +// flows into the gradient without a report, so clip or shift the +// input before it reaches the graph. +func (t *Tensor) Sqrt() (*Tensor, error) { + if err := t.checkFloat("Sqrt"); err != nil { + return nil, err + } + out, err := core.Sqrt(t.data) + if err != nil { + return nil, err + } + const eps = 1e-12 + return t.unaryResult("Sqrt", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, out) + dx := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh} + // dx = g/(2√x) with the same epsilon floor, swept from the raw + // payloads; dx stays float64 whatever the operand width. + if !g.arr.Strided() && !out.Strided() && g.arr.Dtype() == core.Float && out.Dtype() == core.Float { + gs, os, ds := g.arr.RawFloats(), out.RawFloats(), dx.arr.RawFloats()[:dx.arr.Len()] + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + sv := os[i] + if sv < eps { + sv = eps + } + ds[i] = gs[i] / (2 * sv) + } + }) + dst[0] = dx + return nil + } + if !g.arr.Strided() && !out.Strided() && g.arr.Dtype() == core.Float32 && out.Dtype() == core.Float32 { + gs, os, ds := g.arr.RawFloat32s(), out.RawFloat32s(), dx.arr.RawFloats()[:dx.arr.Len()] + elemPass(len(ds), func(s, e int) { + for i := s; i < e; i++ { + sv := float64(os[i]) + if sv < eps { + sv = eps + } + ds[i] = float64(gs[i]) / (2 * sv) + } + }) + dst[0] = dx + return nil + } + ds := dx.arr.RawFloats()[:dx.arr.Len()] + engine.Parallel(out.Len(), func(s, e int) { + for i := s; i < e; i++ { + sv := out.FloatAt(i) + if sv < eps { + sv = eps + } + ds[i] = g.arr.FloatAt(i) / (2 * sv) + } + }) + dst[0] = dx + return nil + }), nil +} + +// Floor rounds down; its derivative is zero almost everywhere, so +// Backward on its output contributes no gradient. +func (t *Tensor) Floor() (*Tensor, error) { + if err := t.checkFloat("Floor"); err != nil { + return nil, err + } + out, err := core.Floor(t.data) + if err != nil { + return nil, err + } + return t.unaryResult("Floor", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + // All-zero gradient; the borrowed buffer arrives zeroed. + sh := g.sh + dst[0] = gradSlot{arr: ar.borrowGrad(g.arr.Dtype(), sh), sh: sh} + return nil + }), nil +} + +// SetGrad replaces the accumulated gradient array: a low-level hook +// for gradient-manipulation utilities (clipping, accumulation resets, +// test harnesses) to drive through directly. +func (t *Tensor) SetGrad(g *core.Array) { t.grad = g } + +// scaleGradSlot writes g·f into a borrowed buffer shaped sh with the +// expressions mapReal writes per dtype, so the bits are MulF's own. +// The second return is false outside the dense float, float32 and +// complex paths, and the caller keeps core.MulF. +func scaleGradSlot(ar *gradArena, g gradSlot, sh []int, f float64) (gradSlot, bool) { + if g.arr == nil || g.arr.Strided() || len(sh) == 0 || len(sh) > 6 { + return gradSlot{}, false + } + n := g.arr.Len() + switch g.arr.Dtype() { + case core.Float: + out := ar.borrowGrad(core.Float, sh) + gs, os := g.arr.RawFloats()[:n], out.RawFloats()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = gs[i] * f + } + } else { + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = gs[i] * f + } + }) + } + return gradSlot{arr: out, sh: sh}, true + case core.Float32: + out := ar.borrowGrad(core.Float32, sh) + gs, os := g.arr.RawFloat32s()[:n], out.RawFloat32s()[:n] + if n < elemSweepMin { + for i := range os { + os[i] = float32(float64(gs[i]) * f) + } + } else { + elemPass(n, func(s, e int) { + for i := s; i < e; i++ { + os[i] = float32(float64(gs[i]) * f) + } + }) + } + return gradSlot{arr: out, sh: sh}, true + case core.Complex: + out := ar.borrowGrad(core.Complex, sh) + gs, os := g.arr.RawComplexes()[:n], out.RawComplexes()[:n] + for i := range os { + os[i] = gs[i] * complex(f, 0) + } + return gradSlot{arr: out, sh: sh}, true + } + return gradSlot{}, false +} + +// Scale multiplies every element by a constant factor; backward +// multiplies the incoming gradient by the same factor. +func (t *Tensor) Scale(f float64) (*Tensor, error) { + if err := t.checkDiff("Scale"); err != nil { + return nil, err + } + a := t.data + return t.unaryResult("Scale", core.MulF(t.data, f), func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { + sh := gradShape(g, a) + da, ok := scaleGradSlot(ar, g, sh, f) + if !ok { + da = gradSlot{arr: core.MulF(g.arr, f), sh: sh} + } + dst[0] = da + return nil + }), nil +} diff --git a/grad/tensor_test.go b/grad/tensor_test.go new file mode 100644 index 0000000..cfe5b48 --- /dev/null +++ b/grad/tensor_test.go @@ -0,0 +1,343 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/integrate/bench_perf_test.go b/integrate/bench_perf_test.go new file mode 100644 index 0000000..3e07c82 --- /dev/null +++ b/integrate/bench_perf_test.go @@ -0,0 +1,256 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/integrate/bench_test.go b/integrate/bench_test.go new file mode 100644 index 0000000..035aa8c --- /dev/null +++ b/integrate/bench_test.go @@ -0,0 +1,305 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/integrate/budget_pins_test.go b/integrate/budget_pins_test.go new file mode 100644 index 0000000..361fc5d --- /dev/null +++ b/integrate/budget_pins_test.go @@ -0,0 +1,75 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/integrate/cubature.go b/integrate/cubature.go new file mode 100644 index 0000000..21ca2c4 --- /dev/null +++ b/integrate/cubature.go @@ -0,0 +1,363 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/integrate/cubature_heap_test.go b/integrate/cubature_heap_test.go new file mode 100644 index 0000000..3a7b251 --- /dev/null +++ b/integrate/cubature_heap_test.go @@ -0,0 +1,112 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } + }) +} diff --git a/integrate/cubature_test.go b/integrate/cubature_test.go new file mode 100644 index 0000000..c5923e1 --- /dev/null +++ b/integrate/cubature_test.go @@ -0,0 +1,114 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/integrate/doc.go b/integrate/doc.go new file mode 100644 index 0000000..6630588 --- /dev/null +++ b/integrate/doc.go @@ -0,0 +1,109 @@ +// Copyright (c) 2026 Petr Balvín (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 diff --git a/integrate/dtypes_census_test.go b/integrate/dtypes_census_test.go new file mode 100644 index 0000000..001f537 --- /dev/null +++ b/integrate/dtypes_census_test.go @@ -0,0 +1,536 @@ +// Copyright (c) 2026 Petr Balvín (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) + }) + } + } +} diff --git a/integrate/event_boundary_pins_test.go b/integrate/event_boundary_pins_test.go new file mode 100644 index 0000000..31b8a30 --- /dev/null +++ b/integrate/event_boundary_pins_test.go @@ -0,0 +1,125 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/integrate/event_guard_pins_test.go b/integrate/event_guard_pins_test.go new file mode 100644 index 0000000..ca2706c --- /dev/null +++ b/integrate/event_guard_pins_test.go @@ -0,0 +1,236 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/integrate/example_test.go b/integrate/example_test.go new file mode 100644 index 0000000..336da59 --- /dev/null +++ b/integrate/example_test.go @@ -0,0 +1,225 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/integrate/fem2d.go b/integrate/fem2d.go new file mode 100644 index 0000000..e79b840 --- /dev/null +++ b/integrate/fem2d.go @@ -0,0 +1,376 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/integrate/fem2d_test.go b/integrate/fem2d_test.go new file mode 100644 index 0000000..7b1001b --- /dev/null +++ b/integrate/fem2d_test.go @@ -0,0 +1,477 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/integrate/fem2dnilmesh_test.go b/integrate/fem2dnilmesh_test.go new file mode 100644 index 0000000..c700466 --- /dev/null +++ b/integrate/fem2dnilmesh_test.go @@ -0,0 +1,21 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/integrate/fem3d.go b/integrate/fem3d.go new file mode 100644 index 0000000..52006ed --- /dev/null +++ b/integrate/fem3d.go @@ -0,0 +1,552 @@ +// Copyright (c) 2026 Petr Balvín (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) +} diff --git a/integrate/fem3d_test.go b/integrate/fem3d_test.go new file mode 100644 index 0000000..4919cbc --- /dev/null +++ b/integrate/fem3d_test.go @@ -0,0 +1,565 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/integrate/filon.go b/integrate/filon.go new file mode 100644 index 0000000..055b576 --- /dev/null +++ b/integrate/filon.go @@ -0,0 +1,234 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/integrate/filon_accuracy_test.go b/integrate/filon_accuracy_test.go new file mode 100644 index 0000000..514f557 --- /dev/null +++ b/integrate/filon_accuracy_test.go @@ -0,0 +1,330 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/integrate/filon_bounds_regression_test.go b/integrate/filon_bounds_regression_test.go new file mode 100644 index 0000000..b3c4c7d --- /dev/null +++ b/integrate/filon_bounds_regression_test.go @@ -0,0 +1,68 @@ +// Copyright (c) 2026 Petr Balvín (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) + } + } +} diff --git a/integrate/filon_test.go b/integrate/filon_test.go new file mode 100644 index 0000000..4e30757 --- /dev/null +++ b/integrate/filon_test.go @@ -0,0 +1,189 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/integrate/heat2d_pin_test.go b/integrate/heat2d_pin_test.go new file mode 100644 index 0000000..72977e4 --- /dev/null +++ b/integrate/heat2d_pin_test.go @@ -0,0 +1,165 @@ +// Copyright (c) 2026 Petr Balvín (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]) + } + } + } +} diff --git a/integrate/helpers_test.go b/integrate/helpers_test.go new file mode 100644 index 0000000..41b8c4d --- /dev/null +++ b/integrate/helpers_test.go @@ -0,0 +1,24 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/integrate/ode.go b/integrate/ode.go new file mode 100644 index 0000000..36046b1 --- /dev/null +++ b/integrate/ode.go @@ -0,0 +1,878 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/integrate/ode_test.go b/integrate/ode_test.go new file mode 100644 index 0000000..7dc84fe --- /dev/null +++ b/integrate/ode_test.go @@ -0,0 +1,475 @@ +// Copyright (c) 2026 Petr Balvín (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)) + } +} diff --git a/integrate/odebdf2.go b/integrate/odebdf2.go new file mode 100644 index 0000000..9358ebd --- /dev/null +++ b/integrate/odebdf2.go @@ -0,0 +1,260 @@ +// Copyright (c) 2026 Petr Balvín (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)) +} diff --git a/integrate/odebdf2_test.go b/integrate/odebdf2_test.go new file mode 100644 index 0000000..b64edc4 --- /dev/null +++ b/integrate/odebdf2_test.go @@ -0,0 +1,202 @@ +// Copyright (c) 2026 Petr Balvín (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) + } +} diff --git a/integrate/odebdfvar.go b/integrate/odebdfvar.go new file mode 100644 index 0000000..1531ef4 --- /dev/null +++ b/integrate/odebdfvar.go @@ -0,0 +1,414 @@ +// Copyright (c) 2026 Petr Balvín (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 +} diff --git a/integrate/odebdfvar_test.go b/integrate/odebdfvar_test.go new file mode 100644 index 0000000..a9ae15b --- /dev/null +++ b/integrate/odebdfvar_test.go @@ -0,0 +1,298 @@ +// Copyright (c) 2026 Petr Balvín (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") + } +} diff --git a/integrate/odeboundary.go b/integrate/odeboundary.go new file mode 100644 index 0000000..cab490b --- /dev/null +++ b/integrate/odeboundary.go @@ -0,0 +1,139 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/optim" +) + +import "math" + +// Two-point boundary value problems by shooting. The differential +// equation is integrated as an initial value problem whose free +// starting components are chosen so the trajectory lands on the +// prescribed end values: the mismatch at t1 is a function of the free +// starting components alone, and FindRootSystem drives that mismatch +// to zero. Linear problems give a mismatch linear in the unknowns, +// which the damped Newton settles in one step; nonlinear ones cost a +// handful of trajectory integrations more. + +// BoundaryConditions fixes the state of a boundary value problem at +// the two ends of the interval. Start lists the components prescribed +// at t0, whose values are read from the initial state; End lists the +// components prescribed at t1 with their values in EndValues, parallel +// to End. A component may carry a condition at both ends, as a +// second-order equation written as a first-order system does for its +// position; the remaining components are the shooting unknowns, seeded +// from the initial state. +type BoundaryConditions struct { + Start []int + End []int + EndValues []float64 +} + +// IntegrateBoundary solves the two-point boundary value problem y' = +// f(t, y) on [t0, t1] under the given conditions by shooting, and +// returns the trajectory sampled at nSamples evenly spaced times, the +// same contract as IntegrateODEPath. states[0] carries the initial +// state with the shooting unknowns replaced by the values that satisfy +// the end conditions. Backward integration (t1 < t0) works. +// +// The count of conditions decides solvability: len(Start) + len(End) +// must equal the state length, with the components in Start the ones +// whose initial values are known. An index out of range, a repeated +// index within Start or within End, an EndValues of the wrong length, +// a truncated sample count or an exhausted step budget in any trial +// trajectory is an error, never a silent answer. The shooting root +// find is local: a start whose basin holds no matching trajectory, or +// a trial that blows up on the way to t1, reports the failure. +func IntegrateBoundary(f func(t float64, y *core.Array) (*core.Array, error), + t0, t1 float64, y0 *core.Array, bc BoundaryConditions, nSamples int, + opts ODEOptions) ([]float64, []*core.Array, error) { + const name = "IntegrateBoundary" + y, err := odeCheck("IntegrateBoundary", y0, &opts) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + n := len(y) + if nSamples < 2 { + return nil, nil, base.Errf("%s: nSamples must be ≥ 2, got %d", name, nSamples) + } + if len(bc.End) == 0 { + return nil, nil, base.Errf("%s: End must prescribe at least one component at t1", name) + } + if len(bc.Start)+len(bc.End) != n { + return nil, nil, base.Errf("%s: %d conditions at t0 and %d at t1 for a state of length %d, want %d in total", + name, len(bc.Start), len(bc.End), n, n) + } + if len(bc.EndValues) != len(bc.End) { + return nil, nil, base.Errf("%s: EndValues has length %d, want %d to match End", + name, len(bc.EndValues), len(bc.End)) + } + inStart := make(map[int]bool, len(bc.Start)) + for _, j := range bc.Start { + if j < 0 || j >= n { + return nil, nil, base.Errf("%s: Start index %d out of range for a state of length %d", name, j, n) + } + if inStart[j] { + return nil, nil, base.Errf("%s: Start prescribes component %d twice", name, j) + } + inStart[j] = true + } + inEnd := make(map[int]bool, len(bc.End)) + for _, j := range bc.End { + if j < 0 || j >= n { + return nil, nil, base.Errf("%s: End index %d out of range for a state of length %d", name, j, n) + } + if inEnd[j] { + return nil, nil, base.Errf("%s: End prescribes component %d twice", name, j) + } + inEnd[j] = true + } + free := make([]int, 0, n-len(bc.Start)) + for j := range n { + if !inStart[j] { + free = append(free, j) + } + } + // The end-point mismatch as a function of the free starting + // components. Each evaluation is one full trajectory, so the + // integration tolerance bounds the noise the root find must see + // and its target sits just above that noise. + residual := func(x *core.Array) (*core.Array, error) { + u0 := cloneDenseSlice(y) + for k, j := range free { + u0[j] = x.FloatAt(k) + } + u1, err := IntegrateODE(f, t0, t1, wrapVector(u0), opts) + if err != nil { + return nil, err + } + r := make([]float64, len(bc.End)) + for k, j := range bc.End { + r[k] = u1.FloatAt(j) - bc.EndValues[k] + } + return wrapVector(r), nil + } + guess := make([]float64, len(free)) + for k, j := range free { + guess[k] = y[j] + } + tol := math.Max(10*opts.RelTol, 1e-12) + solution, _, err := optim.FindRootSystem(residual, wrapVector(guess), + optim.RootSystemOptions{Tolerance: tol}) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + u0 := cloneDenseSlice(y) + for k, j := range free { + u0[j] = solution.FloatAt(k) + } + times, states, err := IntegrateODEPath(f, t0, t1, wrapVector(u0), nSamples, opts) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + return times, states, nil +} diff --git a/integrate/odeboundary_test.go b/integrate/odeboundary_test.go new file mode 100644 index 0000000..ac5b20b --- /dev/null +++ b/integrate/odeboundary_test.go @@ -0,0 +1,152 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "testing" +) + +// harmonicSystem returns f for y” = −y written as the first-order +// system u' = (v, −u). +func harmonicSystem(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2) +} + +// TestIntegrateBoundaryLinear shoots the classic y” = −y with +// y(0) = 0 and y(π/2) = 1: the free initial slope must come out as 1 +// and the sampled trajectory must trace y = sin t. +func TestIntegrateBoundaryLinear(t *testing.T) { + bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}} + times, states, err := IntegrateBoundary(harmonicSystem, 0, math.Pi/2, + mustFloats(t, []float64{0, 0.5}), bc, 5, ODEOptions{RelTol: 1e-9, AbsTol: 1e-12}) + if err != nil { + t.Fatalf("IntegrateBoundary: %v", err) + } + for i := range times { + want := math.Sin(times[i]) + if math.Abs(states[i].FloatAt(0)-want) > 1e-6 { + t.Fatalf("y(%.6g) = %.12g, want %.12g", times[i], states[i].FloatAt(0), want) + } + } + if math.Abs(states[0].FloatAt(1)-1) > 1e-5 { + t.Fatalf("shooting slope = %.12g, want 1", states[0].FloatAt(1)) + } + if math.Abs(times[4]-math.Pi/2) > 1e-12 { + t.Fatalf("last sample time %.14g, want π/2", times[4]) + } +} + +// TestIntegrateBoundaryNonlinear shoots y” = 2y³ with y(0) = 1 and +// y(1) = 1/2, whose exact solution is y = 1/(t+1) with the initial +// slope −1, and checks the sampled path against it. +func TestIntegrateBoundaryNonlinear(t *testing.T) { + f := func(t float64, y *core.Array) (*core.Array, error) { + u := y.FloatAt(0) + return core.FromFloats([]float64{y.FloatAt(1), 2 * u * u * u}, 2) + } + bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{0.5}} + times, states, err := IntegrateBoundary(f, 0, 1, + mustFloats(t, []float64{1, -0.5}), bc, 5, ODEOptions{RelTol: 1e-10, AbsTol: 1e-13}) + if err != nil { + t.Fatalf("IntegrateBoundary: %v", err) + } + for i := range times { + want := 1 / (times[i] + 1) + if math.Abs(states[i].FloatAt(0)-want) > 1e-6 { + t.Fatalf("y(%.6g) = %.12g, want %.12g", times[i], states[i].FloatAt(0), want) + } + } + if math.Abs(states[0].FloatAt(1)+1) > 1e-5 { + t.Fatalf("shooting slope = %.12g, want −1", states[0].FloatAt(1)) + } +} + +// TestIntegrateBoundaryVelocityEnd prescribes the velocity at t1 +// instead of the position: y” = −y with y(0) = 1 and y'(1) = 0 has +// y = cos t + tan(1)·sin t, and the shooting unknown is the initial +// value of the very component the end condition watches. +func TestIntegrateBoundaryVelocityEnd(t *testing.T) { + bc := BoundaryConditions{Start: []int{0}, End: []int{1}, EndValues: []float64{0}} + times, states, err := IntegrateBoundary(harmonicSystem, 0, 1, + mustFloats(t, []float64{1, 1}), bc, 4, ODEOptions{RelTol: 1e-9, AbsTol: 1e-12}) + if err != nil { + t.Fatalf("IntegrateBoundary: %v", err) + } + slope := math.Tan(1) + for i := range times { + want := math.Cos(times[i]) + slope*math.Sin(times[i]) + if math.Abs(states[i].FloatAt(0)-want) > 1e-6 { + t.Fatalf("y(%.6g) = %.12g, want %.12g", times[i], states[i].FloatAt(0), want) + } + } +} + +// TestIntegrateBoundaryBlowUp pins the honest failure: y” = y³ with +// y(0) = 1 and y(1) = 100 demands a trajectory that reaches 100 only +// by skirting its own finite-time blow-up, so some trial integration +// fails and the shooting reports the error instead of an answer. +func TestIntegrateBoundaryBlowUp(t *testing.T) { + f := func(t float64, y *core.Array) (*core.Array, error) { + u := y.FloatAt(0) + return core.FromFloats([]float64{y.FloatAt(1), u * u * u}, 2) + } + bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{100}} + if _, _, err := IntegrateBoundary(f, 0, 1, mustFloats(t, []float64{1, 0}), bc, 3, + ODEOptions{MaxSteps: 2000}); err == nil { + t.Fatal("expected the shooting to fail on the blow-up problem") + } +} + +// TestIntegrateBoundaryErrors pins the validation contract. +func TestIntegrateBoundaryErrors(t *testing.T) { + y0 := mustFloats(t, []float64{1, 0}) + bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}} + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, y0, bc, 1, ODEOptions{}); err == nil { + t.Fatal("expected an error for a single sample") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, y0, + BoundaryConditions{Start: []int{0}}, 3, ODEOptions{}); err == nil { + t.Fatal("expected an error when nothing is prescribed at t1") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, y0, + BoundaryConditions{Start: []int{0, 1}, End: []int{0}, EndValues: []float64{1}}, + 3, ODEOptions{}); err == nil { + t.Fatal("expected an error for too many conditions") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, y0, + BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{}}, + 3, ODEOptions{}); err == nil { + t.Fatal("expected an error when EndValues does not match End") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, y0, + BoundaryConditions{Start: []int{5}, End: []int{1}, EndValues: []float64{0}}, + 3, ODEOptions{}); err == nil { + t.Fatal("expected an error for a Start index out of range") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, y0, + BoundaryConditions{Start: []int{0}, End: []int{2}, EndValues: []float64{1}}, + 3, ODEOptions{}); err == nil { + t.Fatal("expected an error for an End index out of range") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, mustFloats(t, []float64{1, 0, 2}), + BoundaryConditions{Start: []int{0, 0}, End: []int{2}, EndValues: []float64{1}}, + 3, ODEOptions{}); err == nil { + t.Fatal("expected an error for a repeated Start index") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, y0, + BoundaryConditions{Start: []int{0}, End: []int{1, 1}, EndValues: []float64{0, 1}}, + 3, ODEOptions{}); err == nil { + t.Fatal("expected an error for a repeated End index") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, mustFloats(t, []float64{1}, 1), + bc, 3, ODEOptions{}); err == nil { + t.Fatal("expected an error for a rank-2 state") + } + if _, _, err := IntegrateBoundary(harmonicSystem, 0, 1, mustFloats(t, nil), bc, 3, ODEOptions{}); err == nil { + t.Fatal("expected an error for an empty state") + } +} diff --git a/integrate/odecolloc.go b/integrate/odecolloc.go new file mode 100644 index 0000000..bcb9aa5 --- /dev/null +++ b/integrate/odecolloc.go @@ -0,0 +1,548 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Boundary value problems by collocation, the mesh-based sibling of +// the shooting method in IntegrateBoundary. Instead of marching a +// single trajectory and tuning its free start, the solver lays a mesh +// over [t0, t1], represents the solution by a cubic on every mesh +// interval, drives the whole discrete system to zero by a damped +// Newton iteration over a numerically assembled Jacobian, and then +// halves the intervals whose residual estimate is past tolerance and +// solves again, until every interval sits inside the tolerance or the +// node budget runs out. +// +// The scheme is the classic three-point Lobatto IIIA collocation, not +// the Kierzenka-Shampine variant: the collocation polynomial on each +// interval satisfies the ODE at both endpoints and the midpoint, +// which makes the nodal values fourth-order accurate in the interval +// width. The unknowns follow the shape scipy's solve_bvp solves for: +// the state y at every mesh node and the slope s = f(t, y) at every +// mesh node. Between the nodes the solution is the piecewise cubic +// Hermite through (t, y, s), which is exactly the collocation +// polynomial, so that triple is the solver's S-slope representation +// of the continuous answer. + +// CollocationOptions tunes SolveBoundaryCollocation. RelTol ≤ 0 means +// 1e-6, AbsTol ≤ 0 means 1e-9, InitialNodes ≤ 0 means 10, MaxNodes ≤ 0 +// means 256 and MaxIterations ≤ 0 means 40. +type CollocationOptions struct { + // RelTol and AbsTol scale the mesh-refinement estimate: an + // interval whose root-mean-square of residual over + // AbsTol + RelTol·|slope| stays above 1 is halved. The same pair + // floors the Newton convergence, an order of magnitude below it. + RelTol float64 + AbsTol float64 + // InitialNodes is the interval count of the uniform starting + // mesh. + InitialNodes int + // MaxNodes bounds the refined mesh. The Newton matrix is factored + // by the library's dense LU, so the cap also bounds the per-round + // cost; a two-component system lives comfortably at 256, well + // inside memory. + MaxNodes int + // MaxIterations bounds the damped Newton rounds on each mesh. + MaxIterations int +} + +// CollocationSolution carries the solved problem: Mesh holds the node +// times, Values[k] the state at Mesh[k] and Slopes[k] the derivative +// y' = f(t, y) there. The piecewise cubic Hermite through +// (Mesh, Values, Slopes) is the collocation solution itself, so +// interpolating from that data between the nodes is exact to the +// solver's tolerance. +type CollocationSolution struct { + Mesh []float64 + Values []*core.Array + Slopes []*core.Array +} + +// SolveBoundaryCollocation solves the two-point boundary value +// problem y' = f(t, y) on [t0, t1] by three-point Lobatto IIIA +// collocation on an adaptively refined mesh, returning the mesh, the +// nodal states and the nodal slopes. The boundary conditions follow +// the BoundaryConditions contract of IntegrateBoundary: Start lists +// the components prescribed at t0 with values read from y0, End the +// components prescribed at t1 with EndValues, and exactly n +// conditions must be given in total, because the collocation system +// is square. The initial guess interpolates linearly between the +// prescribed endpoint states and reads its slopes from f. +// +// Refusal is part of the contract: inconsistent boundary conditions +// (fewer or more than n conditions, out-of-range or repeated +// indices, EndValues of the wrong length), a non-positive interval, +// a starting mesh past MaxNodes, refinement that would grow past +// MaxNodes, a singular Newton matrix or an iteration that cannot +// converge are errors, never silent answers. +func SolveBoundaryCollocation(f func(t float64, y *core.Array) (*core.Array, error), + t0, t1 float64, y0 *core.Array, bc BoundaryConditions, + opts CollocationOptions) (*CollocationSolution, error) { + const name = "SolveBoundaryCollocation" + if opts.RelTol <= 0 { + opts.RelTol = 1e-6 + } + if opts.AbsTol <= 0 { + opts.AbsTol = 1e-9 + } + if opts.InitialNodes <= 0 { + opts.InitialNodes = 10 + } + if opts.MaxNodes <= 0 { + opts.MaxNodes = 256 + } + if opts.MaxIterations <= 0 { + opts.MaxIterations = 40 + } + 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 !(t1 > t0) { + return nil, base.Errf("%s: the interval must have positive length, got [%g, %g]", name, t0, t1) + } + n := y0.Len() + seed := make([]float64, n) + for i := range n { + seed[i] = y0.FloatAt(i) + if math.IsNaN(seed[i]) || math.IsInf(seed[i], 0) { + return nil, base.Errf("%s: the state holds the non-finite value %g at %d", name, seed[i], i) + } + } + if len(bc.End) == 0 { + return nil, base.Errf("%s: End must prescribe at least one component at t1", name) + } + if len(bc.Start)+len(bc.End) != n { + return nil, base.Errf("%s: %d conditions at t0 and %d at t1 for a state of length %d, want %d in total", + name, len(bc.Start), len(bc.End), n, n) + } + if len(bc.EndValues) != len(bc.End) { + return nil, base.Errf("%s: EndValues has length %d, want %d to match End", + name, len(bc.EndValues), len(bc.End)) + } + inStart := make(map[int]bool, len(bc.Start)) + for _, j := range bc.Start { + if j < 0 || j >= n { + return nil, base.Errf("%s: Start index %d out of range for a state of length %d", name, j, n) + } + if inStart[j] { + return nil, base.Errf("%s: Start prescribes component %d twice", name, j) + } + inStart[j] = true + } + inEnd := make(map[int]bool, len(bc.End)) + for _, j := range bc.End { + if j < 0 || j >= n { + return nil, base.Errf("%s: End index %d out of range for a state of length %d", name, j, n) + } + if inEnd[j] { + return nil, base.Errf("%s: End prescribes component %d twice", name, j) + } + inEnd[j] = true + } + endState := cloneDenseSlice(seed) + for q, j := range bc.End { + endState[j] = bc.EndValues[q] + } + if opts.InitialNodes < 1 { + return nil, base.Errf("%s: InitialNodes must be ≥ 1, got %d", name, opts.InitialNodes) + } + if opts.InitialNodes+1 > opts.MaxNodes { + return nil, base.Errf("%s: the starting mesh of %d intervals already exceeds MaxNodes=%d", + name, opts.InitialNodes, opts.MaxNodes) + } + + // The starting mesh and guess: uniform in t, linear between the + // prescribed endpoint states, slopes read from f. + mesh := make([]float64, opts.InitialNodes+1) + for k := range mesh { + mesh[k] = t0 + (t1-t0)*float64(k)/float64(opts.InitialNodes) + } + stride := 2 * n + z := make([]float64, stride*(opts.InitialNodes+1)) + for k := range opts.InitialNodes + 1 { + theta := (mesh[k] - t0) / (t1 - t0) + for i := range n { + z[k*stride+i] = seed[i] + theta*(endState[i]-seed[i]) + } + } + for k := range opts.InitialNodes + 1 { + sv, err := odeEval(name, f, mesh[k], z[k*stride:k*stride+n], n, nil) + if err != nil { + return nil, err + } + copy(z[k*stride+n:(k+1)*stride], sv) + } + + // collocResidual writes the discrete system for the unknown vector + // zz into dst: one slope-definition block per node, one + // collocation block per interval (the midpoint value eliminated + // through the cubic Hermite it belongs to), then the boundary + // rows. The row count equals the unknown count exactly. + collocResidual := func(dst, zz, msh []float64) error { + nodes := len(msh) + for k := range nodes { + fv, err := odeEval(name, f, msh[k], zz[k*stride:k*stride+n], n, nil) + if err != nil { + return err + } + for i := range n { + dst[k*n+i] = zz[k*stride+n+i] - fv[i] + } + } + baseRow := nodes * n + ym := make([]float64, n) + for i := range nodes - 1 { + h := msh[i+1] - msh[i] + for j := range n { + ym[j] = (zz[i*stride+j]+zz[(i+1)*stride+j])/2 + h*(zz[i*stride+n+j]-zz[(i+1)*stride+n+j])/8 + } + fm, err := odeEval(name, f, msh[i]+h/2, ym, n, nil) + if err != nil { + return err + } + for j := range n { + dst[baseRow+i*n+j] = zz[(i+1)*stride+j] - zz[i*stride+j] - + h*(zz[i*stride+n+j]+4*fm[j]+zz[(i+1)*stride+n+j])/6 + } + } + last := len(dst) - n + for p, j := range bc.Start { + dst[last+p] = zz[j] - seed[j] + } + for q, j := range bc.End { + dst[last+len(bc.Start)+q] = zz[(nodes-1)*stride+j] - bc.EndValues[q] + } + return nil + } + + // collocJac assembles the Newton matrix by central differences: + // the slope rows differentiate y − f(t, y) over their node's y + // block (their slope columns are the exact −I), the collocation + // rows differentiate the interval map over the four blocks + // (yᵢ, sᵢ, yᵢ₊₁, sᵢ₊₁), and the boundary rows enter exactly. + collocJac := func(zz, msh []float64) ([][]float64, error) { + nodes := len(msh) + size := stride * nodes + jac := make([][]float64, size) + for i := range jac { + jac[i] = make([]float64, size) + } + for k := range nodes { + // The slope rows are s − f(y): as a function of the node's + // y block the residual is −f(y), and the slope columns + // carry the +I. + block := func(y []float64) ([]float64, error) { + fv, err := odeEval(name, f, msh[k], y, n, nil) + if err != nil { + return nil, err + } + out := make([]float64, n) + for i := range n { + out[i] = -fv[i] + } + return out, nil + } + if err := collocNumJac(name, block, zz[k*stride:k*stride+n], k*n, k*stride, n, n, jac); err != nil { + return nil, err + } + for i := range n { + // The slope rows are s − f(y), so the slope columns + // carry +I. + jac[k*n+i][k*stride+n+i] = 1 + } + } + baseRow := nodes * n + w := make([]float64, 4*n) + ym := make([]float64, n) + for i := range nodes - 1 { + h := msh[i+1] - msh[i] + copy(w[0:n], zz[i*stride:i*stride+n]) + copy(w[n:2*n], zz[i*stride+n:(i+1)*stride]) + copy(w[2*n:3*n], zz[(i+1)*stride:(i+1)*stride+n]) + copy(w[3*n:4*n], zz[(i+1)*stride+n:(i+2)*stride]) + interval := func(x []float64) ([]float64, error) { + for j := range n { + ym[j] = (x[j]+x[2*n+j])/2 + h*(x[n+j]-x[3*n+j])/8 + } + fm, err := odeEval(name, f, msh[i]+h/2, ym, n, nil) + if err != nil { + return nil, err + } + out := make([]float64, n) + for j := range n { + out[j] = x[2*n+j] - x[j] - h*(x[n+j]+4*fm[j]+x[3*n+j])/6 + } + return out, nil + } + if err := collocNumJac(name, interval, w, baseRow+i*n, i*stride, n, 4*n, jac); err != nil { + return nil, err + } + } + last := size - n + for p, j := range bc.Start { + jac[last+p][j] = 1 + } + for q, j := range bc.End { + jac[last+len(bc.Start)+q][(nodes-1)*stride+j] = 1 + } + return jac, nil + } + + // newtonSolve drives the damped Newton on the fixed mesh: the + // Jacobian is frozen from the seed and rebuilt twice when + // convergence drags, as odeNewton does, and each round backtracks + // along the step until the residual infinity norm actually falls. + newtonSolve := func(zz, msh []float64) error { + size := stride * len(msh) + r := make([]float64, size) + trialZ := make([]float64, size) + trialR := make([]float64, size) + col := make([]float64, size) + if err := collocResidual(r, zz, msh); err != nil { + return err + } + for i := range r { + if math.IsNaN(r[i]) || math.IsInf(r[i], 0) { + return base.Errf("%s: the residual returned the non-finite value %g at row %d", name, r[i], i) + } + } + scale := normInfOfStep(zz) + // The algebraic floor sits three orders below the mesh + // tolerance: the refinement estimator reads the true ODE + // residual of the interpolant, and that reading must not be + // dominated by the residual the Newton iteration left. + limit := math.Max(0.001*(opts.AbsTol+opts.RelTol*scale), 8*base.EpsF*(scale+1)) + worst := normInfOfStep(r) + var work [][]float64 + var perm []int + for iteration := 0; iteration < opts.MaxIterations; iteration++ { + if worst <= limit { + return nil + } + if iteration == 0 || iteration == 4 || iteration == 10 { + jac, err := collocJac(zz, msh) + if err != nil { + return err + } + // Factor a working copy: base.Factor consumes its + // argument in place, and the pristine matrix is not + // needed again before the next rebuild. + work = make([][]float64, size) + for i := range jac { + work[i] = cloneDenseSlice(jac[i]) + } + perm, _ = base.Factor(work) + if err := base.CheckSingular(name, work); err != nil { + return base.Errf("%s: %w, singular collocation Newton matrix", name, errNewtonStalled) + } + } + for i := range col { + col[i] = -r[i] + } + base.PermuteColumn(col, perm) + base.SolveColumn(work, col) + accepted := false + factor := 1.0 + for range 40 { + for i := range zz { + trialZ[i] = zz[i] + factor*col[i] + } + if err := collocResidual(trialR, trialZ, msh); err == nil { + finite := true + cand := 0.0 + for i := range trialR { + v := trialR[i] + if math.IsNaN(v) || math.IsInf(v, 0) { + finite = false + break + } + cand = math.Max(cand, math.Abs(v)) + } + if finite && cand <= (1-1e-4*factor)*worst { + copy(zz, trialZ) + // The residual travels with the accepted point, + // so the next round solves against the state + // this round left behind. + copy(r, trialR) + worst = cand + accepted = true + break + } + } + factor /= 2 + } + if !accepted { + return base.Errf("%s: %w: the residual cannot be reduced below %g by damping", + name, errNewtonStalled, worst) + } + } + return base.Errf("%s: %w after %d rounds, residual %g", name, errNewtonStalled, opts.MaxIterations, worst) + } + + // estimateRefinement returns the intervals whose root-mean-square + // scaled residual is past 1, measured from the collocation + // solution's own cubic Hermite: the ODE is evaluated at three + // interior quadrature points of every interval (the nodes one half + // plus or minus half the square root of three sevenths and the + // midpoint, weights 49/180, 16/45, 49/180) and the mismatch with + // the Hermite derivative is quadrature-weighted (the endpoints + // contribute exactly nothing, the Hermite slope is the ODE slope + // there by construction). + estimateRefinement := func(zz, msh []float64) ([]int, error) { + nodes := len(msh) + theta := [3]float64{0.5 * (1 - math.Sqrt(3.0/7)), 0.5, 0.5 * (1 + math.Sqrt(3.0/7))} + weight := [3]float64{49.0 / 180, 16.0 / 45, 49.0 / 180} + var bad []int + val := make([]float64, n) + for i := range nodes - 1 { + h := msh[i+1] - msh[i] + sum := 0.0 + for pt := range 3 { + th := theta[pt] + th2 := th * th + th3 := th2 * th + h00 := 2*th3 - 3*th2 + 1 + h10 := th3 - 2*th2 + th + h01 := -2*th3 + 3*th2 + h11 := th3 - th2 + hd00 := 6*th2 - 6*th + hd10 := 3*th2 - 4*th + 1 + hd01 := -6*th2 + 6*th + hd11 := 3*th2 - 2*th + for j := range n { + val[j] = h00*zz[i*stride+j] + h*h10*zz[i*stride+n+j] + + h01*zz[(i+1)*stride+j] + h*h11*zz[(i+1)*stride+n+j] + } + fv, err := odeEval(name, f, msh[i]+th*h, val, n, nil) + if err != nil { + return nil, err + } + for j := range n { + der := hd00*zz[i*stride+j]/h + hd10*zz[i*stride+n+j] + + hd01*zz[(i+1)*stride+j]/h + hd11*zz[(i+1)*stride+n+j] + slope := math.Max(math.Abs(zz[i*stride+n+j]), math.Abs(zz[(i+1)*stride+n+j])) + sc := opts.AbsTol + opts.RelTol*math.Max(slope, math.Abs(fv[j])) + ratio := (der - fv[j]) / sc + sum += weight[pt] * ratio * ratio + } + } + if math.Sqrt(sum) > 1 { + bad = append(bad, i) + } + } + return bad, nil + } + + // refineMesh inserts the midpoint of every interval in bad, with + // the new node's state and slope taken from the collocation + // solution's own cubic Hermite at the midpoint. + refineMesh := func(zz, msh []float64, bad []int) ([]float64, []float64, error) { + nodes := len(msh) + if nodes+len(bad) > opts.MaxNodes { + return nil, nil, base.Errf("%s: refining %d intervals would grow the mesh to %d nodes past MaxNodes=%d", + name, len(bad), nodes+len(bad), opts.MaxNodes) + } + badSet := make(map[int]bool, len(bad)) + for _, i := range bad { + badSet[i] = true + } + newMesh := make([]float64, 0, nodes+len(bad)) + newZ := make([]float64, 0, len(zz)+2*n*len(bad)) + push := func(t float64, y, s []float64) { + newMesh = append(newMesh, t) + newZ = append(newZ, y...) + newZ = append(newZ, s...) + } + for i := range nodes - 1 { + push(msh[i], zz[i*stride:i*stride+n], zz[i*stride+n:(i+1)*stride]) + if !badSet[i] { + continue + } + h := msh[i+1] - msh[i] + mid := (msh[i] + msh[i+1]) / 2 + ym := make([]float64, n) + sm := make([]float64, n) + for j := range n { + ym[j] = (zz[i*stride+j]+zz[(i+1)*stride+j])/2 + h*(zz[i*stride+n+j]-zz[(i+1)*stride+n+j])/8 + sm[j] = 1.5*(zz[(i+1)*stride+j]-zz[i*stride+j])/h - (zz[i*stride+n+j]+zz[(i+1)*stride+n+j])/4 + } + push(mid, ym, sm) + } + push(msh[nodes-1], zz[(nodes-1)*stride:(nodes-1)*stride+n], zz[(nodes-1)*stride+n:nodes*stride]) + return newMesh, newZ, nil + } + + for pass := 0; ; pass++ { + if pass >= 100 { + return nil, base.Errf("%s: refinement did not settle within 100 passes", name) + } + if err := newtonSolve(z, mesh); err != nil { + return nil, err + } + bad, err := estimateRefinement(z, mesh) + if err != nil { + return nil, err + } + if len(bad) == 0 { + break + } + mesh, z, err = refineMesh(z, mesh, bad) + if err != nil { + return nil, err + } + } + + solution := &CollocationSolution{ + Mesh: slices.Clone(mesh), + Values: make([]*core.Array, len(mesh)), + Slopes: make([]*core.Array, len(mesh)), + } + for k := range mesh { + solution.Values[k] = arrayFromVector(z[k*stride : k*stride+n]) + solution.Slopes[k] = arrayFromVector(z[k*stride+n : (k+1)*stride]) + } + return solution, nil +} + +// collocNumJac fills jac[row0+i][col0+c] with the central-difference +// derivative of g's i-th output against x's c-th entry, one column +// per entry of x. g receives its own perturbed copy of x and returns +// a fresh output slice, so nothing aliases. +func collocNumJac(name string, g func(x []float64) ([]float64, error), x []float64, + row0, col0, rows, cols int, jac [][]float64) error { + xp := make([]float64, cols) + xm := make([]float64, cols) + for c := range cols { + eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(x[c])) + copy(xp, x) + copy(xm, x) + xp[c] += eps + xm[c] -= eps + rp, e1 := g(xp) + if e1 != nil { + return base.Errf("%s: %w", name, e1) + } + rm, e2 := g(xm) + if e2 != nil { + return base.Errf("%s: %w", name, e2) + } + for i := range rows { + jac[row0+i][col0+c] = (rp[i] - rm[i]) / (2 * eps) + } + } + return nil +} diff --git a/integrate/odecolloc_test.go b/integrate/odecolloc_test.go new file mode 100644 index 0000000..bb85daf --- /dev/null +++ b/integrate/odecolloc_test.go @@ -0,0 +1,273 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// collocCubicSystem is y” = 6t written as u' = v, v' = 6t, whose +// exact solution with u(0) = 0, u(1) = 1 is u = t³, v = 3t². The +// three-point Lobatto IIIA collocation reproduces a cubic exactly, so +// the discrete solve is the exact answer, not an approximation. +func collocCubicSystem(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), 6 * t}, 2) +} + +// collocHermite evaluates the solution's piecewise cubic Hermite +// through (mesh, values, slopes) at time tau, the documented +// continuous representation of the collocation answer. +func collocHermite(t *testing.T, sol *CollocationSolution, tau float64, component int) float64 { + t.Helper() + if tau < sol.Mesh[0] || tau > sol.Mesh[len(sol.Mesh)-1] { + t.Fatalf("time %g outside the mesh", tau) + } + lo, hi := 0, len(sol.Mesh)-1 + for hi-lo > 1 { + mid := (lo + hi) / 2 + if sol.Mesh[mid] <= tau { + lo = mid + } else { + hi = mid + } + } + h := sol.Mesh[lo+1] - sol.Mesh[lo] + th := (tau - sol.Mesh[lo]) / h + th2 := th * th + th3 := th2 * th + y0 := sol.Values[lo].FloatAt(component) + y1 := sol.Values[lo+1].FloatAt(component) + s0 := sol.Slopes[lo].FloatAt(component) + s1 := sol.Slopes[lo+1].FloatAt(component) + return (2*th3-3*th2+1)*y0 + h*(th3-2*th2+th)*s0 + + (-2*th3+3*th2)*y1 + h*(th3-th2)*s1 +} + +// TestSolveBoundaryCollocationCubicExact solves the linear problem +// u” = 6t with u(0) = 0, u(1) = 1 on a uniform mesh: the cubic +// collocation reproduces t³ exactly, the Newton residual drops to +// machine precision in one step, and no refinement is needed, so the +// mesh keeps its initial size. +func TestSolveBoundaryCollocationCubicExact(t *testing.T) { + sol, err := SolveBoundaryCollocation(collocCubicSystem, 0, 1, + mustFloats(t, []float64{0, 0}), + BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}}, + CollocationOptions{RelTol: 1e-7, InitialNodes: 8, MaxNodes: 64}) + if err != nil { + t.Fatalf("SolveBoundaryCollocation: %v", err) + } + if len(sol.Mesh) != 9 { + t.Fatalf("mesh grew to %d nodes on a cubic-exact problem, want the initial 9", len(sol.Mesh)) + } + for k, tk := range sol.Mesh { + if math.Abs(sol.Values[k].FloatAt(0)-tk*tk*tk) > 1e-12 { + t.Fatalf("u(%.6g) = %.14g, want %.14g", tk, sol.Values[k].FloatAt(0), tk*tk*tk) + } + if math.Abs(sol.Slopes[k].FloatAt(0)-3*tk*tk) > 1e-12 { + t.Fatalf("u'(%.6g) = %.14g, want %.14g", tk, sol.Slopes[k].FloatAt(0), 3*tk*tk) + } + if math.Abs(sol.Values[k].FloatAt(1)-3*tk*tk) > 1e-12 { + t.Fatalf("v(%.6g) = %.14g, want %.14g", tk, sol.Values[k].FloatAt(1), 3*tk*tk) + } + } + // A one-component state carries exactly one condition: an + // endpoint-only prescription solves y' = y backward from t1. + sol1, err := SolveBoundaryCollocation( + func(t float64, y *core.Array) (*core.Array, error) { return core.MulF(y, 1), nil }, + 0, 1, mustFloats(t, []float64{1}), + BoundaryConditions{End: []int{0}, EndValues: []float64{math.E}}, + CollocationOptions{}) + if err != nil { + t.Fatalf("one-component solve: %v", err) + } + for k, tk := range sol1.Mesh { + if math.Abs(sol1.Values[k].FloatAt(0)-math.Exp(tk)) > 1e-6 { + t.Fatalf("y(%.6g) = %.14g, want %.14g", tk, sol1.Values[k].FloatAt(0), math.Exp(tk)) + } + } +} + +// TestSolveBoundaryCollocationBratuMatchesShooting solves Bratu's +// equation u” + e^u = 0 with u(0) = u(1) = 0, λ = 1, and requires +// the collocation answer to agree with the shooting method's answer +// through the existing IntegrateBoundary. +func TestSolveBoundaryCollocationBratuMatchesShooting(t *testing.T) { + bratu := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), -math.Exp(y.FloatAt(0))}, 2) + } + sol, err := SolveBoundaryCollocation(bratu, 0, 1, + mustFloats(t, []float64{0, 0.4}), + BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{0}}, + CollocationOptions{RelTol: 1e-7, AbsTol: 1e-10, InitialNodes: 10, MaxNodes: 600, MaxIterations: 60}) + if err != nil { + t.Fatalf("collocation: %v", err) + } + if len(sol.Mesh) <= 10 { + t.Fatalf("the mesh never grew past the initial 10 intervals (refinement instrument): %d nodes", len(sol.Mesh)) + } + times, states, err := IntegrateBoundary(bratu, 0, 1, + mustFloats(t, []float64{0, 0.4}), + BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{0}}, + 9, ODEOptions{RelTol: 1e-10, AbsTol: 1e-13}) + if err != nil { + t.Fatalf("shooting: %v", err) + } + worstU, worstS := 0.0, 0.0 + for i := 1; i < len(times)-1; i++ { + if d := math.Abs(collocHermite(t, sol, times[i], 0) - states[i].FloatAt(0)); d > worstU { + worstU = d + } + if d := math.Abs(collocHermite(t, sol, times[i], 1) - states[i].FloatAt(1)); d > worstS { + worstS = d + } + } + t.Logf("Bratu λ=1: worst state difference %.3g, worst slope difference %.3g", worstU, worstS) + if worstU > 1e-5 { + t.Fatalf("collocation and shooting disagree on u by %.3g", worstU) + } + if worstS > 1e-4 { + t.Fatalf("collocation and shooting disagree on u' by %.3g", worstS) + } +} + +// TestSolveBoundaryCollocationRefinementLayer pins the refinement +// loop with a linear boundary-layer problem u” = −100·u' scaled as +// u' = v, v' = −100v, whose solution 1 − e^(−100t) needs intervals +// clustered near t = 0. A loose tolerance must leave the initial +// mesh alone; a tight one must refine it and land on the solution. +func TestSolveBoundaryCollocationRefinementLayer(t *testing.T) { + layer := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), -100 * y.FloatAt(1)}, 2) + } + exact := func(t float64) float64 { return 1 - math.Exp(-100*t) } + bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}} + y0 := mustFloats(t, []float64{0, 0}) + loose, err := SolveBoundaryCollocation(layer, 0, 1, y0, bc, + CollocationOptions{AbsTol: 1e6, RelTol: 1, InitialNodes: 10, MaxNodes: 4000}) + if err != nil { + t.Fatalf("loose solve: %v", err) + } + if len(loose.Mesh) != 11 { + t.Fatalf("a loose tolerance still refined to %d nodes", len(loose.Mesh)) + } + tight, err := SolveBoundaryCollocation(layer, 0, 1, y0, bc, + CollocationOptions{RelTol: 1e-6, AbsTol: 1e-9, InitialNodes: 10, MaxNodes: 4000, MaxIterations: 60}) + if err != nil { + t.Fatalf("tight solve: %v", err) + } + t.Logf("layer problem: mesh refined from 11 to %d nodes", len(tight.Mesh)) + if len(tight.Mesh) <= 11 { + t.Fatal("the tight solve never refined the mesh (refinement instrument)") + } + worst := 0.0 + for k, tk := range tight.Mesh { + if d := math.Abs(tight.Values[k].FloatAt(0) - exact(tk)); d > worst { + worst = d + } + } + t.Logf("layer problem: worst nodal error %.3g", worst) + if worst > 1e-3 { + t.Fatalf("refined layer error %.3g too large", worst) + } + for k := range tight.Mesh { + if k > 0 && tight.Mesh[k] <= tight.Mesh[k-1] { + t.Fatalf("the mesh is not increasing at %d", k) + } + } +} + +// TestSolveBoundaryCollocationErrors pins the refusal contract. +func TestSolveBoundaryCollocationErrors(t *testing.T) { + const name = "SolveBoundaryCollocation" + good := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2) + } + y0 := mustFloats(t, []float64{0, 0.5}) + bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}} + // Inconsistent boundary conditions: wrong counts, repeated and + // out-of-range indices, wrong EndValues length. + if _, err := SolveBoundaryCollocation(good, 0, 1, y0, + BoundaryConditions{End: []int{0}, EndValues: []float64{1}}, CollocationOptions{}); err == nil || !stringsContains(err, "in total") { + t.Fatalf("%s: one condition for a two-component state: %v", name, err) + } + if _, err := SolveBoundaryCollocation(good, 0, 1, y0, + BoundaryConditions{Start: []int{0, 1}, End: []int{0, 1}, EndValues: []float64{1, 2}}, + CollocationOptions{}); err == nil || !stringsContains(err, "in total") { + t.Fatalf("%s: four conditions for a two-component state: %v", name, err) + } + if _, err := SolveBoundaryCollocation(good, 0, 1, y0, + BoundaryConditions{Start: []int{0}, End: []int{0}}, CollocationOptions{}); err == nil || !stringsContains(err, "EndValues") { + t.Fatalf("%s: EndValues mismatch: %v", name, err) + } + if _, err := SolveBoundaryCollocation(good, 0, 1, y0, + BoundaryConditions{Start: []int{0}, End: []int{2}, EndValues: []float64{1}}, + CollocationOptions{}); err == nil || !stringsContains(err, "out of range") { + t.Fatalf("%s: End index out of range: %v", name, err) + } + // The duplicate-index cases need the total count to be right + // first, so they run on a three-component state. + good3 := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0), y.FloatAt(2)}, 3) + } + if _, err := SolveBoundaryCollocation(good3, 0, 1, mustFloats(t, []float64{0, 0.5, 1}), + BoundaryConditions{Start: []int{0, 0}, End: []int{2}, EndValues: []float64{1}}, + CollocationOptions{}); err == nil || !stringsContains(err, "twice") { + t.Fatalf("%s: repeated Start index: %v", name, err) + } + if _, err := SolveBoundaryCollocation(good3, 0, 1, mustFloats(t, []float64{0, 0.5, 1}), + BoundaryConditions{Start: []int{0}, End: []int{2, 2}, EndValues: []float64{1, 2}}, + CollocationOptions{}); err == nil || !stringsContains(err, "twice") { + t.Fatalf("%s: repeated End index: %v", name, err) + } + if _, err := SolveBoundaryCollocation(good, 0, 1, y0, + BoundaryConditions{}, CollocationOptions{}); err == nil || !stringsContains(err, "End must prescribe") { + t.Fatalf("%s: nothing prescribed at t1: %v", name, err) + } + // State and interval gates. + if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, []float64{1, 0, 2}), + bc, CollocationOptions{}); err == nil || !stringsContains(err, "in total") { + t.Fatalf("%s: three components with two conditions: %v", name, err) + } + if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, nil), bc, CollocationOptions{}); err == nil { + t.Fatalf("%s: empty state accepted", name) + } + if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, []float64{1, 0, 2}, 1, 3), bc, CollocationOptions{}); err == nil || !stringsContains(err, "vector") { + t.Fatalf("%s: a rank-2 state accepted", name) + } + if _, err := SolveBoundaryCollocation(good, 0, 0, y0, bc, CollocationOptions{}); err == nil || !stringsContains(err, "positive length") { + t.Fatalf("%s: an empty interval accepted", name) + } + if _, err := SolveBoundaryCollocation(good, 1, 0, y0, bc, CollocationOptions{}); err == nil || !stringsContains(err, "positive length") { + t.Fatalf("%s: a backward interval accepted", name) + } + if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, []float64{math.NaN(), 0}), + bc, CollocationOptions{}); err == nil || !stringsContains(err, "non-finite") { + t.Fatalf("%s: a NaN state accepted", name) + } + // Mesh gates. + if _, err := SolveBoundaryCollocation(good, 0, 1, y0, bc, + CollocationOptions{InitialNodes: 300, MaxNodes: 200}); err == nil || !stringsContains(err, "MaxNodes") { + t.Fatalf("%s: a starting mesh past MaxNodes accepted", name) + } + // A right-hand side of the wrong shape surfaces with its name. + bad := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 2, 3}, 3) + } + if _, err := SolveBoundaryCollocation(bad, 0, 1, y0, bc, CollocationOptions{}); err == nil || !stringsContains(err, "want a vector") { + t.Fatalf("%s: a wrong-shaped f accepted", name) + } + // Refinement past MaxNodes is refused with its name: the layer + // problem demands far more than 24 intervals at this tolerance. + layer := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), -100 * y.FloatAt(1)}, 2) + } + if _, err := SolveBoundaryCollocation(layer, 0, 1, mustFloats(t, []float64{0, 0}), + BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}}, + CollocationOptions{RelTol: 1e-4, AbsTol: 1e-6, InitialNodes: 8, MaxNodes: 24}); err == nil || !stringsContains(err, "MaxNodes") { + t.Fatalf("%s: refinement past MaxNodes accepted", name) + } +} diff --git a/integrate/odedae.go b/integrate/odedae.go new file mode 100644 index 0000000..6deb6cb --- /dev/null +++ b/integrate/odedae.go @@ -0,0 +1,331 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// Differential-algebraic equations in the semi-implicit mass-matrix +// form M·y' = f(t, y) with a singular M whose rank deficiency sits in +// whole zero rows (and, by the same counts, whole zero columns). The +// differential rows carry the implicit Euler scheme, the algebraic +// rows are constraints the implicit relation enforces at every step, +// and the coupled nonlinear system of each step is solved by Newton +// over the numerical Jacobian and the library's dense LU, the +// machinery the ODE stiff solvers already carry. +// +// The honest contract. The solver is first order and takes equal +// steps, the DAE twin of IntegrateBackwardEuler. Consistent initial +// values are the caller's contract: the solver verifies the initial +// residual on the algebraic rows and refuses when it sits beyond the +// tolerance, naming the row, but it does not project a general start +// onto the constraint manifold (a consistent initialiser is a +// documented follow-up). A second-order formula is not offered yet +// either: BDF2 keeps index 1 only when the algebraic start satisfies +// the constraint to second order, which again needs the missing +// projection step. And the index is certified at t0: the solver factors +// the Jacobian of the algebraic rows against the algebraic variables +// and refuses when it is singular, which is the textbook refutation of +// the index-1 contract and is exactly what catches the Cartesian +// pendulum with multipliers, an index-3 system whose per-step solves +// would otherwise converge while the trajectory drifted. + +// DAEOptions tunes IntegrateDAE. RelTol ≤ 0 means 1e-6, AbsTol ≤ 0 +// means 1e-9, the ODEOptions defaults; they size the Newton tolerance +// of the per-step solves. +type DAEOptions struct { + RelTol float64 + AbsTol float64 +} + +// IntegrateDAE integrates the mass-matrix differential-algebraic +// system M·y' = f(t, y) from t0 to t1 over the given number of equal +// implicit Euler steps and returns y(t1). The mass matrix must be +// square and singular, with its rank deficiency carried by whole zero +// rows matched by whole zero columns; the zero rows are the algebraic +// constraints and the zero columns the algebraic variables. Backward +// integration works: a t1 < t0 simply integrates in the negative +// direction. A wrong-shaped state or matrix, a nonsingular matrix, a +// rank deficiency outside whole zero rows, an inconsistent initial +// residual on an algebraic row, a singular algebraic block at t0 (the +// index-1 refutation), or a Newton iteration that cannot converge is +// an error, never a silently truncated trajectory. +func IntegrateDAE(f func(t float64, y *core.Array) (*core.Array, error), + m *core.Array, t0, t1 float64, y0 *core.Array, steps int, opts DAEOptions) (*core.Array, error) { + const name = "IntegrateDAE" + if steps <= 0 { + return nil, base.Errf("%s: steps must be ≥ 1, got %d", name, steps) + } + y, err := odeCheck(name, y0, nil) + if err != nil { + return nil, err + } + n := len(y) + mat, err := daeMassMatrix(name, m, n) + if err != nil { + return nil, err + } + relTol, absTol := opts.RelTol, opts.AbsTol + if relTol <= 0 { + relTol = 1e-6 + } + if absTol <= 0 { + absTol = 1e-9 + } + rows, cols, err := daeAlgebraicSets(name, mat) + if err != nil { + return nil, err + } + // The initial residual on the algebraic rows is the constraint the + // steps will enforce; a start that violates it has no trajectory. + fy0, err := odeEval(name, f, t0, y, n, nil) + if err != nil { + return nil, err + } + for _, row := range rows { + residual := fy0[row] + limit := 100 * (absTol + relTol*math.Abs(residual)) + if math.Abs(residual) > limit { + return nil, base.Errf("%s: the initial residual on algebraic row %d is %g, beyond the consistency tolerance %g: consistent initial values are the caller's contract", + name, row, residual, limit) + } + } + // The index-1 certificate: the algebraic rows must determine the + // algebraic variables, which is the nonsingularity of their + // Jacobian block. A singular block means the mass-matrix rank + // alone admitted index 1 but the system behaves above it. + w := &odeWork{} + jac, err := odeJacobian(name, f, t0, y, w) + if err != nil { + return nil, err + } + block := make([][]float64, len(rows)) + for i, r := range rows { + block[i] = make([]float64, len(cols)) + for j, c := range cols { + block[i][j] = jac[r*n+c] + } + } + base.Factor(block) + if err := base.CheckSingular(name, block); err != nil { + return nil, base.Errf("%s: %w: the Jacobian of the algebraic rows against the algebraic variables is singular at t=%g, the system behaves above index 1", + name, err, t0) + } + h := (t1 - t0) / float64(steps) + myn := make([]float64, n) + // One Newton result buffer serves every step: daeNewton overwrites + // it fully before the step copies it into the state, so no step + // allocates its own. + zbuf := make([]float64, n) + t := t0 + if odeArrived(t, t1) { + return arrayFromVector(y), nil + } + // 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 + daeMatVec(myn, mat, y) + if nerr := daeNewton(name, f, w, mat, tNext, h, myn, y, zbuf, absTol, relTol); nerr != nil { + return nil, nerr + } + copy(y, zbuf) + } + return arrayFromVector(y), nil +} + +// daeMassMatrix reads the mass matrix into a dense float64 row-major +// form and validates its shape against the state length. +func daeMassMatrix(name string, m *core.Array, n int) ([][]float64, error) { + if m.Dtype() == core.Complex { + return nil, base.Errf("%s: complex mass matrices are not supported", name) + } + shape := m.Shape() + if m.NDim() != 2 || shape[0] != n || shape[1] != n { + return nil, base.Errf("%s: the mass matrix must be square of the state's length %d, got shape %s", + name, n, base.ShapeText(shape)) + } + mat := make([][]float64, n) + for i := range n { + mat[i] = make([]float64, n) + for j := range n { + v := m.FloatAt(i*n + j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: the mass matrix carries the non-finite entry %g at row %d, column %d", name, v, i, j) + } + mat[i][j] = v + } + } + return mat, nil +} + +// daeAlgebraicSets returns the algebraic row and column indices: the +// zero rows and the zero columns of the mass matrix, verified to be +// equally numerous and to carry the whole rank deficiency. +func daeAlgebraicSets(name string, mat [][]float64) ([]int, []int, error) { + n := len(mat) + worst := 0.0 + for i := range n { + for j := range n { + worst = math.Max(worst, math.Abs(mat[i][j])) + } + } + tiny := worst * 1e-14 + var rows, cols []int + for i := range n { + zero := true + for j := range n { + if math.Abs(mat[i][j]) > tiny { + zero = false + break + } + } + if zero { + rows = append(rows, i) + } + } + for j := range n { + zero := true + for i := range n { + if math.Abs(mat[i][j]) > tiny { + zero = false + break + } + } + if zero { + cols = append(cols, j) + } + } + if len(rows) == 0 { + return nil, nil, base.Errf("%s: the mass matrix has no zero rows, want a singular mass matrix", name) + } + if len(rows) != len(cols) { + return nil, nil, base.Errf("%s: the mass matrix carries %d zero rows and %d zero columns, want equal counts for the semi-explicit index-1 form", + name, len(rows), len(cols)) + } + if rank := daeRank(mat, worst*1e-12); rank != n-len(rows) { + return nil, nil, base.Errf("%s: the mass matrix carries rank deficiency outside whole zero rows (rank %d with %d zero rows), which the index-1 contract refutes", + name, rank, len(rows)) + } + return rows, cols, nil +} + +// daeRank counts the pivots Gaussian elimination with partial pivoting +// leaves above the tolerance. +func daeRank(mat [][]float64, tol float64) int { + n := len(mat) + a := make([][]float64, n) + for i := range n { + a[i] = cloneDenseSlice(mat[i]) + } + rank, row := 0, 0 + for col := 0; col < n && row < n; col++ { + piv := row + for i := row + 1; i < n; i++ { + if math.Abs(a[i][col]) > math.Abs(a[piv][col]) { + piv = i + } + } + if math.Abs(a[piv][col]) <= tol { + continue + } + a[piv], a[row] = a[row], a[piv] + for i := row + 1; i < n; i++ { + f := a[i][col] / a[row][col] + for j := col; j < n; j++ { + a[i][j] -= f * a[row][j] + } + } + rank++ + row++ + } + return rank +} + +// daeMatVec writes the product m·x into dst. +func daeMatVec(dst []float64, m [][]float64, x []float64) { + for i := range dst { + s := 0.0 + for j, v := range m[i] { + s += v * x[j] + } + dst[i] = s + } +} + +// daeNewton solves the implicit Euler step equation M·z − M·y_n = +// h·f(tNext, z) for z by Newton over a numerical Jacobian and the +// library's LU solver, the mass-matrix twin of odeNewton. The +// converged state lands in dst, which must not alias seed: a driver +// that hands one buffer down every step allocates it once per run. +// The Jacobian is frozen from the seed and rebuilt twice when +// convergence drags; 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. An iteration that outlives twenty rounds, or +// a singular Newton matrix, surfaces as errNewtonStalled; an f that +// fails is the fatal error it is. +func daeNewton(name string, f func(t float64, y *core.Array) (*core.Array, error), + w *odeWork, m [][]float64, tNext, h float64, myn, seed, dst []float64, + absTol, relTol float64) error { + n := len(seed) + w.use(n) + z := dst + copy(z, seed) + mz := w.mz + for iteration := range 20 { + daeMatVec(mz, m, z) + 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] = mz[i] - h*w.fzs[i] - myn[i] + worst = math.Max(worst, math.Abs(w.g[i])) + terms = math.Max(terms, math.Abs(mz[i])+math.Abs(h*w.fzs[i])+math.Abs(myn[i])) + } + limit := math.Max(0.1*(absTol+relTol*normInfOfStep(z)), 8*base.EpsF*terms) + if worst <= limit { + return nil + } + if iteration == 0 || iteration == 4 || iteration == 10 { + jac, jerr := odeJacobian(name, f, tNext, z, w) + if jerr != nil { + return jerr + } + // Newton matrix M − 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] = m[i][j] - h*jac[i*n+j] + } + } + 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) +} diff --git a/integrate/odedae_test.go b/integrate/odedae_test.go new file mode 100644 index 0000000..f3be191 --- /dev/null +++ b/integrate/odedae_test.go @@ -0,0 +1,267 @@ +// Copyright (c) 2026 Petr Balvín (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" +) + +// daeCircuit returns the source, mass matrix and state of a linear +// index-1 circuit: a one-volt source feeds a unit resistor into node +// v1 (unit capacitor to ground), an inductor of one henry carries i +// on to node v2, and node v2 dumps through a unit resistor with no +// capacitor, so its KCL row 0 = i − v2 is the algebraic constraint +// and i the algebraic variable. With C = L = R = 1 the differential +// pair is x' = Ax + (1, 0) with A = [[−1, −1], [1, −1]], whose +// solution is elementary. +func daeCircuit(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{ + 1 - y.FloatAt(0) - y.FloatAt(2), + y.FloatAt(2) - y.FloatAt(1), + y.FloatAt(0) - y.FloatAt(1), + }, 3) +} + +func daeCircuitEnd(t *testing.T, steps int, y0 []float64, t0, t1 float64) []float64 { + t.Helper() + m := mustFloats(t, []float64{1, 0, 0, 0, 0, 0, 0, 0, 1}, 3, 3) + end, err := IntegrateDAE(daeCircuit, m, t0, t1, mustFloats(t, y0), steps, DAEOptions{}) + if err != nil { + t.Fatalf("IntegrateDAE: %v", err) + } + return []float64{end.FloatAt(0), end.FloatAt(1), end.FloatAt(2)} +} + +// daeCircuitExact evaluates the exact v1, v2, i at time t: the +// equilibrium (0.5, 0.5) plus the elementary homogeneous part. +func daeCircuitExact(t float64) []float64 { + c := math.Exp(-t) / 2 + return []float64{ + 0.5 + c*(math.Cos(t)+math.Sin(t)), + 0.5 + c*(math.Sin(t)-math.Cos(t)), + 0.5 + c*(math.Sin(t)-math.Cos(t)), + } +} + +// TestIntegrateDAECircuit is the linear index-1 pin: the differential +// nodes track the elementary solution, and the algebraic variable i +// satisfies the KCL constraint i = v2 to rounding at the end state, +// because every step enforces the constraint row exactly. +func TestIntegrateDAECircuit(t *testing.T) { + end := daeCircuitEnd(t, 200, []float64{1, 0, 0}, 0, 1) + want := daeCircuitExact(1) + for k, band := range []float64{0.01, 0.01, 0.01} { + if math.Abs(end[k]-want[k]) > band { + t.Fatalf("circuit[%d] = %.14g, want %.14g ± %g", k, end[k], want[k], band) + } + } + if math.Abs(end[1]-end[2]) > 1e-10 { + t.Fatalf("the constraint i = v2 drifted to %g at the end state", end[1]-end[2]) + } +} + +// TestIntegrateDAEScalarConstraint pins the algebraic variable on the +// exact constraint to rounding: with w' unconstrained by M's zero row +// and 0 = w − cos t, the solved w must equal cos at every step, so +// certainly at the end. +func TestIntegrateDAEScalarConstraint(t *testing.T) { + f := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-y.FloatAt(0), y.FloatAt(1) - math.Cos(t)}, 2) + } + m := mustFloats(t, []float64{1, 0, 0, 0}, 2, 2) + end, err := IntegrateDAE(f, m, 0, 1, mustFloats(t, []float64{1, 1}), 25, DAEOptions{}) + if err != nil { + t.Fatalf("IntegrateDAE: %v", err) + } + if math.Abs(end.FloatAt(1)-math.Cos(1)) > 1e-12 { + t.Fatalf("algebraic w(1) = %.16g, want cos(1) = %.16g to rounding", + end.FloatAt(1), math.Cos(1)) + } + if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 0.05 { + t.Fatalf("differential u(1) = %.14g, want %.14g ± 0.05", end.FloatAt(0), math.Exp(-1)) + } +} + +// TestIntegrateDAEBackward integrates the circuit backwards from the +// exact end state; the signed-step formulation must return the start. +func TestIntegrateDAEBackward(t *testing.T) { + want := daeCircuitExact(1) + end := daeCircuitEnd(t, 200, want, 1, 0) + start := daeCircuitExact(0) + for k := range 3 { + if math.Abs(end[k]-start[k]) > 0.01 { + t.Fatalf("backward circuit[%d] = %.14g, want %.14g ± 0.01", k, end[k], start[k]) + } + } +} + +// TestIntegrateDAEConsistencyRefused pins the initial-residual check: +// a start violating the KCL row by one full unit is refused, with the +// row named. +func TestIntegrateDAEConsistencyRefused(t *testing.T) { + m := mustFloats(t, []float64{1, 0, 0, 0, 0, 0, 0, 0, 1}, 3, 3) + _, err := IntegrateDAE(daeCircuit, m, 0, 1, mustFloats(t, []float64{1, 1, 0}), 200, DAEOptions{}) + if err == nil { + t.Fatal("expected an error for an inconsistent start") + } + if !strings.Contains(err.Error(), "row 1") { + t.Fatalf("want the algebraic row named, got %v", err) + } +} + +// TestIntegrateDAEPendulumRefused pins the honest index-3 refusal: the +// Cartesian pendulum with multipliers has a mass matrix that admits +// index 1 by rank alone, but its algebraic rows (the constraints) do +// not depend on the algebraic variables (the multipliers) at all, so +// the certified block is singular and the solver refuses, naming the +// detection. +func TestIntegrateDAEPendulumRefused(t *testing.T) { + // y = (x, ypos, u, v, lambda, mu): position, velocity, multipliers. + // m = 1, g = 0, length 1: the unit circle, started at (1, 0) with + // unit tangential speed and the multiplier that holds it there. + f := func(t float64, y *core.Array) (*core.Array, error) { + x, ypos, u, v, lambda := y.FloatAt(0), y.FloatAt(1), y.FloatAt(2), y.FloatAt(3), y.FloatAt(4) + return core.FromFloats([]float64{ + u, v, + -2 * x * lambda, + -2 * ypos * lambda, + x*x + ypos*ypos - 1, + x*u + ypos*v, + }, 6) + } + m := mustFloats(t, []float64{ + 1, 0, 0, 0, 0, 0, + 0, 1, 0, 0, 0, 0, + 0, 0, 1, 0, 0, 0, + 0, 0, 0, 1, 0, 0, + 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, + }, 6, 6) + y0 := mustFloats(t, []float64{1, 0, 0, 1, 0.5, 0}) + _, err := IntegrateDAE(f, m, 0, 0.1, y0, 10, DAEOptions{}) + if err == nil { + t.Fatal("expected the index-3 pendulum to be refused") + } + if !strings.Contains(err.Error(), "index 1") || !strings.Contains(err.Error(), "singular") { + t.Fatalf("want the index detection stated, got %v", err) + } +} + +// TestIntegrateDAEIndexTwoStall pins the other honest refusal: an +// ordinary stiff ODE whose per-step Newton matrix is exactly singular +// at the chosen step (y0' = 100·y0 with h·100 = 1) fails loudly +// through the Newton solve, not through silent drift; the index-1 +// certificate itself passes because the algebraic row y1 − y0 does +// depend on the algebraic variable, so the refusal here comes from +// the differential row's pathology and must be named as such. +func TestIntegrateDAEIndexTwoStall(t *testing.T) { + // y0' = 100·y0 with h·100 = 1 makes the Newton matrix singular. + f := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{100 * y.FloatAt(0), y.FloatAt(1) - y.FloatAt(0)}, 2) + } + m := mustFloats(t, []float64{1, 0, 0, 0}, 2, 2) + _, err := IntegrateDAE(f, m, 0, 1, mustFloats(t, []float64{1, 1}), 100, DAEOptions{}) + if err == nil { + t.Fatal("expected the singular per-step solve to be refused") + } + if !strings.Contains(err.Error(), "Newton") { + t.Fatalf("want a Newton failure, got %v", err) + } +} + +// TestIntegrateDAEErrors pins the structural error contract: a zero +// step count, a rank-1 mass matrix of the wrong shape, a nonsingular +// matrix, a rank deficiency without whole zero rows, mismatched zero +// row and column counts, an empty state, a non-finite matrix entry and +// a failing f are all errors; a degenerate span returns the start. +func TestIntegrateDAEErrors(t *testing.T) { + good := mustFloats(t, []float64{1, 0, 0, 0}, 2, 2) + simple := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-y.FloatAt(0), y.FloatAt(1)}, 2) + } + simple3 := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-y.FloatAt(0), y.FloatAt(1), -y.FloatAt(2)}, 3) + } + if _, err := IntegrateDAE(simple, good, 0, 1, mustFloats(t, []float64{1, 1}), 0, DAEOptions{}); err == nil { + t.Fatal("expected an error for zero steps") + } + badShape := mustFloats(t, []float64{1, 0}, 1, 2) + if _, err := IntegrateDAE(simple, badShape, 0, 1, mustFloats(t, []float64{1, 1}), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for a non-square mass matrix") + } + identity := mustFloats(t, []float64{1, 0, 0, 1}, 2, 2) + if _, err := IntegrateDAE(simple, identity, 0, 1, mustFloats(t, []float64{1, 1}), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for a nonsingular mass matrix") + } + noZeroRows := mustFloats(t, []float64{1, 1, 1, 1}, 2, 2) + if _, err := IntegrateDAE(simple, noZeroRows, 0, 1, mustFloats(t, []float64{1, 1}), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for rank deficiency without zero rows") + } + rowColMismatch := mustFloats(t, []float64{1, 1, 0, 0}, 2, 2) + if _, err := IntegrateDAE(simple, rowColMismatch, 0, 1, mustFloats(t, []float64{1, 1}), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for mismatched zero row and column counts") + } + if _, err := IntegrateDAE(simple, good, 0, 1, mustFloats(t, nil), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for an empty state") + } + nonFinite := mustFloats(t, []float64{1, 0, 0, math.Inf(1)}, 2, 2) + if _, err := IntegrateDAE(simple, nonFinite, 0, 1, mustFloats(t, []float64{1, 1}), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for a non-finite mass matrix entry") + } + complexM, _ := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + if _, err := IntegrateDAE(simple, complexM, 0, 1, mustFloats(t, []float64{1, 1}), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for a complex mass matrix") + } + // One zero row but a second dependent row: the rank deficiency + // exceeds the zero rows and the contract is refused. + hiddenDeficiency := mustFloats(t, []float64{1, 1, 0, 1, 1, 0, 0, 0, 0}, 3, 3) + _, err := IntegrateDAE(simple3, hiddenDeficiency, 0, 1, mustFloats(t, []float64{1, 1, 1}), 10, DAEOptions{}) + if err == nil || !strings.Contains(err.Error(), "rank deficiency") { + t.Fatalf("expected the hidden rank deficiency to be refused, got %v", err) + } + // A degenerate span answers the validated start unchanged. The + // start must satisfy the algebraic row of this system, y1 = 0. + same, err := IntegrateDAE(simple, good, 1, 1, mustFloats(t, []float64{1, 0}), 10, DAEOptions{}) + if err != nil { + t.Fatalf("zero span: %v", err) + } + if same.FloatAt(0) != 1 || same.FloatAt(1) != 0 { + t.Fatalf("zero span moved the state to (%v, %v)", same.FloatAt(0), same.FloatAt(1)) + } + boom := func(t float64, y *core.Array) (*core.Array, error) { + if t > 0.5 { + return nil, base.Errf("detector tripped") + } + return core.FromFloats([]float64{-y.FloatAt(0), y.FloatAt(1)}, 2) + } + if _, err := IntegrateDAE(boom, good, 0, 1, mustFloats(t, []float64{1, 0}), 100, DAEOptions{}); err == nil { + t.Fatal("expected the operator error to propagate") + } + // An f failing on the initial evaluation, the initial Jacobian and + // inside the first Newton iteration is refused at once. + always := func(t float64, y *core.Array) (*core.Array, error) { + return nil, base.Errf("detector tripped") + } + if _, err := IntegrateDAE(always, good, 0, 1, mustFloats(t, []float64{1, 0}), 10, DAEOptions{}); err == nil { + t.Fatal("expected an error for an f that always fails") + } + // An f that only tolerates the exact seed fails when the Newton + // iteration perturbs the state for its numerical Jacobian. + touchy := func(t float64, y *core.Array) (*core.Array, error) { + if y.FloatAt(0) != 1 { + return nil, base.Errf("detector tripped") + } + return core.FromFloats([]float64{1 - y.FloatAt(0), y.FloatAt(1) - y.FloatAt(0)}, 2) + } + if _, err := IntegrateDAE(touchy, good, 0, 1, mustFloats(t, []float64{1, 0}), 10, DAEOptions{}); err == nil { + t.Fatal("expected the Jacobian perturbation to trip the f error") + } +} diff --git a/integrate/odeevent.go b/integrate/odeevent.go new file mode 100644 index 0000000..cb1211e --- /dev/null +++ b/integrate/odeevent.go @@ -0,0 +1,208 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import ( + "cmp" + "math" + "slices" +) + +// Event detection along an ODE trajectory. Impact times, resonance +// crossings and threshold passages are all the same question: when +// does a scalar watch function g(t, y) cross zero on the way to t1? +// The integrator checks every accepted step for a sign change of each +// watch and, when one appears, narrows the crossing by bisection that +// re-integrates the step's interval, so the event time is as accurate +// as the integrator itself and needs no dense-output machinery. +// +// Watches are only compared across the boundaries of accepted steps: +// a watch that touches zero and returns to its sign inside one step +// goes unnoticed, the same blind spot every step-based detector has. +// Sign changes are searched strictly after t0, so a watch sitting on +// zero at the initial state does not fire until it leaves and returns. + +// ODEWatch reports a scalar quantity to watch along the trajectory. +// Direction filters the crossings: +1 records only rising crossings +// (watch going from negative to non-negative), −1 only falling ones, +// 0 both. +type ODEWatch struct { + Function func(t float64, y *core.Array) (float64, error) + Direction int +} + +// ODEEventHit records one zero crossing: the time it happened, the +// state there, and which watch fired. +type ODEEventHit struct { + Time float64 + State *core.Array + Watch int + Rising bool +} + +// IntegrateODEEvents integrates y' = f(t, y) from t0 to t1 exactly +// like IntegrateODE and, alongside the final state, returns every +// zero crossing of the watches, sorted by time. The watch functions +// must tolerate being called at any time inside [t0, t1]; they are +// also called at the step boundaries the integrator accepts, and the +// first accepted step is seeded with the watch value at its start +// state, so a crossing inside it is detected like any other. A watch +// error, a non-finite watch value included, aborts the integration. +func IntegrateODEEvents(f func(t float64, y *core.Array) (*core.Array, error), + t0, t1 float64, y0 *core.Array, watches []ODEWatch, opts ODEOptions) ([]ODEEventHit, *core.Array, error) { + if len(watches) == 0 { + return nil, nil, base.Errf("IntegrateODEEvents: at least one watch is needed") + } + for i := range watches { + if watches[i].Function == nil { + return nil, nil, base.Errf("IntegrateODEEvents: watch %d has no function", i) + } + } + // gPrev[i] carries the watch value at the last accepted boundary; + // havePrev becomes true after the first evaluation. + gPrev := make([]float64, len(watches)) + havePrev := false + var hits []ODEEventHit + + watch := func(tPrev, tNow float64, yPrev, yNow []float64) (bool, error) { + for i := range watches { + if !havePrev { + // First accepted step: the watch value at the step's + // start state seeds the comparison, so a crossing + // inside the very first step is seen like any other + // one instead of hiding behind a missing gPrev. + g0, err := watches[i].Function(tPrev, wrapVector(yPrev)) + if err != nil { + return false, base.Errf("IntegrateODEEvents: watch %d: %w", i, err) + } + if math.IsNaN(g0) || math.IsInf(g0, 0) { + return false, base.Errf("IntegrateODEEvents: watch %d returned the non-finite value %g at t=%g", i, g0, tPrev) + } + gPrev[i] = g0 + } + g, err := watches[i].Function(tNow, wrapVector(yNow)) + if err != nil { + return false, base.Errf("IntegrateODEEvents: watch %d: %w", i, err) + } + // A non-finite watch value compares false against every + // sign test and would masquerade as a crossing (or swallow + // one), so it is refused like quad's non-finite integrand. + if math.IsNaN(g) || math.IsInf(g, 0) { + return false, base.Errf("IntegrateODEEvents: watch %d returned the non-finite value %g at t=%g", i, g, tNow) + } + gp := gPrev[i] + if gp == 0 { + gPrev[i] = g + continue + } + if g == 0 { + // The watch landed exactly on zero at the accepted + // boundary: the crossing is here, at (tNow, yNow), no + // bisection needed. The leaving interval starts from + // zero, so the next step stays silent, exactly the + // "does not fire until it leaves and returns" the + // documented zero-at-start rule spells out. The clamp + // h = t1 − t makes the final boundary land here for + // round numbers, so dropping it would lose events. + rising := gp < 0 + if watches[i].Direction > 0 && !rising || watches[i].Direction < 0 && rising { + gPrev[i] = g + continue + } + state := make([]float64, len(yNow)) + copy(state, yNow) + hits = append(hits, ODEEventHit{ + Time: tNow, + State: arrayFromVector(state), + Watch: i, + Rising: rising, + }) + gPrev[i] = g + continue + } + if (gp < 0) == (g < 0) { + gPrev[i] = g + continue + } + // A crossing between tPrev and tNow: bisect on the + // re-integrated watch from the step's start state. + rising := gp < 0 + if watches[i].Direction != 0 { + if watches[i].Direction > 0 && !rising { + gPrev[i] = g + continue + } + if watches[i].Direction < 0 && rising { + gPrev[i] = g + continue + } + } + tHit, yHit, err := refineEvent(f, tPrev, tNow, yPrev, gPrev[i], watches[i].Function, opts) + if err != nil { + return false, base.Errf("IntegrateODEEvents: %w", err) + } + hits = append(hits, ODEEventHit{ + Time: tHit, + State: yHit, + Watch: i, + Rising: rising, + }) + gPrev[i] = g + } + havePrev = true + return false, nil + } + final, err := odeRun(f, t0, t1, y0, opts, watch) + if err != nil { + return nil, nil, err + } + slices.SortFunc(hits, func(a, b ODEEventHit) int { + return cmp.Compare(a.Time, b.Time) + }) + return hits, final, nil +} + +// refineEvent narrows a zero crossing of g between tPrev and tNow by +// bisection, evaluating g by re-integrating from the step's start +// state. Both endpoint values are known to have opposite signs, which +// bisection turns into the crossing at integrator accuracy. +func refineEvent(f func(t float64, y *core.Array) (*core.Array, error), + tPrev, tNow float64, yPrev []float64, gPrev float64, + g func(t float64, y *core.Array) (float64, error), opts ODEOptions) (float64, *core.Array, error) { + lo, hi := tPrev, tNow + for range 100 { + mid := (lo + hi) / 2 + if mid == lo || mid == hi { + break + } + yMid, err := IntegrateODE(f, tPrev, mid, wrapVector(yPrev), opts) + if err != nil { + return 0, nil, err + } + gm, err := g(mid, yMid) + if err != nil { + return 0, nil, err + } + if gm == 0 { + return mid, yMid, nil + } + if (gPrev < 0) == (gm < 0) { + lo = mid + gPrev = gm + } else { + hi = mid + } + } + tHit := (lo + hi) / 2 + yHit, err := IntegrateODE(f, tPrev, tHit, wrapVector(yPrev), opts) + if err != nil { + return 0, nil, err + } + return tHit, yHit, nil +} diff --git a/integrate/odegrid_pin_test.go b/integrate/odegrid_pin_test.go new file mode 100644 index 0000000..d371aef --- /dev/null +++ b/integrate/odegrid_pin_test.go @@ -0,0 +1,41 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestFixedStepTimeGrid pins the exact time grid of the fixed-step +// solvers: the stage times come from t0 + i·h, never from an +// accumulated t += h, whose rounding walks over a long run. On +// y' = cos t from t0 = 1e6 the walk used to hold the answer at an +// error of 1.1e-8 that no step count refined away; the grid puts RK4 +// at its own truncation floor and lets the first-order schemes show +// their h. The tolerances sit an order below the old walk, so the pin +// fails on an accumulated-time build. +func TestFixedStepTimeGrid(t *testing.T) { + const t0, span, steps = 1e6, 1e-1, 1e5 + cos := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{math.Cos(t)}, 1) + } + exact := math.Sin(t0 + span) + end, err := IntegrateRK4(cos, t0, t0+span, mustFloats(t, []float64{math.Sin(t0)}), steps) + if err != nil { + t.Fatalf("IntegrateRK4: %v", err) + } + if got := math.Abs(end.FloatAt(0) - exact); got > 1e-11 { + t.Fatalf("RK4 error %.3e exceeds 1e-11: the stage times have left the exact grid", got) + } + end, err = IntegrateBackwardEuler(cos, t0, t0+span, mustFloats(t, []float64{math.Sin(t0)}), 1e6, ODEOptions{}) + if err != nil { + t.Fatalf("IntegrateBackwardEuler: %v", err) + } + if got := math.Abs(end.FloatAt(0) - exact); got > 5e-9 { + t.Fatalf("backward Euler error %.3e exceeds 5e-9: the step times have left the exact grid", got) + } +} diff --git a/integrate/oderow.go b/integrate/oderow.go new file mode 100644 index 0000000..ad8cdbb --- /dev/null +++ b/integrate/oderow.go @@ -0,0 +1,264 @@ +// Copyright (c) 2026 Petr Balvín (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 Rosenbrock-Wanner workhorse: the four-stage L-stable scheme ROS4 +// of Hairer and Wanner, fourth order accurate with an embedded +// third-order solution driving the step control. Where BDF solves each +// step by Newton, a W method puts the Jacobian into the formula +// itself: every stage is one linear solve against the frozen matrix +// (1/(γh))I − J, so a step costs one numerical Jacobian, one LU +// factorisation and four back-substitutions, and no Newton iteration +// ever stalls. The scheme is L-stable: its stability function vanishes +// at infinity, so a step far beyond the transient's time constant +// damps the stiff mode instead of amplifying it. +// +// The stage form is the divided one the standard implementations use. +// With A = (1/(γh))I − J factored once per step, stage one solves +// A·k₁ = f(t, y) and stage i solves A·kᵢ = f(t + cᵢh, yᵢ) + +// (Σⱼ<ᵢ cᵢⱼkⱼ)/h with the stage value yᵢ = y + Σⱼ<ᵢ aᵢⱼkⱼ; the step +// advances by y + Σ bᵢkᵢ and the embedded estimate is Σ êᵢkᵢ. Stage +// four shares its node and its stage value with stage three, so its f +// evaluation carries over and a step costs three f calls. The digits +// are the published ones, cross-checked against the standard +// implementations of the scheme. +// +// Two honest limits of the tableau. The Jacobian enters at (t, y) and +// the tableau's time-derivative weights are left out, because the +// package's f contract carries no partial derivative in t: the fourth +// order therefore holds for autonomous systems (every problem in this +// package's stiff tests is one), while a genuinely time-dependent f +// loses the second-order local terms those weights carry and can +// degrade to first order globally; measured on y' = −y + t with +// uniform steps the global ratios come out near 2, not near 16. And +// the scheme is L-stable but not stiffly accurate: the fourth stage's +// weights differ from the solution weights, the stiff damping coming +// from the stability limit rather than from stage-end agreement. + +var ( + // rowGamma is the abscissa of the stage solves and the source of + // the L-stability: the stability function's denominator is + // (1 − γz)⁴ and its numerator vanishes at infinity. + rowGamma = 0.57282 + + // rowNodes are the stage times as multiples of h. The second node + // is the published tableau's c2 = 0.114564 (Hairer and Wanner's + // ROS4); its doubled form 2·gamma appears in some reprints, and the + // digit never enters the arithmetic of an autonomous problem, where + // stage times cancel, so either form integrates identically here. + rowNodes = [4]float64{0, 0.114564, 0.65521686381559, 0.65521686381559} + + // rowA holds the stage-value weights a[i][j], the coefficient of + // k_j inside stage i's value. + rowA = [4][3]float64{ + {}, + {2}, + {1.867943637803922, 0.2344449711399156}, + {1.867943637803922, 0.2344449711399156, 0}, + } + + // rowC holds the Jacobian-coupling weights c[i][j], the coefficient + // of k_j inside stage i's right side, divided by h. + rowC = [4][3]float64{ + {}, + {-7.13761503641231}, + {2.580708087951457, 0.6515950076447975}, + {-2.137148994382534, -0.3214669691237626, -0.6949742501781779}, + } + + // rowB advances the solution; rowE carries the embedded third-order + // estimate (the difference between the fourth-order and the + // embedded third-order weights). + rowB = [4]float64{2.255570073418735, 0.2870493262186792, 0.435317943184018, 1.093502252409163} + rowE = [4]float64{-0.2815431932141155, -0.0727619912493892, -0.1082196201495311, -1.093502252409163} +) + +// IntegrateROS4 integrates y' = f(t, y) from t0 to t1 with the +// four-stage L-stable Rosenbrock-Wanner scheme ROS4 and returns y(t1). +// The fourth order holds for autonomous systems; a genuinely +// time-dependent f loses the second-order local terms the tableau's +// time-derivative weights would carry (the package's f contract has no +// partial derivative in t), and can degrade to first order globally: +// measured on y' = −y + t with uniform steps the ratios come out near +// 2, not near 16. The adaptive controller still holds its tolerance +// there, at the cost of more steps. The step size follows the embedded +// error under the classic accept-or-shrink control, with the numerical +// Jacobian taken once per step through the same central-difference +// helper the Newton path uses; there is no user Jacobian parameter. +// 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 stage that leaves the +// finite range even as the step shrinks is an error, never a silently +// truncated trajectory. +func IntegrateROS4(f func(t float64, y *core.Array) (*core.Array, error), + t0, t1 float64, y0 *core.Array, opts ODEOptions) (*core.Array, error) { + const name = "IntegrateROS4" + y, err := odeCheck(name, y0, &opts) + if err != nil { + return nil, err + } + h, err := bdf2InitialStep(name, f, t0, t1, y, &opts) + if err != nil { + return nil, err + } + budget := odeBudget{max: opts.MaxSteps} + w := &odeWork{} + t := t0 + 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) + yEnd, errNorm, nerr := ros4Step(name, f, w, t, y, h, opts.AbsTol, opts.RelTol) + if nerr != nil { + if errors.Is(nerr, errNewtonStalled) { + // The stage matrix closed onto singularity: halve the + // step and retry the same interval, within the budget. + h *= 0.5 + continue + } + return nil, nerr + } + if math.IsNaN(errNorm) || math.IsInf(errNorm, 0) { + // A stage left the finite range: shrink hard and retry. + h *= 0.25 + continue + } + factor := min(5, max(0.2, 0.9*math.Pow(1/errNorm, 1.0/4))) + if errNorm <= 1 { + copy(y, yEnd) + prevT := t + t += h + 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) + } + } else { + // Rejected: retry the same interval with the smaller step. + h *= max(0.2, factor) + } + } + return arrayFromVector(y), nil +} + +// ros4Step attempts one ROS4 step of the given size from (t, y) and +// returns the candidate end state, the embedded error norm against the +// mixed absolute and relative tolerance, and a nil error when the +// stages all solved. Every buffer belongs to the workspace, so a step +// allocates nothing of its own: the caller's y is left untouched, and +// the returned slice stays valid until the next step. +func ros4Step(name string, f func(t float64, y *core.Array) (*core.Array, error), + w *odeWork, t float64, y []float64, h, absTol, relTol float64) ([]float64, float64, error) { + n := len(y) + w.use(n) + // One scratch buffer per role: the stage value, the stage right + // side, f's result and the candidate end state are rebuilt at the + // top of their use and read only by it. + ks, stage, rhs, fy, yEnd := w.ks, w.stage, w.rhs, w.fy, w.yEnd + // One numerical Jacobian per step, the house central-difference + // helper the Newton path shares, and one factorisation of I − hγJ, + // which is (1/(γh))I − J scaled by γh: the scale folds into the + // stage right sides instead of the matrix. + jac, jerr := odeJacobian(name, f, t, y, w) + if jerr != nil { + return nil, 0, jerr + } + lu := w.mat + // The step-scaled matrix weight and the stage weight are one + // product each, the same (−h·γ) and (h·γ) groupings the element + // loops evaluated. + hgr := -h * rowGamma + hg := h * rowGamma + for i := range n { + row := lu[i] + for j := range n { + row[j] = hgr * jac[i*n+j] + } + row[i]++ + } + w.perm, _ = base.Factor(lu) + if err := base.CheckSingular(name, lu); err != nil { + return nil, 0, base.Errf("%s: %w, singular stage matrix at t=%g", name, errNewtonStalled, t) + } + // Stage one. + out, ferr := odeCall(name, f, t, y, n, &w.views) + if ferr != nil { + return nil, 0, ferr + } + readVector(fy, out) + for i := range n { + rhs[i] = hg * fy[i] + } + odePermuteColumn(rhs, w.perm, w.visited) + base.SolveColumn(lu, rhs) + copy(ks[0], rhs) + // Stages two through four. Stage four repeats stage three's node + // and stage value, so its derivative carries over. + for s := 1; s < 4; s++ { + if s < 3 { + copy(stage, y) + for j := range s { + aj := rowA[s][j] + if aj == 0 { + continue + } + kjs := ks[j] + for i := range n { + stage[i] += aj * kjs[i] + } + } + out, ferr = odeCall(name, f, t+rowNodes[s]*h, stage, n, &w.views) + if ferr != nil { + return nil, 0, ferr + } + readVector(fy, out) + } + for i := range n { + rhs[i] = hg * fy[i] + } + for j := range s { + cj := rowC[s][j] + if cj == 0 { + continue + } + gcj := rowGamma * cj + kjs := ks[j] + for i := range n { + rhs[i] += gcj * kjs[i] + } + } + odePermuteColumn(rhs, w.perm, w.visited) + base.SolveColumn(lu, rhs) + copy(ks[s], rhs) + } + // The embedded pair: the b weights advance, the gap to the + // embedded third-order solution estimates the local error. + errNorm := 0.0 + for i := range n { + e, advance := 0.0, 0.0 + for s := range 4 { + e += rowE[s] * ks[s][i] + advance += rowB[s] * ks[s][i] + } + yEnd[i] = y[i] + advance + scale := absTol + relTol*math.Max(math.Abs(y[i]), math.Abs(yEnd[i])) + ratio := e / scale + errNorm += ratio * ratio + } + return yEnd, math.Sqrt(errNorm/float64(n)) + 1e-10, nil +} diff --git a/integrate/oderow_test.go b/integrate/oderow_test.go new file mode 100644 index 0000000..860e1d8 --- /dev/null +++ b/integrate/oderow_test.go @@ -0,0 +1,266 @@ +// Copyright (c) 2026 Petr Balvín (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" +) + +// TestROS4FixedStepOrder pins the fourth order of the scheme: on the +// oscillator, whose exact rotation is known, uniform steps must shrink +// the global error by roughly sixteen per halving. The convergence is +// measured against the driven step, the way the order is defined. +func TestROS4FixedStepOrder(t *testing.T) { + oscillator := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2) + } + errAt := func(steps int) float64 { + h := 1.0 / float64(steps) + y := []float64{1, 0} + now := 0.0 + w := &odeWork{} + for range steps { + yEnd, _, err := ros4Step("TestROS4FixedStepOrder", oscillator, w, now, y, h, 1e-300, 1e-300) + if err != nil { + t.Fatalf("ros4Step: %v", err) + } + copy(y, yEnd) + now += h + } + return math.Max(math.Abs(y[0]-math.Cos(1)), math.Abs(y[1]+math.Sin(1))) + } + e8, e16, e32 := errAt(8), errAt(16), errAt(32) + if e8 < 1e-13 { + t.Skipf("error already at round-off (%v)", e8) + } + for _, r := range []float64{e8 / e16, e16 / e32} { + if r < 12 || r > 20 { + t.Fatalf("error ratio over a halved step = %.2g, want ≈ 16 for a fourth-order scheme", r) + } + } +} + +// TestROS4NonAutonomousDegradation pins the documented limit the +// autonomous fourth order carries with it: on y' = −y + t, whose exact +// answer y = t − 1 + e^{−t} is known, the tableau's missing +// time-derivative weights cost the second-order local terms and the +// uniform-step ratios sit near 2, first order, not near 16. A change +// that lifts this must move the doc comment with it. +func TestROS4NonAutonomousDegradation(t *testing.T) { + forced := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-y.FloatAt(0) + t}, 1) + } + errAt := func(steps int) float64 { + h := 1.0 / float64(steps) + y := []float64{0} + now := 0.0 + w := &odeWork{} + for range steps { + yEnd, _, err := ros4Step("TestROS4NonAutonomousDegradation", forced, w, now, y, h, 1e-300, 1e-300) + if err != nil { + t.Fatalf("ros4Step: %v", err) + } + copy(y, yEnd) + now += h + } + return math.Abs(y[0] - (1 - 1/math.E)) + } + e8, e16, e32 := errAt(8), errAt(16), errAt(32) + for _, r := range []float64{e8 / e16, e16 / e32} { + if r >= 4 { + t.Fatalf("error ratio over a halved step = %.2g, the forced system runs at first order (the doc comment names this limit)", r) + } + } +} + +// TestIntegrateROS4Quadrature pins exactness on y' = 1: the constant +// right side is reproduced to rounding whatever the accepted steps do. +func TestIntegrateROS4Quadrature(t *testing.T) { + one := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1}, 1) + } + end, err := IntegrateROS4(one, 0, 1, mustFloats(t, []float64{0}), ODEOptions{}) + if err != nil { + t.Fatalf("IntegrateROS4: %v", err) + } + if math.Abs(end.FloatAt(0)-1) > 1e-12 { + t.Fatalf("y(1) = %.16g, want 1 to rounding", end.FloatAt(0)) + } +} + +// TestIntegrateROS4Accuracy checks the adaptive driver on the analytic +// decay and over a full oscillator period with a two-dimensional state. +func TestIntegrateROS4Accuracy(t *testing.T) { + end, err := IntegrateROS4(decay, 0, 1, mustFloats(t, []float64{1}), + ODEOptions{RelTol: 1e-8, AbsTol: 1e-12}) + if err != nil { + t.Fatalf("IntegrateROS4: %v", err) + } + if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-7 { + t.Fatalf("y(1) = %.14g, want %.14g ± 1e-7", 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 := IntegrateROS4(oscillator, 0, 2*math.Pi, mustFloats(t, []float64{1, 0}), + ODEOptions{RelTol: 1e-8, AbsTol: 1e-12}) + if err != nil { + t.Fatalf("IntegrateROS4 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 := IntegrateROS4(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}), + ODEOptions{RelTol: 1e-8, AbsTol: 1e-12}) + if err != nil { + t.Fatalf("IntegrateROS4 backward: %v", err) + } + if math.Abs(back.FloatAt(0)-1) > 1e-7 { + t.Fatalf("backward y(0) = %.14g, want 1 ± 1e-7", back.FloatAt(0)) + } +} + +// TestIntegrateROS4LStability pins the L-stable damping: on y' = +// −10^8(y − 1) and y' = −10^8 y the steps are far beyond the transient +// and a non-L-stable scheme blows up, while the W scheme lands on the +// forcing, respectively on zero, with a bounded step count. +func TestIntegrateROS4LStability(t *testing.T) { + rise := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-1e8 * (y.FloatAt(0) - 1)}, 1) + } + end, err := IntegrateROS4(rise, 0, 1, mustFloats(t, []float64{0}), ODEOptions{MaxSteps: 1000}) + if err != nil { + t.Fatalf("IntegrateROS4 stiff rise: %v", err) + } + if v := end.FloatAt(0); math.IsNaN(v) || math.Abs(v-1) > 1e-9 { + t.Fatalf("stiff rise y(1) = %.14g, want 1", v) + } + decayStiff := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-1e8 * y.FloatAt(0)}, 1) + } + end, err = IntegrateROS4(decayStiff, 0, 1, mustFloats(t, []float64{1}), ODEOptions{MaxSteps: 1000}) + if err != nil { + t.Fatalf("IntegrateROS4 stiff decay: %v", err) + } + if v := end.FloatAt(0); math.IsNaN(v) || math.Abs(v) > 1e-9 { + t.Fatalf("stiff decay y(1) = %.14g, want 0", v) + } + // The one-step amplification at a stiff eigenvalue must die out: + // a giant step on the decay damps the state by orders of magnitude + // instead of amplifying it. + yEnd, _, err := ros4Step("TestIntegrateROS4LStability", decay, &odeWork{}, 0, []float64{1}, 1e6, 1e-3, 1e-3) + if err != nil { + t.Fatalf("ros4Step at z = 1e6: %v", err) + } + if math.Abs(yEnd[0]) > 1e-4 { + t.Fatalf("one step at h·λ = 1e6 multiplied the state by %.3g, want heavy damping", yEnd[0]) + } +} + +// TestIntegrateROS4VanDerPol integrates the Van der Pol oscillator in +// the stiff relaxation regime: μ = 1000 carries a transient of width +// 10^−3 under a slow motion, and the W scheme must cross it and follow +// the slow branch inside the step budget. +func TestIntegrateROS4VanDerPol(t *testing.T) { + const mu = 1000.0 + vdp := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{ + y.FloatAt(1), + mu*(1-y.FloatAt(0)*y.FloatAt(0))*y.FloatAt(1) - y.FloatAt(0), + }, 2) + } + end, err := IntegrateROS4(vdp, 0, 2, mustFloats(t, []float64{2, 0}), + ODEOptions{MaxSteps: 100000}) + if err != nil { + t.Fatalf("IntegrateROS4 Van der Pol: %v", err) + } + for i := range 2 { + if math.IsNaN(end.FloatAt(i)) || math.IsInf(end.FloatAt(i), 0) { + t.Fatalf("Van der Pol state[%d] = %g left the finite range", i, end.FloatAt(i)) + } + } + // The trajectory returns onto the slow branch near x = 2 with a + // small velocity; anything else means the jump was not resolved. + if math.Abs(end.FloatAt(0)-2) > 0.01 || math.Abs(end.FloatAt(1)) > 0.01 { + t.Fatalf("Van der Pol end = (%.10g, %.10g), want the slow branch near (2, 0)", + end.FloatAt(0), end.FloatAt(1)) + } +} + +// TestIntegrateROS4Errors pins the error contract: a degenerate span +// returns the initial state unchanged, a wrong-shaped f, a rank-2 +// state, an empty state, an exhausted step budget and an f that blows +// up mid-span are errors. +func TestIntegrateROS4Errors(t *testing.T) { + y0 := mustFloats(t, []float64{1}) + same, err := IntegrateROS4(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 := IntegrateROS4(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 := IntegrateROS4(decay, 0, 1, matrixState, ODEOptions{}); err == nil { + t.Fatal("expected an error for a rank-2 state") + } + if _, err := IntegrateROS4(decay, 0, 1, mustFloats(t, nil), ODEOptions{}); err == nil { + t.Fatal("expected an error for an empty state") + } + if _, err := IntegrateROS4(decay, 0, 1, y0, ODEOptions{MaxSteps: 2}); err == nil { + t.Fatal("expected an error for an exhausted step budget") + } else if !strings.Contains(err.Error(), "MaxSteps=2") { + t.Fatalf("want a step-budget error, got %v", err) + } + 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 := IntegrateROS4(boom, 0, 1, y0, ODEOptions{}); err == nil { + t.Fatal("expected the operator error to propagate") + } + nan := func(t float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{math.NaN()}, 1) + } + if _, err := IntegrateROS4(nan, 0, 1, y0, ODEOptions{MaxSteps: 200}); err == nil { + t.Fatal("expected an error when f returns NaN throughout") + } +} + +// TestROS4StepErrors pins the error paths of a single attempted step: +// an f failing inside the numerical Jacobian and an f failing at the +// later stage times are both fatal to the step. +func TestROS4StepErrors(t *testing.T) { + always := func(t float64, y *core.Array) (*core.Array, error) { + return nil, base.Errf("detector tripped") + } + if _, _, err := ros4Step("TestROS4StepErrors", always, &odeWork{}, 0, []float64{1}, 0.1, 1e-6, 1e-9); err == nil { + t.Fatal("expected the Jacobian's f error to propagate") + } + gated := func(t float64, y *core.Array) (*core.Array, error) { + if t > 0 { + return nil, base.Errf("detector tripped") + } + return core.FromFloats([]float64{0}, 1) + } + if _, _, err := ros4Step("TestROS4StepErrors", gated, &odeWork{}, 0, []float64{1}, 0.1, 1e-6, 1e-9); err == nil { + t.Fatal("expected the stage f error to propagate") + } +} diff --git a/integrate/pde.go b/integrate/pde.go new file mode 100644 index 0000000..9904b57 --- /dev/null +++ b/integrate/pde.go @@ -0,0 +1,256 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Turnkey PDE evolution in one space dimension: the two +// equations half of physics reduces to, wrapped on machinery the +// library already owns. The heat equation runs Crank-Nicolson (the +// unconditionally stable trapezoidal rule) through the shared +// tridiagonal solver; the wave equation runs velocity Verlet on the +// second-order form, kick-drift-kick like the Hamiltonian integrator +// it is. Both return the trajectory sampled on a time grid, +// IntegrateODEPath-style. + +// pdeValidate checks the shared input contract and returns the grid +// size. +func pdeValidate(name string, u0 *core.Array, dx, tFinal, dt float64, samples int) (int, error) { + if u0.NDim() != 1 || u0.Len() == 0 { + return 0, base.Errf("%s: the initial condition must be a non-empty rank-1 array, got shape %s", + name, base.ShapeText(u0.Shape())) + } + if u0.Dtype() == core.Complex { + return 0, base.Errf("%s: complex states are not supported", name) + } + if !(dx > 0) { + return 0, base.Errf("%s: the grid spacing must be positive, got %g", name, dx) + } + if !(tFinal > 0) { + return 0, base.Errf("%s: the integration time must be positive, got %g", name, tFinal) + } + if !(dt > 0) { + return 0, base.Errf("%s: the time step must be positive, got %g", name, dt) + } + // A dt far below tFinal/1e12 cannot be honoured: the step count + // would leave the int range on some platforms and wrap on others, + // and the silently larger step would run past the wave equation's + // CFL check, which runs on the requested dt. + if tFinal/dt > 1e12 { + return 0, base.Errf("%s: dt = %g asks for more than 1e12 steps over %g", name, dt, tFinal) + } + if samples < 2 { + return 0, base.Errf("%s: at least two samples are needed, got %d", name, samples) + } + // A non-finite entry would flow through the stencil and the + // tridiagonal solve's zero-pivot checks compare false against NaN, + // publishing an all-NaN history with no error. + for i := range u0.Len() { + if v := u0.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return 0, base.Errf("%s: the initial condition holds the non-finite value %g at %d", name, v, i) + } + } + return u0.Len(), nil +} + +// pdeSchedule picks the step count and the actual step size for a +// requested dt. The count is rounded up to a multiple of the sampling +// interval, so every published time j·tFinal/(samples−1) is a step +// boundary and the last step lands on tFinal exactly: the returned +// samples are the evenly spaced interior states the documentation +// promises, not the states at multiples of dt. +func pdeSchedule(tFinal, dt float64, samples int) (steps int, h float64) { + steps = max(int(math.Ceil(tFinal/dt)), samples-1) + if rem := steps % (samples - 1); rem != 0 { + steps += samples - 1 - rem + } + return steps, tFinal / float64(steps) +} + +// IntegrateHeat1D evolves u_t = κ·u_xx over [0, L] discretised by the +// interior grid of u0 (n = u0.Len(), dx = L/(n+1)), from t = 0 to +// tFinal in equal steps of at most dt, holding the boundary values +// boundL and boundR (Dirichlet). It returns the (samples, n) array of +// interior states evenly spaced in time, endpoints included. Crank- +// Nicolson is stable for any dt; accuracy wants dt of a few dx²/κ. +func IntegrateHeat1D(u0 *core.Array, kappa, dx, tFinal, dt float64, samples int, boundL, boundR float64) (*core.Array, error) { + const name = "IntegrateHeat1D" + n, err := pdeValidate(name, u0, dx, tFinal, dt, samples) + if err != nil { + return nil, err + } + if !(kappa > 0) || math.IsInf(kappa, 0) { + return nil, base.Errf("%s: the diffusivity must be positive, got %g", name, kappa) + } + // The boundary values enter the right side every step: a non-finite + // one would flow through the stencil and the solve and publish an + // all-NaN history with no error. + if math.IsNaN(boundL) || math.IsInf(boundL, 0) || math.IsNaN(boundR) || math.IsInf(boundR, 0) { + return nil, base.Errf("%s: the boundary values must be finite, got %g and %g", name, boundL, boundR) + } + u := make([]float64, n) + copy(u, denseFloats(u0)) + // Crank-Nicolson: (I − r/2·A)uⁿ⁺¹ = (I + r/2·A)uⁿ with A the + // second-difference stencil and r = κ·h/dx² for the step h the + // schedule actually takes; the Dirichlet neighbours enter the right + // side through the stencil ends. + steps, h := pdeSchedule(tFinal, dt, samples) + r := kappa * h / (dx * dx) + lower := make([]float64, n-1) + diag := make([]float64, n) + upper := make([]float64, n-1) + // The left side carries I − r/2·A: the diagonal gains r (A's −2 + // times −r/2) and the off-diagonals stay −r/2. The system is + // strictly diagonally dominant for every positive r, so the + // elimination's pivots stay finite and non-zero. + for i := range n { + diag[i] = 1 + r + if i < n-1 { + lower[i] = -r / 2 + upper[i] = -r / 2 + } + } + + every := steps / (samples - 1) + out := make([]float64, samples*n) + copy(out, u) + written := 1 + // The right side and the elimination scratch are constants of one + // solve: every step refills the same buffers, the kernel reads its + // inputs without touching them, and the solution is written + // straight into the working state the samples copy from. + var tri triScratch + triSized(&tri, n, n) + rhs := tri.rhs + // The trapezoidal weight is a constant of the scheme: one division + // by two, the same value the expression inside the loop carried. + half := r / 2 + for s := 1; s <= steps; s++ { + for i := range n { + um, up := boundL, boundR + if i > 0 { + um = u[i-1] + } + if i < n-1 { + up = u[i+1] + } + rhs[i] = u[i] + half*(um-2*u[i]+up) + } + // The implicit side's boundary neighbours move across as known + // data: the first and last rows only, in that order. + rhs[0] += half * boundL + rhs[n-1] += half * boundR + if err := base.TriSolve(u, tri.cp, tri.dp, lower, diag, upper, rhs); err != nil { + return nil, base.Errf("%s: %w", name, err) + } + if s%every == 0 && written < samples { + copy(out[written*n:(written+1)*n], u) + written++ + } + } + // The final state is the last sample whatever the grid remainder. + copy(out[(samples-1)*n:], u) + return core.FromFloats(out, samples, n) +} + +// IntegrateWave1D evolves u_tt = c²·u_xx over [0, L] with the grid of +// u0 (dx = L/(n+1), Dirichlet ends held at zero) and the initial +// velocity v0, by velocity Verlet with fixed step dt. The CFL budget +// |c·dt/dx| ≤ 1 is a genuine stability requirement and is enforced as +// an error. The return contract mirrors IntegrateHeat1D. +func IntegrateWave1D(u0, v0 *core.Array, c, dx, tFinal, dt float64, samples int) (*core.Array, error) { + const name = "IntegrateWave1D" + n, err := pdeValidate(name, u0, dx, tFinal, dt, samples) + if err != nil { + return nil, err + } + if v0.NDim() != 1 || v0.Len() != n { + return nil, base.Errf("%s: the initial velocity must match the state shape, got %s", + name, base.ShapeText(v0.Shape())) + } + if v0.Dtype() == core.Complex { + return nil, base.Errf("%s: complex velocities are not supported", name) + } + // As for u0: a non-finite velocity flows through the Verlet kick + // and poisons the trajectory without an error. + for i := range n { + if v := v0.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: the initial velocity holds the non-finite value %g at %d", name, v, i) + } + } + // The CFL ratio compares false against 1 when it is NaN, so a + // non-finite speed is refused before the budget test. + if math.IsNaN(c) || math.IsInf(c, 0) { + return nil, base.Errf("%s: the wave speed must be finite, got %g", name, c) + } + cfl := math.Abs(c * dt / dx) + if cfl > 1 { + return nil, base.Errf("%s: CFL violated, |c·dt/dx| = %.3g > 1", name, cfl) + } + u := make([]float64, n) + v := make([]float64, n) + copy(u, denseFloats(u0)) + copy(v, denseFloats(v0)) + // The stencil's constants: the products are the ones the element + // loop evaluated, built once per run. + cc := c * c + dx2 := dx * dx + accel := func(dst, us []float64) { + // The two ends take the fixed zero neighbour the boundaries + // impose; the interior runs the same stencil over the real + // neighbours, so the two tests leave the element loop. Every + // term keeps the order the uniform loop evaluated. + end := func(i int) { + um, up := 0.0, 0.0 + if i > 0 { + um = us[i-1] + } + if i < n-1 { + up = us[i+1] + } + dst[i] = cc * (um - 2*us[i] + up) / dx2 + } + end(0) + for i := 1; i < n-1; i++ { + dst[i] = cc * (us[i-1] - 2*us[i] + us[i+1]) / dx2 + } + if n > 1 { + end(n - 1) + } + } + // Kick-drift-kick: the modified energy stays within (c·dt/dx)²/8 + // of the true one, which is why the wave equation keeps its shape. + // The schedule's step never exceeds dt, so the CFL check above + // bounds this one too. + steps, h := pdeSchedule(tFinal, dt, samples) + every := steps / (samples - 1) + out := make([]float64, samples*n) + copy(out, u) + written := 1 + a := make([]float64, n) + for s := 1; s <= steps; s++ { + accel(a, u) + for i := range n { + v[i] += 0.5 * h * a[i] + } + for i := range n { + u[i] += h * v[i] + } + accel(a, u) + for i := range n { + v[i] += 0.5 * h * a[i] + } + if s%every == 0 && written < samples { + copy(out[written*n:(written+1)*n], u) + written++ + } + } + copy(out[(samples-1)*n:], u) + return core.FromFloats(out, samples, n) +} diff --git a/integrate/pde2d.go b/integrate/pde2d.go new file mode 100644 index 0000000..16f07df --- /dev/null +++ b/integrate/pde2d.go @@ -0,0 +1,423 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Two-dimensional evolution equations on the rectangle, the +// higher-rank siblings of IntegrateHeat1D and IntegrateWave1D. The +// grid is the rank-2 shape of the initial state: row r samples +// y = r·dy, column c samples x = c·dx, the boundary ring is held +// fixed, and the interior carries the dynamics. + +// pde2dValidate checks the shared rectangle arguments and returns the +// grid shape. +func pde2dValidate(name string, u0 *core.Array, dx, dy, tFinal, dt float64, samples int) (rows, cols int, err error) { + if u0.NDim() != 2 { + return 0, 0, base.Errf("%s: the initial state must be rank 2, got shape %s", name, base.ShapeText(u0.Shape())) + } + if u0.Dtype() == core.Complex { + return 0, 0, base.Errf("%s: complex initial states are not supported", name) + } + rows, cols = u0.Shape()[0], u0.Shape()[1] + if rows < 3 || cols < 3 { + return 0, 0, base.Errf("%s: the grid must be at least 3×3 to hold interior points, got %d×%d", name, rows, cols) + } + if !(dx > 0) || !(dy > 0) { + return 0, 0, base.Errf("%s: the spacings must be positive, got %g and %g", name, dx, dy) + } + if !(tFinal > 0) || !(dt > 0) { + return 0, 0, base.Errf("%s: tFinal and dt must be positive, got %g and %g", name, tFinal, dt) + } + // As in the 1-D validation: a dt this far below tFinal cannot be + // honoured, and the wrapped step count would disable the bound + // silently. + if tFinal/dt > 1e12 { + return 0, 0, base.Errf("%s: dt = %g asks for more than 1e12 steps over %g", name, dt, tFinal) + } + if samples < 2 { + return 0, 0, base.Errf("%s: at least two samples are needed, got %d", name, samples) + } + // A non-finite entry flows through the stencils and the solvers' + // zero-pivot guards compare false against NaN, so it is refused + // up front. + for i := range u0.Len() { + if v := u0.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return 0, 0, base.Errf("%s: the initial state holds the non-finite value %g at %d", name, v, i) + } + } + return rows, cols, nil +} + +// fromSlice wraps a float slice of known length as an array. +func fromSlice(vals []float64, n int) *core.Array { + out := core.New(core.Float, n) + copy(out.RawFloats(), vals[:n]) + return out +} + +// IntegrateHeat2D evolves u_t = κ·Δu on the rectangle by +// Peaceman-Rachford alternating direction implicit steps: one +// half-step implicit in x against an explicit y neighbour sum, one +// implicit in y against an explicit x neighbour sum, each row and +// column a tridiagonal solve through SolveTridiagonal. The scheme is +// second order in space and time and unconditionally stable, so no +// step-size refusal stands between the caller and a coarse first +// look. The boundary ring is held at the four constant edge values. +// The return is a (samples × rows × cols) array: the initial state, +// then the state after each stored interval, the final state forced +// into the last sample. +func IntegrateHeat2D(u0 *core.Array, kappa, dx, dy, tFinal, dt float64, samples int, + boundBottom, boundTop, boundLeft, boundRight float64) (*core.Array, error) { + const name = "IntegrateHeat2D" + rows, cols, err := pde2dValidate(name, u0, dx, dy, tFinal, dt, samples) + if err != nil { + return nil, err + } + if !(kappa > 0) || math.IsInf(kappa, 0) { + return nil, base.Errf("%s: the diffusivity must be positive, got %g", name, kappa) + } + // The four boundary constants enter the explicit neighbour sums + // every step: a non-finite one would flow through the stencils and + // the solves and publish an all-NaN history with no error. + if math.IsNaN(boundBottom) || math.IsInf(boundBottom, 0) || + math.IsNaN(boundTop) || math.IsInf(boundTop, 0) || + math.IsNaN(boundLeft) || math.IsInf(boundLeft, 0) || + math.IsNaN(boundRight) || math.IsInf(boundRight, 0) { + return nil, base.Errf("%s: the boundary values must be finite, got bottom %g, top %g, left %g, right %g", + name, boundBottom, boundTop, boundLeft, boundRight) + } + steps, h := pdeSchedule(tFinal, dt, samples) + rx := kappa * h / (2 * dx * dx) // the implicit half-step weight in x + ry := kappa * h / (2 * dy * dy) // ...and in y + + u := make([]float64, rows*cols) + copy(u, denseFloats(u0)) + every := steps / (samples - 1) + out := core.New(core.Float, samples, rows, cols) + into := out.RawFloats() + // Sample 0 is the initial state exactly, before any ring + // enforcement touches the working state. + copy(into, u) + written := 1 + // The working state's horizontal ring rows carry the boundary + // constants from here on, so the stencils read their up and down + // neighbours straight out of the payload: a ring slot holds exactly + // the value the boundary test it replaces would have substituted. + // The step loop below rewrites both rows after every half-step, + // so the invariant holds for every step, the first included. + for c := range cols { + u[c] = boundBottom + u[(rows-1)*cols+c] = boundTop + } + + star := make([]float64, rows*cols) + var stepErr error + var errMu sync.Mutex + // The implicit diagonals are constants of the scheme, one set per + // orientation, built once with exactly the values the solves read. + 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 + } + // Per-solve line scratch, one set per pdeStepFloor lines of each + // orientation. engine.ParallelMin invokes a sweep closure once per + // chunk; in parallel the chunk starts differ by at least + // pdeStepFloor, so two concurrent chunks never share a set, and a + // sweep the engine runs inline reuses set zero serially. + triX := make([]triScratch, (rows-2-1)/pdeStepFloor+1) + triY := make([]triScratch, (cols-2-1)/pdeStepFloor+1) + for i := range triX { + triSized(&triX[i], cols-2, cols) + } + for i := range triY { + triSized(&triY[i], rows-2, rows) + } + for s := 1; s <= steps; s++ { + // Half-step one: implicit in x along every interior row. Rows are + // independent (a row's solve reads its own and its neighbours' + // values and writes the star row), so the sweep is partitioned + // over them; each chunk carries its own right-hand side and + // solve scratch, and the arithmetic of a row is untouched. + engine.ParallelMin(rows-2, pdeStepFloor, func(rs, re int) { + tri := &triX[rs/pdeStepFloor] + rhs := tri.rhs + for r := rs + 1; r < re+1; r++ { + // The row's own values and its two neighbours are slices + // of the payload, so the stencil reads them by offset + // instead of rebuilding the flat index per element. The + // horizontal ring rows hold the boundary constants, so + // the neighbours need no edge test. + row := u[r*cols : (r+1)*cols] + up := u[(r+1)*cols : (r+2)*cols] + down := u[(r-1)*cols : r*cols] + for c := range cols { + rhs[c] = row[c] + ry*(up[c]-2*row[c]+down[c]) + } + // Boundary neighbours of the implicit solve: the row's + // own left and right ring values enter the right side. + rhs[1] += rx * boundLeft + rhs[cols-2] += rx * boundRight + // The elimination writes the row's interior straight + // into the star state; the ring values follow. + dst := star[r*cols+1 : r*cols+cols-1] + if serr := base.TriSolve(dst, tri.cp, tri.dp, lowerX, diagX, upperX, rhs[1:cols-1]); serr != nil { + errMu.Lock() + if stepErr == nil { + stepErr = base.Errf("%s: %w", name, serr) + } + errMu.Unlock() + return + } + star[r*cols] = boundLeft + star[r*cols+cols-1] = boundRight + } + }) + if stepErr != nil { + return nil, stepErr + } + for c := range cols { + star[c] = boundBottom + star[(rows-1)*cols+c] = boundTop + } + // Half-step two: implicit in y down every interior column, + // with the explicit x neighbours from the star state. Columns are + // independent the same way the rows were, and the walk takes them + // two at a time: the star state's vertical ring columns hold the + // boundary constants, so both lanes read their neighbours + // straight out of the payload, and a pair shares its loads, so + // each line of the state is fetched once per pair instead of + // three times across three column passes. + engine.ParallelMin(cols-2, pdeStepFloor, func(cs, ce int) { + tri := &triY[cs/pdeStepFloor] + rhs := tri.rhs + aux := tri.aux + dst := tri.dst + for c := cs + 1; c < ce+1; c += 2 { + if c+1 < ce+1 { + // Columns c and c+1: one pass over the rows builds + // both right sides, the shared neighbour loads in + // registers. + for r := range rows { + base := r * cols + l := star[base+c-1] + wm := star[base+c] + e := star[base+c+1] + rhs[r] = wm + rx*(e-2*wm+l) + aux[r] = e + rx*(star[base+c+2]-2*e+wm) + } + rhs[1] += ry * boundBottom + rhs[rows-2] += ry * boundTop + if serr := base.TriSolve(dst, tri.cp, tri.dp, lowerY, diagY, upperY, rhs[1:rows-1]); serr != nil { + errMu.Lock() + if stepErr == nil { + stepErr = base.Errf("%s: %w", name, serr) + } + errMu.Unlock() + return + } + u[c] = boundBottom + u[(rows-1)*cols+c] = boundTop + // The column's interior comes straight out of the + // solve's payload: the values are the ones the + // accessor read. + for r := 1; r < rows-1; r++ { + u[r*cols+c] = dst[r-1] + } + aux[1] += ry * boundBottom + aux[rows-2] += ry * boundTop + if serr := base.TriSolve(dst, tri.cp, tri.dp, lowerY, diagY, upperY, aux[1:rows-1]); serr != nil { + errMu.Lock() + if stepErr == nil { + stepErr = base.Errf("%s: %w", name, serr) + } + errMu.Unlock() + return + } + u[c+1] = boundBottom + u[(rows-1)*cols+c+1] = boundTop + for r := 1; r < rows-1; r++ { + u[r*cols+c+1] = dst[r-1] + } + continue + } + // The odd tail column, when the chunk ends on one. + for r := range rows { + base := r * cols + wm := star[base+c] + rhs[r] = wm + rx*(star[base+c+1]-2*wm+star[base+c-1]) + } + rhs[1] += ry * boundBottom + rhs[rows-2] += ry * boundTop + if serr := base.TriSolve(dst, tri.cp, tri.dp, lowerY, diagY, upperY, rhs[1:rows-1]); serr != nil { + errMu.Lock() + if stepErr == nil { + stepErr = base.Errf("%s: %w", name, serr) + } + errMu.Unlock() + return + } + u[c] = boundBottom + u[(rows-1)*cols+c] = boundTop + for r := 1; r < rows-1; r++ { + u[r*cols+c] = dst[r-1] + } + } + }) + if stepErr != nil { + return nil, stepErr + } + for r := range rows { + u[r*cols] = boundLeft + u[r*cols+cols-1] = boundRight + } + if s%every == 0 && written < samples { + copy(into[written*rows*cols:(written+1)*rows*cols], u) + written++ + } + } + copy(into[(samples-1)*rows*cols:], u) + return out, nil +} + +// IntegrateWave2D evolves u_tt = c²·Δu on the rectangle with the +// boundary ring held at zero, by the explicit central-difference +// stencil from the velocity Verlet family the 1-D wave solver uses: +// second order in space and time, with the CFL budget +// c·dt·sqrt(1/dx² + 1/dy²) ≤ 1 enforced as an error, because the +// explicit stencil has no honest answer past it. The return contract +// mirrors IntegrateHeat2D. +func IntegrateWave2D(u0, v0 *core.Array, c, dx, dy, tFinal, dt float64, samples int) (*core.Array, error) { + const name = "IntegrateWave2D" + rows, cols, err := pde2dValidate(name, u0, dx, dy, tFinal, dt, samples) + if err != nil { + return nil, err + } + if v0.NDim() != 2 || v0.Shape()[0] != rows || v0.Shape()[1] != cols { + return nil, base.Errf("%s: the velocity must be rank 2 on the same grid, got shape %s", name, base.ShapeText(v0.Shape())) + } + if v0.Dtype() == core.Complex { + return nil, base.Errf("%s: complex velocities are not supported", name) + } + // As for u0: a non-finite velocity enters the Taylor start and + // poisons the three-level stencil without an error. + for i := range v0.Len() { + if v := v0.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: the velocity holds the non-finite value %g at %d", name, v, i) + } + } + if !(c > 0) { + return nil, base.Errf("%s: the wave speed must be positive, got %g", name, c) + } + cfl := c * dt * math.Sqrt(1/(dx*dx)+1/(dy*dy)) + if cfl > 1 { + return nil, base.Errf("%s: the CFL number %g exceeds 1 (c = %g, dt = %g, dx = %g, dy = %g); the explicit stencil is unstable there", + name, cfl, c, dt, dx, dy) + } + steps, h := pdeSchedule(tFinal, dt, samples) + lapW := (c * h) * (c * h) + // The stencil's spacings are constants of the grid; squaring them + // once per step is the same product the element loop evaluated. + dx2 := dx * dx + dy2 := dy * dy + + prev := make([]float64, rows*cols) // u^{s-1} + cur := make([]float64, rows*cols) // u^{s} + next := make([]float64, rows*cols) // u^{s+1}, the buffer being written + copy(cur, denseFloats(u0)) + every := steps / (samples - 1) + out := core.New(core.Float, samples, rows, cols) + into := out.RawFloats() + // Slot 0 is the initial state at t = 0; slot j is the state after + // j·every steps, at the published time j·tFinal/(samples-1), exactly + // like IntegrateHeat2D. + copy(into, cur) + // From there on the boundary ring is held at zero, so clear it in the + // working buffer the initial state came in (the other two are born + // zeroed): otherwise a recycled buffer would carry the caller's ring + // back into the stencil every third step. + for r := range rows { + cur[r*cols], cur[r*cols+cols-1] = 0, 0 + } + for c := range cols { + cur[c], cur[(rows-1)*cols+c] = 0, 0 + } + written := 1 + // The three buffers cycle through (prev, cur, next), so every stencil + // reads u^{s-1} and u^s while writing u^{s+1} into a third buffer: no + // point is ever overwritten before its neighbours have read it. The + // first iteration is the Taylor start + // u¹ = u⁰ + h·v⁰ + h²c²/2·Δu⁰, second order, so the three-level + // stencil starts honest, and it reads the untouched u⁰ (the initial + // velocity and the Laplacian of u⁰), not a half-updated state. + for s := 1; s <= steps; s++ { + first := s == 1 + for r := 1; r < rows-1; r++ { + // The row's neighbours are the slices above and below it, so + // the stencil reads them by offset instead of rebuilding the + // flat index per element. The terms keep their order. + mid := cur[r*cols : (r+1)*cols] + up := cur[(r+1)*cols : (r+2)*cols] + down := cur[(r-1)*cols : r*cols] + if first { + // The Taylor start reads the untouched u⁰ and its + // initial velocity, not a half-updated state. + for cc := 1; cc < cols-1; cc++ { + i := r*cols + cc + lap := (down[cc]-2*mid[cc]+up[cc])/dy2 + + (mid[cc-1]-2*mid[cc]+mid[cc+1])/dx2 + next[i] = mid[cc] + h*v0.FloatAt(i) + 0.5*lapW*lap + } + continue + } + for cc := 1; cc < cols-1; cc++ { + i := r*cols + cc + lap := (down[cc]-2*mid[cc]+up[cc])/dy2 + + (mid[cc-1]-2*mid[cc]+mid[cc+1])/dx2 + next[i] = 2*mid[cc] - prev[i] + lapW*lap + } + } + prev, cur, next = cur, next, prev + if s%every == 0 && written < samples { + copy(into[written*rows*cols:(written+1)*rows*cols], cur) + written++ + } + } + copy(into[(samples-1)*rows*cols:], cur) + return out, nil +} + +// pdeStepFloor is the per-worker floor for a 2-D step sweep. One item is +// one tridiagonal elimination over a stencil row, measured at roughly +// half a microsecond on a 32-line grid, so a worker needs about eight +// of them before the split pays for its spawn. The floor is a constant +// for that reason, not a length-scaled budget. The line scratch is +// indexed by start/pdeStepFloor on the same constant: parallel chunk +// starts differ by at least the floor, so concurrent chunks never share +// a scratch set. +const pdeStepFloor = 8 diff --git a/integrate/pde2d_test.go b/integrate/pde2d_test.go new file mode 100644 index 0000000..77de448 --- /dev/null +++ b/integrate/pde2d_test.go @@ -0,0 +1,238 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// heatMode builds sin(πx)·sin(πy) on an (n+2)×(n+2) grid over [0,1]² +// including the zero boundary ring, the lowest interior mode. +func heatMode(t *testing.T, n int) *core.Array { + t.Helper() + vals := make([]float64, (n+2)*(n+2)) + for r := range n + 2 { + for c := range n + 2 { + vals[r*(n+2)+c] = math.Sin(math.Pi*float64(c)/float64(n+1)) * + math.Sin(math.Pi*float64(r)/float64(n+1)) + } + } + a, err := core.FromFloats(vals, n+2, n+2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestIntegrateHeat2DModeDecay checks the ADI solver on the lowest +// mode: with zero boundaries the amplitude decays like +// exp(−κ·2π²·t), and the scheme's O(dt²+h²) error must stay inside a +// one percent band on a 32-interior grid. +func TestIntegrateHeat2DModeDecay(t *testing.T) { + const n = 32 + u0 := heatMode(t, n) + const kappa, dt, tFinal = 1.0, 0.005, 0.1 + history, err := IntegrateHeat2D(u0, kappa, 1.0/float64(n+1), 1.0/float64(n+1), + tFinal, dt, 2, 0, 0, 0, 0) + if err != nil { + t.Fatalf("IntegrateHeat2D: %v", err) + } + if history.Shape()[0] != 2 || history.Shape()[1] != n+2 { + t.Fatalf("shape %v, want [%d %d %d]", history.Shape(), 2, n+2, n+2) + } + final := history.Shape()[0] - 1 + // The interior peak of the final state versus the exact decay. + peak := 0.0 + for r := 1; r <= n; r++ { + for c := 1; c <= n; c++ { + if v := history.FloatAt(final*(n+2)*(n+2) + r*(n+2) + c); v > peak { + peak = v + } + } + } + want := math.Exp(-kappa * 2 * math.Pi * math.Pi * tFinal) + if math.Abs(peak-want) > 0.01 { + t.Fatalf("final peak %.5g, want %.5g", peak, want) + } + // The boundary ring is held at zero. + for c := range n + 2 { + if history.FloatAt(final*(n+2)*(n+2)+c) != 0 || + history.FloatAt(final*(n+2)*(n+2)+(n+1)*(n+2)+c) != 0 { + t.Fatalf("boundary ring moved at column %d", c) + } + } +} + +// TestIntegrateWave2DStandingWave checks the explicit solver on the +// lowest standing mode with zero initial velocity against the closed +// form of the leapfrog itself. The mode is an exact eigenfunction of the +// five-point Laplacian, with eigenvalue mu, and the discrete +// characteristic of the scheme is cos theta = 1 - (c*h)²·mu/2, so after +// s steps the state is cos(s·theta)·u0. The run below takes 400 steps of +// period/400, so its amplitude is cos(400·theta) = -0.4155, not 1: the +// mode has not returned to its start at this time. +func TestIntegrateWave2DStandingWave(t *testing.T) { + const n = 32 + u0 := heatMode(t, n) // sin(πx)sin(πy) with zero boundary ring + v0 := core.New(core.Float, n+2, n+2) + const c = 1.0 + dx := 1.0 / float64(n+1) + period := math.Sqrt2 / (c * math.Pi) + const steps = 400 // samples = 2 and dt = period/400 force this many steps + dt := period / float64(steps) + history, err := IntegrateWave2D(u0, v0, c, dx, dx, period, dt, 2) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + h := period / float64(steps) + // The discrete eigenvalue of the (1,1) mode, mu = 8/dx²·sin²(π·dx/2) + // on this square grid, and the phase 400 steps accumulate. + sine := math.Sin(math.Pi * dx / 2) + mu := 8 / (dx * dx) * sine * sine + amp := math.Cos(float64(steps) * math.Acos(1-0.5*(c*h)*(c*h)*mu)) + final := history.Shape()[0] - 1 + worst := 0.0 + for r := range n + 2 { + for cc := range n + 2 { + i := final*(n+2)*(n+2) + r*(n+2) + cc + if e := math.Abs(history.FloatAt(i) - amp*u0.FloatAt(r*(n+2)+cc)); e > worst { + worst = e + } + } + } + if worst > 1e-11 { + t.Fatalf("after %.0f steps the worst deviation from cos(%.6f)·u0 is %.4g, want the leapfrog characteristic", + float64(steps), amp, worst) + } +} + +// TestIntegrateWave2DCFLRefusal checks the stability budget: a step +// past the CFL limit is an error, not a silent blow-up. +func TestIntegrateWave2DCFLRefusal(t *testing.T) { + const n = 32 + u0 := heatMode(t, n) + v0 := core.New(core.Float, n+2, n+2) + dx := 1.0 / float64(n+1) + // c·dt·sqrt(1/dx²+1/dy²) = 1·0.05·45.25 ≈ 2.26 > 1. + if _, err := IntegrateWave2D(u0, v0, 1, dx, dx, 0.05, 0.05, 2); err == nil { + t.Fatal("a CFL-violating step accepted") + } + if _, err := IntegrateWave2D(u0, core.New(core.Float, 3, 3), 1, dx, dx, 0.05, 0.001, 2); err == nil { + t.Fatal("mismatched velocity grid accepted") + } +} + +// pde2dAnisoMode builds the (p, q) discrete Dirichlet eigenmode of the +// five-point Laplacian on a rows×cols grid: sin(π·p·c/(cols−1)) · +// sin(π·q·r/(rows−1)), which vanishes on all four boundary lines. +func pde2dAnisoMode(t *testing.T, rows, cols, p, q int) *core.Array { + t.Helper() + vals := make([]float64, rows*cols) + for r := range rows { + for c := range cols { + vals[r*cols+c] = math.Sin(math.Pi*float64(p)*float64(c)/float64(cols-1)) * + math.Sin(math.Pi*float64(q)*float64(r)/float64(rows-1)) + } + } + a, err := core.FromFloats(vals, rows, cols) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// anisoMu returns the dimensionless eigenvalues of the undivided second +// difference along each axis for the (p, q) mode above. +func anisoMu(rows, cols, p, q int) (mx, my float64) { + return 4 * math.Pow(math.Sin(math.Pi*float64(p)/2/float64(cols-1)), 2), + 4 * math.Pow(math.Sin(math.Pi*float64(q)/2/float64(rows-1)), 2) +} + +// TestIntegrateWave2DAnisotropicGrid pins the explicit solver on a grid +// whose spacings differ between the axes, the case the square-grid tests +// cannot see. The (1, 2) mode is an exact eigenfunction of the +// five-point Laplacian with eigenvalue λx + λy, where λx carries dx and +// λy carries dy, so the leapfrog state after s steps is cos(s·θ)·u0 with +// cos θ = 1 − (c·h)²·(λx + λy)/2. A stencil that divides the y +// neighbours by dx² instead of dy², or the reverse, moves those +// eigenvalues and the amplitude with them. +func TestIntegrateWave2DAnisotropicGrid(t *testing.T) { + const ( + rows, cols = 26, 34 + dx, dy = 0.02, 0.05 + c = 1.0 + steps = 60 + ) + if dx == dy { + t.Fatal("the case needs spacings that differ between the axes") + } + mx, my := anisoMu(rows, cols, 1, 2) + lx, ly := mx/(dx*dx), my/(dy*dy) + dt := 0.5 / (c * math.Sqrt(1/(dx*dx)+1/(dy*dy))) // CFL = 1/2 + tFinal := dt * float64(steps) + u0 := pde2dAnisoMode(t, rows, cols, 1, 2) + v0 := core.New(core.Float, rows, cols) + history, err := IntegrateWave2D(u0, v0, c, dx, dy, tFinal, dt, 2) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + h := tFinal / float64(steps) + amp := math.Cos(float64(steps) * math.Acos(1-0.5*(c*h)*(c*h)*(lx+ly))) + if math.Abs(amp) < 0.1 { + t.Fatalf("the run decays to %.3g; the case needs an amplitude the comparison can see", amp) + } + final := history.Shape()[0] - 1 + worst := 0.0 + for i := range rows * cols { + worst = math.Max(worst, math.Abs(history.FloatAt(final*rows*cols+i)-amp*u0.FloatAt(i))) + } + if worst > 1e-11 { + t.Fatalf("after %d steps the worst deviation from cos(θ·%d)·u0 is %.4g (amplitude %.4g), want the leapfrog characteristic", + steps, steps, worst, amp) + } +} + +// TestIntegrateHeat2DAnisotropicGrid pins the ADI solver the same way: +// on an eigenmode the two half steps compose into one amplification +// factor per step, (1 − rx·μx)(1 − ry·μy)/((1 + rx·μx)(1 + ry·μy)), +// with rx = κ·h/(2·dx²) and ry = κ·h/(2·dy²) against the dimensionless +// second-difference eigenvalues. Exchanging the two spacings moves the +// factor by orders of magnitude, so a mislabelled axis cannot pass. +func TestIntegrateHeat2DAnisotropicGrid(t *testing.T) { + const ( + rows, cols = 26, 34 + dx, dy = 0.02, 0.05 + kappa = 1.0 + steps = 10 + ) + if dx == dy { + t.Fatal("the case needs spacings that differ between the axes") + } + mx, my := anisoMu(rows, cols, 1, 2) + h := 0.0044 + tFinal := h * float64(steps) + u0 := pde2dAnisoMode(t, rows, cols, 1, 2) + history, err := IntegrateHeat2D(u0, kappa, dx, dy, tFinal, h, 2, 0, 0, 0, 0) + if err != nil { + t.Fatalf("IntegrateHeat2D: %v", err) + } + rx := kappa * h / (2 * dx * dx) + ry := kappa * h / (2 * dy * dy) + amp := math.Pow((1-rx*mx)*(1-ry*my)/((1+rx*mx)*(1+ry*my)), float64(steps)) + if math.Abs(amp) < 0.05 { + t.Fatalf("the run decays to %.3g; the case needs an amplitude the comparison can see", amp) + } + final := history.Shape()[0] - 1 + worst := 0.0 + for i := range rows * cols { + worst = math.Max(worst, math.Abs(history.FloatAt(final*rows*cols+i)-amp*u0.FloatAt(i))) + } + if worst > 1e-12 { + t.Fatalf("after %d steps the worst deviation from the ADI amplification %.6g·u0 is %.4g, want the anisotropic factor", + steps, amp, worst) + } +} diff --git a/integrate/pde_scratch_pin_test.go b/integrate/pde_scratch_pin_test.go new file mode 100644 index 0000000..4bdf567 --- /dev/null +++ b/integrate/pde_scratch_pin_test.go @@ -0,0 +1,64 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The 2-D step sweeps index their per-line tridiagonal scratch by +// chunk start divided by the spawn floor, which is only exercised when +// the engine actually spawns: on a many-core runner the fixture grids +// run inline and a scratch-indexing break passes silently. This pin +// forces both regimes and compares every output bit. + +func pinHeat2D(workers int) ([]float64, error) { + prev := engine.SetNumWorkers(workers) + defer engine.SetNumWorkers(prev) + const rows, cols = 20, 20 + 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 { + return nil, err + } + out, err := IntegrateHeat2D(state, 1, 1.0/21, 1.0/21, 0.02, 0.002, 2, 0, 0, 0, 0) + if err != nil { + return nil, err + } + return append([]float64{}, out.RawFloats()...), nil +} + +func TestHeat2DScratchIndexingPinnedAcrossWorkers(t *testing.T) { + serial, err := pinHeat2D(1) + if err != nil { + t.Fatalf("serial: %v", err) + } + // Two workers spawn on an 18-line sweep (chunk 9 at the floor of + // 8); four and thirty-two collapse back to inline runs. + for _, w := range []int{2, 4, 32} { + got, err := pinHeat2D(w) + if err != nil { + t.Fatalf("workers=%d: %v", w, err) + } + if len(got) != len(serial) { + t.Fatalf("workers=%d: length %d, want %d", w, len(got), len(serial)) + } + for i := range serial { + if math.Float64bits(got[i]) != math.Float64bits(serial[i]) { + t.Fatalf("workers=%d element %d: %#x, want %#x", + w, i, math.Float64bits(got[i]), math.Float64bits(serial[i])) + } + } + } +} diff --git a/integrate/pde_test.go b/integrate/pde_test.go new file mode 100644 index 0000000..862ef09 --- /dev/null +++ b/integrate/pde_test.go @@ -0,0 +1,145 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestHeatEigenmodeDecay pins the analytic solution: the first sine +// eigenmode decays as exp(−κ·π²·t/L²). +func TestHeatEigenmodeDecay(t *testing.T) { + const ( + n = 49 + L = 1.0 + kappa = 0.1 + ) + dx := L / float64(n+1) + u0 := make([]float64, n) + for i := range n { + u0[i] = math.Sin(math.Pi * float64(i+1) * dx / L) + } + u0Arr, _ := core.FromFloats(u0, n) + const tFinal = 1.0 + states, err := IntegrateHeat1D(u0Arr, kappa, dx, tFinal, 0.002, 3, 0, 0) + if err != nil { + t.Fatalf("IntegrateHeat1D: %v", err) + } + decay := math.Exp(-kappa * math.Pi * math.Pi * tFinal / (L * L)) + final := states.Shape()[0]*n - n + for i := range n { + got := states.FloatAt(final + i) + want := u0[i] * decay + if math.Abs(got-want) > 5e-4*decay { + t.Fatalf("u[%d] = %.8f, eigenmode says %.8f", i, got, want) + } + } +} + +// TestHeatConservesConstantWithZeroBounds pins the fixed-point: a +// constant field with equal Dirichlet bounds never moves. +func TestHeatConservesConstantWithZeroBounds(t *testing.T) { + const n = 20 + u0 := make([]float64, n) + for i := range n { + u0[i] = 2.5 + } + u0Arr, _ := core.FromFloats(u0, n) + states, err := IntegrateHeat1D(u0Arr, 1.0, 0.1, 1.0, 0.05, 5, 2.5, 2.5) + if err != nil { + t.Fatalf("IntegrateHeat1D: %v", err) + } + last := (states.Shape()[0] - 1) * n + for i := range n { + if math.Abs(states.FloatAt(last+i)-2.5) > 1e-12 { + t.Fatalf("constant field drifted: u[%d] = %.12f", i, states.FloatAt(last+i)) + } + } +} + +// TestWaveStandingFrequency pins a standing wave: the fundamental mode +// u(x, t) = sin(πx)·cos(πct) must return to (minus) itself after half +// a period. +func TestWaveStandingFrequency(t *testing.T) { + const ( + n = 99 + L = 1.0 + c = 1.0 + ) + dx := L / float64(n+1) + u0 := make([]float64, n) + for i := range n { + u0[i] = math.Sin(math.Pi * float64(i+1) * dx / L) + } + u0Arr, _ := core.FromFloats(u0, n) + v0, _ := core.FromFloats(make([]float64, n), n) + // Half period of the fundamental: T1/2 = L/c. + states, err := IntegrateWave1D(u0Arr, v0, c, dx, 1.0, 0.001, 2) + if err != nil { + t.Fatalf("IntegrateWave1D: %v", err) + } + last := (states.Shape()[0] - 1) * n + for i := range n { + got := states.FloatAt(last + i) + want := -u0[i] + if math.Abs(got-want) > 2e-3 { + t.Fatalf("standing wave off after half period: u[%d] = %.6f, want %.6f", i, got, want) + } + } +} + +// TestWaveEnergyBand pins Verlet's bounded energy over many periods. +func TestWaveEnergyBand(t *testing.T) { + const ( + n = 79 + L = 1.0 + c = 1.0 + ) + dx := L / float64(n+1) + u0 := make([]float64, n) + for i := range n { + u0[i] = math.Sin(math.Pi*float64(i+1)*dx/L) + 0.3*math.Sin(3*math.Pi*float64(i+1)*dx/L) + } + u0Arr, _ := core.FromFloats(u0, n) + v0, _ := core.FromFloats(make([]float64, n), n) + states, err := IntegrateWave1D(u0Arr, v0, c, dx, 10.0, 0.002, 11) + if err != nil { + t.Fatalf("IntegrateWave1D: %v", err) + } + energy := func(row int) float64 { + e := 0.0 + for i := range n { + e += states.FloatAt(row*n+i) * states.FloatAt(row*n+i) + } + return e + } + e0 := energy(0) + for row := 1; row < 11; row++ { + e := energy(row) + if math.Abs(e-e0) > 1e-3*e0 { + t.Fatalf("energy drifted: row %d has %.8f vs %.8f", row, e, e0) + } + } +} + +// TestPDEErrors pins the input gates. +func TestPDEErrors(t *testing.T) { + u, _ := core.FromFloats([]float64{1, 2, 3}, 3) + if _, err := IntegrateHeat1D(u, -1, 0.1, 1, 0.01, 2, 0, 0); err == nil { + t.Error("negative diffusivity accepted") + } + if _, err := IntegrateHeat1D(u, 1, 0.1, 1, 0.01, 1, 0, 0); err == nil { + t.Error("one sample accepted") + } + v, _ := core.FromFloats([]float64{1, 2}, 2) + if _, err := IntegrateWave1D(u, v, 1, 0.1, 1, 0.01, 2); err == nil { + t.Error("velocity shape mismatch accepted") + } + if _, err := IntegrateWave1D(u, u, 1, 0.1, 1, 0.2, 2); err == nil { + t.Error("CFL violation accepted") + } +} diff --git a/integrate/pdeadvect.go b/integrate/pdeadvect.go new file mode 100644 index 0000000..b32bed0 --- /dev/null +++ b/integrate/pdeadvect.go @@ -0,0 +1,291 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Advection in one space dimension, the transport siblings of +// IntegrateHeat1D and IntegrateWave1D: u_t + a·u_x = 0 and the +// advection-diffusion equation u_t + a·u_x = D·u_xx. Transport is +// where the cheap stencils fail in public: the first-order upwind +// flux is monotone but smears a front every step it takes, and the +// centred flux that heat uses is oscillatory or worse here. The +// middle ground is a flux-limited scheme: an upwind flux whose face +// value is raised toward the third-order one by a slope limiter that +// switches itself off across discontinuities, keeping the scheme +// total variation diminishing. +// +// The limiter is Koren's third-order one, φ(θ) = max(0, min(2θ, +// (2+θ)/3, 2)), stated here as the choice. The face value is the +// Sweby flux-limited form u_upwind + ½φ(θ)·(1−|ν|)·Δ, with ν = +// |a|·dt/dx the CFL number and Δ the forward difference across the +// face: the (1−|ν|) factor is what makes the explicit update total +// variation diminishing in the Sweby sense for CFL ≤ 1, and the +// third-order branch of φ is what keeps smooth profiles sharp (a +// marginally small overshoot past the strict Sweby bound survives as +// roundoff-scale noise). The grid convention is the house one: u0 +// carries the interior cells with dx = L/(n+1), and boundL, boundR +// are the fixed values of the ghost cells on the ends. The face +// adjacent to the inflow end takes its prescribed value directly +// (first order there); on every other face, including the outflow +// end, the limiter runs at full strength with the prescribed ghost +// value entering its smoothness ratio, and a window too short for a +// three-cell stencil forces first order. + +// korenSlope returns the Koren-limited normalised slope φ(θ) for the +// smoothness ratio θ: the three branches are the Sweby bounds with +// the third-order diagonal (2+θ)/3. +func korenSlope(theta float64) float64 { + phi := min(2*theta, (2+theta)/3, 2) + if phi < 0 { + return 0 + } + return phi +} + +// advectFaceFlux returns the numerical flux a·u_face across the face +// between cells ul (left) and ur (right), with the second neighbours +// ull (left of ul) and urr (right of ur) feeding the limiter and nu +// the CFL number feeding the time factor. The wind decides the side +// the face value is reconstructed from; with limited false the flux +// is plain first-order upwind. +func advectFaceFlux(a, ull, ul, ur, urr, nu float64, limited bool) float64 { + d := ur - ul + if a >= 0 { + if !limited || d == 0 { + return a * ul + } + theta := (ul - ull) / d + return a * (ul + 0.5*korenSlope(theta)*(1-nu)*d) + } + if !limited || d == 0 { + return a * ur + } + theta := (ur - urr) / d + return a * (ur - 0.5*korenSlope(theta)*(1-nu)*d) +} + +// advectStep advances dst from src by one explicit step of width h +// of the conservative update u −= (h/dx)·(F₊ − F₋), the faces built +// from src with the ghost values boundL and boundR on the ends. dst +// and src must not alias. Face j sits between cell j−1 and cell j, +// so face 0 borders the left ghost and face n the right one. +func advectStep(dst, src []float64, a, dx, h, boundL, boundR float64, limited bool, faces []float64) { + n := len(src) + lambda := h / dx + nu := math.Abs(a) * lambda + for j := range n + 1 { + switch { + case j == 0: + // The inflow face takes the prescribed ghost value; on an + // outflow left end the upwind cell is cell 0, whose + // limited face value reaches one cell into the interior. + if a >= 0 { + faces[j] = a * boundL + } else if n < 2 || !limited { + faces[j] = a * src[0] + } else { + faces[j] = advectFaceFlux(a, boundL, boundL, src[0], src[1], nu, true) + } + case j == n: + if a >= 0 { + if n < 2 || !limited { + faces[j] = a * src[n-1] + } else { + faces[j] = advectFaceFlux(a, src[n-2], src[n-1], boundR, boundR, nu, true) + } + } else { + faces[j] = a * boundR + } + default: + ull := boundL + if j >= 2 { + ull = src[j-2] + } + urr := boundR + if j <= n-2 { + urr = src[j+1] + } + faces[j] = advectFaceFlux(a, ull, src[j-1], src[j], urr, nu, limited) + } + } + for i := range n { + dst[i] = src[i] - lambda*(faces[i+1]-faces[i]) + } +} + +// advectValidate checks the shared input contract of the transport +// solvers and returns the cell count. The grid, step bound, sample +// and finiteness gates are pdeValidate's; transport adds a finite +// speed and finite ghost values, and enforces the CFL budget +// |a|·dt/dx ≤ 1 the explicit update needs, exactly like the wave +// solver enforces its own. +func advectValidate(name string, u0 *core.Array, a, dx, tFinal, dt float64, samples int, boundL, boundR float64) (int, error) { + n, err := pdeValidate(name, u0, dx, tFinal, dt, samples) + if err != nil { + return 0, err + } + if math.IsNaN(a) || math.IsInf(a, 0) { + return 0, base.Errf("%s: the transport speed must be finite, got %g", name, a) + } + if math.IsNaN(boundL) || math.IsInf(boundL, 0) || math.IsNaN(boundR) || math.IsInf(boundR, 0) { + return 0, base.Errf("%s: the ghost values must be finite, got %g and %g", name, boundL, boundR) + } + cfl := math.Abs(a * dt / dx) + if cfl > 1 { + return 0, base.Errf("%s: CFL violated, |a·dt/dx| = %.3g > 1", name, cfl) + } + return n, nil +} + +// IntegrateAdvection1D evolves u_t + a·u_x = 0 over the grid of u0 +// from t = 0 to tFinal in equal steps of at most dt, and returns the +// (samples, n) array of interior states evenly spaced in time, +// endpoints included, exactly like IntegrateHeat1D. The flux is the +// Koren-limited upwind one described at the top of the file: total +// variation diminishing under CFL ≤ 1, third-order at smooth faces +// and first order next to the inflow boundary, so a front is carried +// sharply where plain upwind would smear it away. +func IntegrateAdvection1D(u0 *core.Array, a, dx, tFinal, dt float64, samples int, boundL, boundR float64) (*core.Array, error) { + const name = "IntegrateAdvection1D" + n, err := advectValidate(name, u0, a, dx, tFinal, dt, samples, boundL, boundR) + if err != nil { + return nil, err + } + return advectRun(name, u0, a, dx, tFinal, dt, samples, boundL, boundR, n, true) +} + +// IntegrateUpwindAdvection1D evolves the same equation with the +// plain first-order upwind flux: monotone under CFL ≤ 1 (no new +// extrema, ever) and diffuse, the baseline the limited scheme is +// measured against. The return contract mirrors IntegrateAdvection1D. +func IntegrateUpwindAdvection1D(u0 *core.Array, a, dx, tFinal, dt float64, samples int, boundL, boundR float64) (*core.Array, error) { + const name = "IntegrateUpwindAdvection1D" + n, err := advectValidate(name, u0, a, dx, tFinal, dt, samples, boundL, boundR) + if err != nil { + return nil, err + } + return advectRun(name, u0, a, dx, tFinal, dt, samples, boundL, boundR, n, false) +} + +// advectRun is the shared stepping loop of the two pure-transport +// solvers: fixed steps on the pdeSchedule grid, states sampled every +// steps/(samples−1) steps with the final state forced into the last +// sample. +func advectRun(name string, u0 *core.Array, a, dx, tFinal, dt float64, samples int, boundL, boundR float64, n int, limited bool) (*core.Array, error) { + u := make([]float64, n) + for i := range n { + u[i] = u0.FloatAt(i) + } + steps, h := pdeSchedule(tFinal, dt, samples) + every := steps / (samples - 1) + out := core.New(core.Float, samples, n) + into := out.RawFloats() + copy(into, u) + written := 1 + faces := make([]float64, n+1) + scratch := make([]float64, n) + for s := 1; s <= steps; s++ { + advectStep(scratch, u, a, dx, h, boundL, boundR, limited, faces) + u, scratch = scratch, u + if s%every == 0 && written < samples { + copy(into[written*n:(written+1)*n], u) + written++ + } + } + copy(into[(samples-1)*n:], u) + return out, nil +} + +// IntegrateAdvectionDiffusion1D evolves u_t + a·u_x = D·u_xx over the +// grid of u0 with the Dirichlet ghost values boundL and boundR. Each +// step combines the Koren-limited advection flux, advanced +// explicitly, with the Crank-Nicolson second-difference diffusion the +// heat solver runs through the shared tridiagonal solve, so the +// composition is first order in time and second order in space and +// the diffusion side is unconditionally stable. The explicit +// advection still answers for its own CFL budget |a|·dt/dx ≤ 1 and +// the step is refused past it. With a = 0 the scheme reduces exactly +// to IntegrateHeat1D. +func IntegrateAdvectionDiffusion1D(u0 *core.Array, a, kappa, dx, tFinal, dt float64, samples int, boundL, boundR float64) (*core.Array, error) { + const name = "IntegrateAdvectionDiffusion1D" + n, err := advectValidate(name, u0, a, dx, tFinal, dt, samples, boundL, boundR) + if err != nil { + return nil, err + } + if !(kappa > 0) || math.IsInf(kappa, 0) { + return nil, base.Errf("%s: the diffusivity must be positive, got %g", name, kappa) + } + u := make([]float64, n) + for i := range n { + u[i] = u0.FloatAt(i) + } + steps, h := pdeSchedule(tFinal, dt, samples) + r := kappa * h / (dx * dx) + lower := make([]float64, n-1) + diag := make([]float64, n) + upper := make([]float64, n-1) + // The implicit left side I − r/2·A is a constant of the scheme, + // built once exactly as IntegrateHeat1D builds it, strictly + // diagonally dominant for every positive r like the heat system. + for i := range n { + diag[i] = 1 + r + if i < n-1 { + lower[i] = -r / 2 + upper[i] = -r / 2 + } + } + // The trapezoidal weight and the elimination scratch are constants + // of one solve: every step refills the same right side and the + // solution is written straight into the working state. + half := r / 2 + var tri triScratch + triSized(&tri, n, n) + + every := steps / (samples - 1) + out := core.New(core.Float, samples, n) + into := out.RawFloats() + copy(into, u) + written := 1 + faces := make([]float64, n+1) + advected := make([]float64, n) + rhs := tri.rhs + for s := 1; s <= steps; s++ { + // Explicit advection sub-step on the limited fluxes. + advectStep(advected, u, a, dx, h, boundL, boundR, true, faces) + // Crank-Nicolson diffusion sub-step on the advected state, + // the Dirichlet neighbours entering as known data on both + // sides, in the heat solver's own arithmetic. + for i := range n { + um, up := boundL, boundR + if i > 0 { + um = advected[i-1] + } + if i < n-1 { + up = advected[i+1] + } + rhs[i] = advected[i] + half*(um-2*advected[i]+up) + if i == 0 { + rhs[i] += half * boundL + } + if i == n-1 { + rhs[i] += half * boundR + } + } + if serr := base.TriSolve(u, tri.cp, tri.dp, lower, diag, upper, rhs); serr != nil { + return nil, base.Errf("%s: %w", name, serr) + } + if s%every == 0 && written < samples { + copy(into[written*n:(written+1)*n], u) + written++ + } + } + copy(into[(samples-1)*n:], u) + return out, nil +} diff --git a/integrate/pdeadvect_test.go b/integrate/pdeadvect_test.go new file mode 100644 index 0000000..1c36831 --- /dev/null +++ b/integrate/pdeadvect_test.go @@ -0,0 +1,387 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// advectGrid builds the interior cells of [0, 1] with n cells and the +// matching cell centres. +func advectGrid(n int) (dx float64, centres []float64) { + dx = 1 / float64(n+1) + centres = make([]float64, n) + for i := range n { + centres[i] = float64(i+1) * dx + } + return dx, centres +} + +// advectL1 returns the L1 error of the final sample against exact. +func advectL1(t *testing.T, states *core.Array, n int, exact func(x float64) float64) float64 { + t.Helper() + last := (states.Shape()[0] - 1) * n + _, centres := advectGrid(n) + sum := 0.0 + for i := range n { + sum += math.Abs(states.FloatAt(last+i) - exact(centres[i])) + } + return sum / float64(n) +} + +// TestAdvectionUpwindMonotone pins the discrete maximum principle of +// both schemes: a monotone profile transported at CFL 0.9 stays +// monotone and inside its initial range, sample after sample. +func TestAdvectionUpwindMonotone(t *testing.T) { + for _, limited := range []bool{false, true} { + name := "upwind" + if limited { + name = "Koren" + } + n := 96 + dx, centres := advectGrid(n) + u0 := make([]float64, n) + for i := range n { + u0[i] = 1 - centres[i] + } + u0Arr, err := core.FromFloats(u0, n) + if err != nil { + t.Fatal(err) + } + cfl := 0.9 + a := 1.0 + dt := cfl * dx + run := func() (*core.Array, error) { + if limited { + return IntegrateAdvection1D(u0Arr, a, dx, 0.3, dt, 4, 1, 0) + } + return IntegrateUpwindAdvection1D(u0Arr, a, dx, 0.3, dt, 4, 1, 0) + } + states, err := run() + if err != nil { + t.Fatalf("%s: %v", name, err) + } + rows := states.Shape()[0] + for r := range rows { + prev := math.Inf(1) + for i := range n { + v := states.FloatAt(r*n + i) + if v < -1e-12 || v > 1+1e-12 { + t.Fatalf("%s: sample %d cell %d left the range [0, 1]: %g", name, r, i, v) + } + if v > prev+1e-12 { + t.Fatalf("%s: sample %d stops being non-increasing at cell %d: %g after %g", + name, r, i, v, prev) + } + prev = v + } + } + // The transported ramp also keeps its shape: the last sample + // tracks the shifted exact ramp to within the scheme's smear. + last := (rows - 1) * n + shift := 0.3 + for i := range n { + x := centres[i] - shift + want := 1.0 + if x > 0 { + want = 1 - x + } + if d := math.Abs(states.FloatAt(last+i) - want); d > 0.02 { + t.Fatalf("%s: ramp cell %d: %g, want %.6g", name, i, states.FloatAt(last+i), want) + } + } + } +} + +// TestAdvectionSquareWaveLimiterBeatsUpwind pins the limiter's reason +// for being: a square pulse carried at CFL 0.9 comes out visibly +// sharper under the Koren flux than under plain upwind, measured as +// the L1 error ratio against the exact shifted pulse. +func TestAdvectionSquareWaveLimiterBeatsUpwind(t *testing.T) { + n := 256 + dx, centres := advectGrid(n) + u0 := make([]float64, n) + for i := range n { + u0[i] = 0.0 + if centres[i] >= 0.3 && centres[i] <= 0.7 { + u0[i] = 1.0 + } + } + u0Arr, err := core.FromFloats(u0, n) + if err != nil { + t.Fatal(err) + } + a := 1.0 + cfl := 0.9 + dt := cfl * dx + const tFinal = 0.2 + upwind, err := IntegrateUpwindAdvection1D(u0Arr, a, dx, tFinal, dt, 2, 0, 0) + if err != nil { + t.Fatalf("upwind: %v", err) + } + koren, err := IntegrateAdvection1D(u0Arr, a, dx, tFinal, dt, 2, 0, 0) + if err != nil { + t.Fatalf("Koren: %v", err) + } + exact := func(x float64) float64 { + if x >= 0.5 && x <= 0.9 { + return 1.0 + } + return 0.0 + } + errUpwind := advectL1(t, upwind, n, exact) + errKoren := advectL1(t, koren, n, exact) + t.Logf("square pulse: upwind L1 %.4g, Koren L1 %.4g, ratio %.2f", errUpwind, errKoren, errUpwind/errKoren) + if errUpwind/errKoren < 2.0 { + t.Fatalf("limiter advantage %.2f, want at least 2x over upwind", errUpwind/errKoren) + } + // Both stay monotone in the sense that matters for a pulse: no + // undershoot below the initial range. + last := n + for i := range n { + for _, states := range []*core.Array{upwind, koren} { + if v := states.FloatAt(last + i); v < -1e-12 || v > 1+1e-12 { + t.Fatalf("scheme left the pulse range: %g at cell %d", v, i) + } + } + } +} + +// TestAdvectionSmoothTransportOrder pins the documented orders: at +// fixed CFL 0.9 the upwind L1 error halves as the grid halves (first +// order) and the Koren error improves second order or better (measured +// ratio past 3 at every pair), staying below the upwind error +// throughout. Three grid levels separate the two rates: a pair alone +// cannot tell a second-order Koren from a degraded one. +func TestAdvectionSmoothTransportOrder(t *testing.T) { + a := 1.0 + cfl := 0.9 + const tFinal = 0.3 + gaussian := func(x float64) float64 { return math.Exp(-math.Pow((x-0.35)/0.1, 2)) } + previousUpwind, previousKoren := 0.0, 0.0 + for _, n := range []int{100, 200, 400} { + dx, centres := advectGrid(n) + u0 := make([]float64, n) + for i := range n { + u0[i] = gaussian(centres[i]) + } + u0Arr, err := core.FromFloats(u0, n) + if err != nil { + t.Fatal(err) + } + dt := cfl * dx / math.Abs(a) + upwind, err := IntegrateUpwindAdvection1D(u0Arr, a, dx, tFinal, dt, 2, 0, 0) + if err != nil { + t.Fatalf("upwind n=%d: %v", n, err) + } + koren, err := IntegrateAdvection1D(u0Arr, a, dx, tFinal, dt, 2, 0, 0) + if err != nil { + t.Fatalf("Koren n=%d: %v", n, err) + } + shift := a * tFinal + exact := func(x float64) float64 { return gaussian(x - shift) } + eUpwind := advectL1(t, upwind, n, exact) + eKoren := advectL1(t, koren, n, exact) + t.Logf("n=%3d: upwind L1 %.3g, Koren L1 %.3g", n, eUpwind, eKoren) + if eKoren > eUpwind { + t.Fatalf("n=%d: Koren error %.3g above upwind %.3g", n, eKoren, eUpwind) + } + if previousUpwind > 0 { + if r := previousUpwind / eUpwind; r < 1.5 || r > 3.0 { + t.Fatalf("n=%d: upwind refinement ratio %.2f, want about 2", n, r) + } + if r := previousKoren / eKoren; r < 3.0 || r > 6.0 { + t.Fatalf("n=%d: Koren refinement ratio %.2f, want the second-order rate the limiter carries", n, r) + } + } + previousUpwind, previousKoren = eUpwind, eKoren + } +} + +// TestAdvectionDiffusionMatchesHeatWhenAZero pins the reduction: with +// a = 0 the advection-diffusion solver performs exactly the +// Crank-Nicolson steps of IntegrateHeat1D, so the two histories agree +// to the last bit. +func TestAdvectionDiffusionMatchesHeatWhenAZero(t *testing.T) { + n := 32 + dx := 1 / float64(n+1) + u0 := make([]float64, n) + for i := range n { + u0[i] = math.Sin(math.Pi * float64(i+1) * dx) + } + u0Arr, err := core.FromFloats(u0, n) + if err != nil { + t.Fatal(err) + } + heat, err := IntegrateHeat1D(u0Arr, 0.05, dx, 0.5, 0.004, 3, 0, 0) + if err != nil { + t.Fatalf("IntegrateHeat1D: %v", err) + } + adv, err := IntegrateAdvectionDiffusion1D(u0Arr, 0, 0.05, dx, 0.5, 0.004, 3, 0, 0) + if err != nil { + t.Fatalf("IntegrateAdvectionDiffusion1D: %v", err) + } + for i := range heat.Len() { + if heat.FloatAt(i) != adv.FloatAt(i) { + t.Fatalf("sample %d differs: heat %.17g, advection-diffusion %.17g", + i, heat.FloatAt(i), adv.FloatAt(i)) + } + } +} + +// TestAdvectionDiffusionConvergence pins the documented orders of the +// combination: the split step is first order in time and second order +// in space, so the coupled refinement at fixed CFL shows the two +// mixed, an L1 error shrinking by roughly 2.5 to 3 per halving. The +// exact solution is the drifting heat kernel +// sqrt(w0/w)·exp(−(x−x0−at)²/w), w = w0 + 4Dt, whose boundary values +// are zero to well below the measured errors. +func TestAdvectionDiffusionConvergence(t *testing.T) { + const ( + a = 0.3 + probD = 0.005 + x0 = 0.3 + tFinal = 0.6 + ) + exact := func(t, x float64) float64 { + w := 0.01 + 4*probD*t + return math.Sqrt(0.01/w) * math.Exp(-math.Pow(x-x0-a*t, 2)/w) + } + run := func(n int, cfl float64) float64 { + dx := 1 / float64(n+1) + u0 := make([]float64, n) + for i := range n { + u0[i] = exact(0, float64(i+1)*dx) + } + u0Arr, err := core.FromFloats(u0, n) + if err != nil { + t.Fatal(err) + } + states, err := IntegrateAdvectionDiffusion1D(u0Arr, a, probD, dx, tFinal, cfl*dx/a, 2, 0, 0) + if err != nil { + t.Fatalf("n=%d: %v", n, err) + } + return advectL1(t, states, n, func(x float64) float64 { return exact(tFinal, x) }) + } + for _, cfl := range []float64{0.9, 0.3} { + previous := 0.0 + for _, n := range []int{64, 128, 256} { + e := run(n, cfl) + t.Logf("CFL %.2f n=%3d: L1 %.4g", cfl, n, e) + if previous > 0 { + if r := previous / e; r < 2.2 || r > 3.7 { + t.Fatalf("CFL %.2f: refinement ratio %.2f at n=%d, want the mixed time-space rate about 2.7", + cfl, r, n) + } + } + previous = e + } + } +} + +func TestAdvectionCFLRefusal(t *testing.T) { + n := 10 + dx := 1 / float64(n+1) + u0, _ := core.FromFloats(make([]float64, n), n) + // CFL = 1.5 for a = 1. + dt := 1.5 * dx + if _, err := IntegrateAdvection1D(u0, 1, dx, 0.1, dt, 2, 0, 0); err == nil || !stringsContains(err, "CFL violated") { + t.Fatalf("Koren CFL violation: %v", err) + } + if _, err := IntegrateUpwindAdvection1D(u0, -1, dx, 0.1, dt, 2, 0, 0); err == nil || !stringsContains(err, "CFL violated") { + t.Fatalf("upwind CFL violation: %v", err) + } + if _, err := IntegrateAdvectionDiffusion1D(u0, 1, 0.01, dx, 0.1, dt, 2, 0, 0); err == nil || !stringsContains(err, "CFL violated") { + t.Fatalf("advection-diffusion CFL violation: %v", err) + } + // The boundary value at the CFL edge is accepted. + if _, err := IntegrateAdvection1D(u0, 1, dx, 0.1, dx, 2, 0, 0); err != nil { + t.Fatalf("CFL = 1 refused: %v", err) + } +} + +func TestAdvectionErrors(t *testing.T) { + if _, err := IntegrateAdvection1D(mustFloats(t, []float64{1, 2, 3}, 3, 1), 1, 0.1, 1, 0.01, 2, 0, 0); err == nil || !stringsContains(err, "rank-1") { + t.Fatalf("a rank-2 initial state: %v", err) + } + if _, err := IntegrateAdvection1D(mustFloats(t, []float64{1, 2, 3}), math.NaN(), 0.1, 1, 0.01, 2, 0, 0); err == nil || !stringsContains(err, "finite") { + t.Fatalf("a NaN speed: %v", err) + } + if _, err := IntegrateAdvection1D(mustFloats(t, []float64{1, 2, 3}), 1, 0.1, 1, 0.01, 2, math.Inf(1), 0); err == nil || !stringsContains(err, "finite") { + t.Fatalf("an infinite ghost value: %v", err) + } + if _, err := IntegrateUpwindAdvection1D(mustFloats(t, []float64{1, 2, 3}), 1, -0.1, 1, 0.01, 2, 0, 0); err == nil || !stringsContains(err, "positive") { + t.Fatalf("a negative spacing: %v", err) + } + if _, err := IntegrateAdvection1D(mustFloats(t, []float64{1, 2, 3}), 1, 0.1, 1, 0.01, 1, 0, 0); err == nil || !stringsContains(err, "two samples") { + t.Fatalf("one sample: %v", err) + } + if _, err := IntegrateAdvection1D(mustFloats(t, []float64{math.NaN()}), 1, 0.1, 1, 0.01, 2, 0, 0); err == nil || !stringsContains(err, "non-finite") { + t.Fatalf("a NaN initial cell: %v", err) + } + if _, err := IntegrateAdvectionDiffusion1D(mustFloats(t, []float64{1, 2, 3}), 1, 0, 0.1, 1, 0.01, 2, 0, 0); err == nil || !stringsContains(err, "diffusivity") { + t.Fatalf("zero diffusivity: %v", err) + } + // A single interior cell is refused by the limiter's stencil need + // only for the limited boundary faces, so n = 1 must still run on + // the upwind path with first order there. + if _, err := IntegrateUpwindAdvection1D(mustFloats(t, []float64{0.5}), 1, 0.1, 0.05, 0.05, 2, 1, 0); err != nil { + t.Fatalf("a single-cell upwind run: %v", err) + } +} + +// TestAdvectionLeftwardTransport pins the mirror branch: with a < 0 +// the inflow is the right boundary, and both schemes must transport a +// monotone leftward ramp without new extrema, with the ghost value +// feeding in from the right. +func TestAdvectionLeftwardTransport(t *testing.T) { + n := 96 + dx, centres := advectGrid(n) + u0 := make([]float64, n) + for i := range n { + u0[i] = centres[i] + } + u0Arr, err := core.FromFloats(u0, n) + if err != nil { + t.Fatal(err) + } + a := -1.0 + dt := 0.9 * dx + for _, limited := range []bool{false, true} { + name := "upwind" + run := func() (*core.Array, error) { + if limited { + name = "Koren" + return IntegrateAdvection1D(u0Arr, a, dx, 0.2, dt, 3, 0, 1) + } + return IntegrateUpwindAdvection1D(u0Arr, a, dx, 0.2, dt, 3, 0, 1) + } + states, err := run() + if err != nil { + t.Fatalf("%s: %v", name, err) + } + for r := range states.Shape()[0] { + for i := range n { + v := states.FloatAt(r*n + i) + if v < -1e-9 || v > 1+1e-9 { + t.Fatalf("%s: sample %d cell %d left the range: %g", name, r, i, v) + } + if i > 0 && states.FloatAt(r*n+i) < states.FloatAt(r*n+i-1)-1e-9 { + t.Fatalf("%s: sample %d grows a new extremum at cell %d: %g after %g", name, r, i, v, states.FloatAt(r*n+i-1)) + } + } + } + // The exact ramp is x + 0.2, cut at the inflow value 1. + last := (states.Shape()[0] - 1) * n + for i := range n { + want := math.Min(centres[i]+0.2, 1) + if d := math.Abs(states.FloatAt(last+i) - want); d > 0.02 { + t.Fatalf("%s: cell %d: %.6g, want %.6g", name, i, states.FloatAt(last+i), want) + } + } + } +} diff --git a/integrate/quad.go b/integrate/quad.go new file mode 100644 index 0000000..8bf6a87 --- /dev/null +++ b/integrate/quad.go @@ -0,0 +1,269 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +import ( + "math" + "sync" +) + +// Numerical quadrature: the definite integral of a function over an +// interval. The library integrates ODEs but until now not plain +// integrals, which every data-reduction pipeline needs. +// +// The scheme is adaptive Gauss-Legendre with an honest internal +// error estimate: every subinterval is evaluated by a 21-point rule +// and a 10-point rule over the same interval, the gap between the two +// is that subinterval's error, and the subinterval with the largest +// error is bisected until the summed error meets the tolerance. The +// 21 and 10 point nodes come from Newton's method on the Legendre +// recurrence, so the whole construction is derived in the library +// rather than imported as a table. Infinite bounds map onto (0, 1) by +// rational substitution before the rule runs; an integrand over an +// infinite interval has to decay to zero for this to converge. + +// QuadratureOptions tunes Integrate. RelTol ≤ 0 means 1e-10, AbsTol +// ≤ 0 means 1e-12, MaxIntervals ≤ 0 means 256. +type QuadratureOptions struct { + RelTol float64 + AbsTol float64 + MaxIntervals int +} + +// gaussLegendreEntry holds the cached nodes and weights of the +// n-point Gauss-Legendre rule over [-1, 1]. +type gaussLegendreEntry struct { + nodes []float64 + weights []float64 +} + +var ( + gaussLegendreCacheMu sync.Mutex + gaussLegendreCache = map[int]gaussLegendreEntry{} +) + +// GaussLegendreNodes returns the nodes and weights of the n-point +// Gauss-Legendre rule over [-1, 1], exact for polynomials up to +// degree 2n−1. Nodes come out ascending. n must be between 1 and 128. +// The nodes are the roots of the n-th Legendre polynomial, found by +// Newton's method on the three-term recurrence, which reaches +// rounding-level accuracy in a handful of iterations per node. The +// returned slices are a shared cache: they must be treated as +// read-only, because a write would poison every later quadrature run +// on the same node count. +func GaussLegendreNodes(n int) (nodes, weights []float64, err error) { + if n < 1 || n > 128 { + return nil, nil, base.Errf("GaussLegendreNodes: n must be between 1 and 128, got %d", n) + } + gaussLegendreCacheMu.Lock() + entry, ok := gaussLegendreCache[n] + gaussLegendreCacheMu.Unlock() + if ok { + return entry.nodes, entry.weights, nil + } + nodes = make([]float64, n) + weights = make([]float64, n) + for i := range n { + // The trigonometric starting guess separates the roots well + // enough that Newton never jumps to a neighbour. + x := -math.Cos(math.Pi * (float64(i) + 0.75) / (float64(n) + 0.5)) + var p, dp float64 + for range 100 { + p, dp = legendrePair(n, x) + dx := p / dp + x -= dx + if math.Abs(dx) <= 1e-16*(1+math.Abs(x)) { + break + } + } + if math.Abs(p) > 1e-10 { + return nil, nil, base.Errf("GaussLegendreNodes: Newton failed to converge for n=%d", n) + } + nodes[i] = x + weights[i] = 2 / ((1 - x*x) * dp * dp) + } + entry = gaussLegendreEntry{nodes: nodes, weights: weights} + gaussLegendreCacheMu.Lock() + gaussLegendreCache[n] = entry + gaussLegendreCacheMu.Unlock() + return nodes, weights, nil +} + +// legendrePair evaluates P_n(x) and P_n'(x) by the three-term +// recurrence and its derivative identity. +func legendrePair(n int, x float64) (p, dp float64) { + p = 1 + if n == 0 { + return 1, 0 + } + pm := 0.0 + for k := 1; k <= n; k++ { + pm, p = p, ((2*float64(k)-1)*x*p-float64(k-1)*pm)/float64(k) + } + dp = float64(n) * (x*p - pm) / (x*x - 1) + return p, dp +} + +// IntegrateFunction returns the definite integral of f over [a, b] with an +// estimate of the absolute error. Infinite bounds are accepted: a = −∞ +// or b = +∞ (or both) integrate over the whole tail under a rational +// substitution, which asks the integrand to decay to zero. A reversed +// interval (a > b) integrates in the negative direction. A tolerance +// that cannot be met within MaxIntervals subintervals is an error, +// never a silent approximation. +// +// Every sampled scheme has one blind spot: a feature entirely inside +// the gaps of the first rule's nodes, say a peak far narrower than +// (b−a)/n, produces small values everywhere it samples and is missed +// with a small error estimate. Split known sharp features into their +// own IntegrateFunction calls. +func IntegrateFunction(f func(x float64) (float64, error), a, b float64, opts QuadratureOptions) (float64, float64, error) { + if opts.RelTol <= 0 { + opts.RelTol = 1e-10 + } + if opts.AbsTol <= 0 { + opts.AbsTol = 1e-12 + } + if opts.MaxIntervals <= 0 { + opts.MaxIntervals = 256 + } + if math.IsNaN(a) || math.IsNaN(b) { + return 0, 0, base.Errf("Integrate: bounds must not be NaN") + } + sign := 1.0 + if b < a { + a, b = b, a + sign = -1 + } + if a == b { + return 0, 0, nil + } + // Map every bound combination onto a plain finite interval with a + // wrapped integrand carrying the substitution's Jacobian. + lo, hi := 0.0, 1.0 + g := f + switch { + case a == math.Inf(-1) && b == math.Inf(1): + g = func(t float64) (float64, error) { + u := t - 0.5 + den := 0.25 - u*u + if den <= 0 { + return 0, nil + } + x := 2 * u / den + fx, ferr := f(x) + if ferr != nil { + return 0, ferr + } + return fx * (0.5 + 2*u*u) / (den * den), nil + } + case a == math.Inf(-1): + g = func(t float64) (float64, error) { + x := b - t/(1-t) + fx, ferr := f(x) + if ferr != nil { + return 0, ferr + } + return fx / (1 - t) / (1 - t), nil + } + case b == math.Inf(1): + g = func(t float64) (float64, error) { + x := a + t/(1-t) + fx, ferr := f(x) + if ferr != nil { + return 0, ferr + } + return fx / (1 - t) / (1 - t), nil + } + default: + lo, hi = a, b + } + + nodes10, w10, err := GaussLegendreNodes(10) + if err != nil { + return 0, 0, err + } + nodes21, w21, err := GaussLegendreNodes(21) + if err != nil { + return 0, 0, err + } + rule := func(g func(float64) (float64, error), l, r float64, xs, ws []float64) (float64, error) { + mid, half := (l+r)/2, (r-l)/2 + total := 0.0 + for i := range xs { + fx, ferr := g(mid + half*xs[i]) + if ferr != nil { + return 0, ferr + } + total += ws[i] * fx + } + return total * half, nil + } + + // Each leaf carries its 21-point value and its error, the gap to + // the 10-point rule over the same interval. + type leaf struct { + l, r, value, err float64 + } + measure := func(l, r float64) (leaf, error) { + v21, err := rule(g, l, r, nodes21, w21) + if err != nil { + return leaf{}, err + } + v10, err := rule(g, l, r, nodes10, w10) + if err != nil { + return leaf{}, err + } + // An infinite integrand value would poison errSum with NaN and + // slip through the loop condition (NaN comparisons are false), + // publishing a bogus integral, so both non-finite kinds are + // refused exactly as IntegrateND refuses them. + if math.IsNaN(v21) || math.IsNaN(v10) || math.IsInf(v21, 0) || math.IsInf(v10, 0) { + return leaf{}, base.Errf("the integrand returned a non-finite value on [%g, %g]", l, r) + } + return leaf{l: l, r: r, value: v21, err: math.Abs(v21 - v10)}, nil + } + + first, err := measure(lo, hi) + if err != nil { + return 0, 0, base.Errf("Integrate: %w", err) + } + leaves := []leaf{first} + budget := func() (total, errSum float64) { + for _, s := range leaves { + total += s.value + errSum += s.err + } + return total, errSum + } + value, errSum := budget() + for errSum > math.Max(opts.AbsTol, opts.RelTol*math.Abs(value)) { + if len(leaves) >= opts.MaxIntervals { + return 0, 0, base.Errf("Integrate: error estimate %g exceeds the tolerance within %d subintervals", + errSum, opts.MaxIntervals) + } + // Bisect the leaf that contributes the most error. + worst := 0 + for i, s := range leaves { + if s.err > leaves[worst].err { + worst = i + } + } + w := leaves[worst] + left, lerr := measure(w.l, (w.l+w.r)/2) + if lerr != nil { + return 0, 0, base.Errf("Integrate: %w", lerr) + } + right, rerr := measure((w.l+w.r)/2, w.r) + if rerr != nil { + return 0, 0, base.Errf("Integrate: %w", rerr) + } + leaves[worst] = left + leaves = append(leaves, right) + value, errSum = budget() + } + return sign * value, errSum, nil +} diff --git a/integrate/quad_cub_bench_test.go b/integrate/quad_cub_bench_test.go new file mode 100644 index 0000000..1a1357d --- /dev/null +++ b/integrate/quad_cub_bench_test.go @@ -0,0 +1,60 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" +) + +// Benchmarks for the adaptive cubature's low-dimensional workloads, +// the ones whose bisection loop turns often enough against a cheap +// integrand that the worst-box selection is a visible share of the +// cost. The three-dimensional workload lives in bench_test.go. + +// BenchmarkIntegrateND1D runs the oscillatory task of +// TestCubatureMatches1D: a degenerate dimension whose boxes bisect +// for sixteen evaluations apiece, so the selection runs at its +// cheapest per box. +func BenchmarkIntegrateND1D(b *testing.B) { + f := func(x []float64) float64 { return math.Exp(-x[0]) * math.Cos(3*x[0]) } + lo := []float64{0} + hi := []float64{5} + b.ReportAllocs() + for b.Loop() { + if _, err := IntegrateND(f, lo, hi, CubatureOptions{Tolerance: 1e-12}); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkIntegrateND2D runs the Gaussian of TestCubatureGaussian at +// its test tolerance: the subdivision reaches into the thousands of +// boxes, the regime where the worst-box choice dominates over the +// arithmetic inside each box. +func BenchmarkIntegrateND2D(b *testing.B) { + f := func(x []float64) float64 { return math.Exp(-x[0]*x[0] - x[1]*x[1]) } + lo := []float64{-3, -3} + hi := []float64{3, 3} + b.ReportAllocs() + for b.Loop() { + if _, err := IntegrateND(f, lo, hi, CubatureOptions{Tolerance: 1e-11}); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkIntegrateFilon runs the oscillatory integral over enough +// carrier wavelengths that the automatic panel count keeps the weights +// busy: the amplitude stays smooth, so the per-panel cost is the node +// sweep and the phase rotation. +func BenchmarkIntegrateFilon(b *testing.B) { + f := func(x float64) (float64, error) { return math.Exp(-0.1 * x), nil } + b.ReportAllocs() + for b.Loop() { + if _, _, err := IntegrateFilon(f, 0, 100, 40*math.Pi, FilonOptions{}); err != nil { + b.Fatal(err) + } + } +} diff --git a/integrate/quad_test.go b/integrate/quad_test.go new file mode 100644 index 0000000..3a0dbcc --- /dev/null +++ b/integrate/quad_test.go @@ -0,0 +1,201 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +import ( + "math" + "testing" +) + +func quadValue(t *testing.T, f func(float64) (float64, error), a, b float64, opts QuadratureOptions) (float64, float64) { + t.Helper() + v, errEst, err := IntegrateFunction(f, a, b, opts) + if err != nil { + t.Fatalf("Integrate: %v", err) + } + return v, errEst +} + +func TestIntegratePolynomials(t *testing.T) { + // Degree 14 on 21 points: exact by construction. + v, _ := quadValue(t, func(x float64) (float64, error) { return math.Pow(x, 7), nil }, 0, 1, QuadratureOptions{}) + if math.Abs(v-1.0/8) > 1e-14 { + t.Fatalf("integral of x^7 = %.16g, want 0.125", v) + } + v, _ = quadValue(t, func(x float64) (float64, error) { return math.Pow(x, 14), nil }, -1, 1, QuadratureOptions{}) + if math.Abs(v-2.0/15) > 1e-14 { + t.Fatalf("integral of x^14 = %.16g, want %g", v, 2.0/15) + } +} + +func TestIntegrateSmooth(t *testing.T) { + cases := []struct { + name string + f func(float64) (float64, error) + a, b float64 + want float64 + }{ + {"sine over a period", func(x float64) (float64, error) { return math.Sin(x), nil }, 0, math.Pi, 2}, + {"exponential", func(x float64) (float64, error) { return math.Exp(x), nil }, -1, 1, 2 * math.Sinh(1)}, + {"arctangent derivative", func(x float64) (float64, error) { return 1 / (1 + x*x), nil }, 0, 1, math.Pi / 4}, + {"gaussian", func(x float64) (float64, error) { return math.Exp(-x * x), nil }, 0, 5, 0.5 * math.Sqrt(math.Pi)}, + } + for _, c := range cases { + v, est := quadValue(t, c.f, c.a, c.b, QuadratureOptions{}) + if math.Abs(v-c.want) > 1e-11 { + t.Fatalf("%s: %.14g, want %.14g", c.name, v, c.want) + } + if math.Abs(v-c.want) > 10*est+1e-14 { + t.Fatalf("%s: error estimate %g understates the true error %g", c.name, est, math.Abs(v-c.want)) + } + } +} + +func TestIntegrateOscillatory(t *testing.T) { + v, _ := quadValue(t, func(x float64) (float64, error) { return math.Sin(x), nil }, 0, 10*math.Pi, QuadratureOptions{}) + if math.Abs(v) > 1e-9 { + t.Fatalf("ten sine periods integrate to %g, want 0", v) + } + v, _ = quadValue(t, func(x float64) (float64, error) { return math.Sin(30*x) * math.Exp(-x), nil }, + 0, 40, QuadratureOptions{RelTol: 1e-9}) + want := 30.0 / 901.0 // exact over [0, inf): 30/(30^2 + 1) + if math.Abs(v-want) > 1e-8 { + t.Fatalf("damped oscillation = %.14g, want %.14g", v, want) + } +} + +func TestIntegrateSharpPeak(t *testing.T) { + // A Lorentzian a hundred times narrower than the interval forces + // deep subdivision; the answer must land on the exact closed form + // 2·arctan(100). (A peak narrower than the first rule's node + // spacing would be invisible to any sampled scheme, which the doc + // contract states.) + v, _ := quadValue(t, func(x float64) (float64, error) { + return 10 / (1 + 100*x*x), nil + }, -10, 10, QuadratureOptions{RelTol: 1e-10}) + if want := 2 * math.Atan(100); math.Abs(v-want) > 1e-9 { + t.Fatalf("narrow Lorentzian = %.12g, want %.12g", v, want) + } +} + +func TestIntegrateInfiniteBounds(t *testing.T) { + cases := []struct { + name string + f func(float64) (float64, error) + a, b float64 + want float64 + }{ + {"exponential tail", func(x float64) (float64, error) { return math.Exp(-x), nil }, 0, math.Inf(1), 1}, + {"cauchy tail", func(x float64) (float64, error) { return 1 / (1 + x*x), nil }, 0, math.Inf(1), math.Pi / 2}, + {"full gaussian", func(x float64) (float64, error) { return math.Exp(-x * x), nil }, math.Inf(-1), math.Inf(1), math.Sqrt(math.Pi)}, + {"negative tail", func(x float64) (float64, error) { return math.Exp(x), nil }, math.Inf(-1), 0, 1}, + } + for _, c := range cases { + v, _ := quadValue(t, c.f, c.a, c.b, QuadratureOptions{}) + if math.Abs(v-c.want) > 1e-9 { + t.Fatalf("%s: %.14g, want %.14g", c.name, v, c.want) + } + } +} + +func TestIntegrateOrientation(t *testing.T) { + sin := func(x float64) (float64, error) { return math.Sin(x), nil } + forward, _ := quadValue(t, sin, 0, math.Pi, QuadratureOptions{}) + backward, _ := quadValue(t, sin, math.Pi, 0, QuadratureOptions{}) + if math.Abs(backward+forward) > 1e-14 { + t.Fatalf("reversed integral = %g, want %g", backward, -forward) + } + if v, _ := quadValue(t, sin, 2, 2, QuadratureOptions{}); v != 0 { + t.Fatalf("empty interval = %g, want 0", v) + } +} + +func TestIntegrateErrors(t *testing.T) { + boom := func(x float64) (float64, error) { + if x > 0.5 { + return 0, base.Errf("integrand failed at %g", x) + } + return 1, nil + } + if _, _, err := IntegrateFunction(boom, 0, 1, QuadratureOptions{}); err == nil { + t.Fatal("integrand error: want an error") + } + if _, _, err := IntegrateFunction(func(x float64) (float64, error) { return x, nil }, + math.NaN(), 1, QuadratureOptions{}); err == nil { + t.Fatal("NaN bound: want an error") + } + // A flat integrand over an infinite interval diverges: the + // adaptation must report it, not return a number. + if _, _, err := IntegrateFunction(func(x float64) (float64, error) { return 1, nil }, + 0, math.Inf(1), QuadratureOptions{MaxIntervals: 8}); err == nil { + t.Fatal("divergent integral: want an error") + } + // A tolerance no budget can meet must be reported. + if _, _, err := IntegrateFunction(func(x float64) (float64, error) { return math.Exp(-1000 * (x - 0.5) * (x - 0.5)), nil }, + 0, 1, QuadratureOptions{RelTol: 1e-16, AbsTol: 0, MaxIntervals: 4}); err == nil { + t.Fatal("exhausted budget: want an error") + } +} + +func TestGaussLegendreNodes(t *testing.T) { + nodes, weights, err := GaussLegendreNodes(8) + if err != nil { + t.Fatalf("GaussLegendreNodes: %v", err) + } + total := 0.0 + for i := range 8 { + total += weights[i] + if math.Abs(nodes[i]+nodes[7-i]) > 1e-14 { + t.Fatalf("nodes %d and %d are not symmetric", i, 7-i) + } + if i > 0 && nodes[i] <= nodes[i-1] { + t.Fatalf("nodes are not ascending at %d", i) + } + } + if math.Abs(total-2) > 1e-14 { + t.Fatalf("weight sum = %.16g, want 2", total) + } + // A 16th degree polynomial integrates exactly on 8 points. + moment := 0.0 + for i := range 8 { + moment += weights[i] * math.Pow(nodes[i], 14) + } + if math.Abs(moment-2.0/15) > 1e-14 { + t.Fatalf("moment of x^14 = %.16g, want %g", moment, 2.0/15) + } + // The cache returns the same arrays. + again, _, err := GaussLegendreNodes(8) + if err != nil { + t.Fatalf("cached call: %v", err) + } + if &again[0] != &nodes[0] { + t.Fatal("cached nodes are not the cached arrays") + } + if _, _, err := GaussLegendreNodes(0); err == nil { + t.Fatal("n = 0: want an error") + } + if _, _, err := GaussLegendreNodes(129); err == nil { + t.Fatal("n = 129: want an error") + } +} + +// TestIntegrateFunctionNaNIsError pins the NaN contract: an integrand +// that yields NaN makes IntegrateFunction return an error instead of +// a silent NaN integral (the first estimate and every bisected leaf +// are checked). +func TestIntegrateFunctionNaNIsError(t *testing.T) { + f := func(x float64) (float64, error) { + if math.Abs(x) > 0.25 { + return math.NaN(), nil + } + return 1, nil + } + if _, _, err := IntegrateFunction(f, -1, 1, QuadratureOptions{}); err == nil { + t.Fatal("expected an error for an integrand that returns NaN") + } +} diff --git a/integrate/symplectic.go b/integrate/symplectic.go new file mode 100644 index 0000000..c1a63b6 --- /dev/null +++ b/integrate/symplectic.go @@ -0,0 +1,129 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Symplectic integration for separable Hamiltonian systems, where the +// energy splits as H(q, p) = T(p) + V(q): Newtonian mechanics, N-body +// gravity, molecular dynamics. The adaptive Runge-Kutta pair that +// drives IntegrateODE is accurate per step but dissipates energy +// systematically, so a two-hundred-period orbit spirals inward or +// outward; the leapfrog structure below is symplectic, which means it +// conserves a shadow Hamiltonian exactly and keeps the true energy +// oscillating in a bounded band forever. That long-time fidelity, not +// per-step accuracy, is what separates integrators for celestial +// mechanics and molecular dynamics from general-purpose ones. +// +// The scheme is velocity Verlet, a kick-drift-kick leapfrog of second +// order: half a momentum kick, a full position drift, another half +// kick with the force at the new position. Unit masses are assumed +// (p is the velocity); scale the momentum by the masses beforehand or +// fold them into the acceleration. + +// IntegrateVerlet integrates a separable Hamiltonian system with unit +// masses over an even time grid: accel returns the acceleration +// −∂V/∂q at a position, q0 and p0 are the initial position and +// momentum (velocity), and steps fixes the number of equal steps, so +// the i-th returned pair sits at t0 + i·h with h = (t1−t0)/steps. +// positions[0] is q0 and momenta[0] is p0. The state may be float64 +// or float32, read through per-element accessors, and the returned +// arrays are float64. The step size stays fixed by design: +// adaptivity would destroy the symplectic property the method exists +// for. A negative or zero-length acceleration vector, a mismatched +// pair, a complex or int state, or an empty state is an error. +func IntegrateVerlet(accel func(q *core.Array) (*core.Array, error), + t0, t1 float64, q0, p0 *core.Array, steps int) (positions, momenta []*core.Array, err error) { + if steps <= 0 { + return nil, nil, base.Errf("IntegrateVerlet: steps must be ≥ 1, got %d", steps) + } + if q0.Dtype() == core.Complex || p0.Dtype() == core.Complex { + return nil, nil, base.Errf("IntegrateVerlet: complex states are not supported") + } + // The whole integer class follows Int into the standing refusal: + // bool and the narrow widths carry a discrete state, which has no + // place in a continuous integrator, and the wording is Int's own. + if integerState(q0.Dtype()) || integerState(p0.Dtype()) { + return nil, nil, base.Errf("IntegrateVerlet: int states cannot integrate, use float or float32 states") + } + if q0.NDim() != 1 || p0.NDim() != 1 || q0.Len() != p0.Len() { + return nil, nil, base.Errf("IntegrateVerlet: position and momentum must be vectors of equal length, got %s and %s", + base.ShapeText(q0.Shape()), base.ShapeText(p0.Shape())) + } + n := q0.Len() + if n == 0 { + return nil, nil, base.Errf("IntegrateVerlet: the state must not be empty") + } + // Read the initial state element-wise: RawFloats backs float64 + // payloads only, so a float32 state would come through as nil. + // A non-finite entry is refused up front: it would propagate + // through every kick and drift silently. + q := make([]float64, n) + p := make([]float64, n) + for i := range n { + q[i] = q0.FloatAt(i) + p[i] = p0.FloatAt(i) + if math.IsNaN(q[i]) || math.IsInf(q[i], 0) { + return nil, nil, base.Errf("IntegrateVerlet: q0 holds the non-finite value %g at %d", q[i], i) + } + if math.IsNaN(p[i]) || math.IsInf(p[i], 0) { + return nil, nil, base.Errf("IntegrateVerlet: p0 holds the non-finite value %g at %d", p[i], i) + } + } + a := make([]float64, n) + // One cached read-only view serves every acceleration call: the + // position slice is the run's own buffer, stable for the whole + // integration, so the wrapper is built once per run. + var views odeViews + eval := func(x []float64, out []float64) error { + v, err := accel(views.of(x)) + if err != nil { + return base.Errf("IntegrateVerlet: %w", err) + } + if v.NDim() != 1 || v.Len() != n { + return base.Errf("IntegrateVerlet: accel returned shape %s, want a vector of length %d", + base.ShapeText(v.Shape()), n) + } + readVector(out, v) + // Like RK4: a non-finite acceleration would flow through the + // kicks silently, and the published trajectory would be NaN + // with a nil error. + for i := range n { + if math.IsNaN(out[i]) || math.IsInf(out[i], 0) { + return base.Errf("IntegrateVerlet: accel returned the non-finite value %g at coordinate %d", out[i], i) + } + } + return nil + } + if err := eval(q, a); err != nil { + return nil, nil, err + } + h := (t1 - t0) / float64(steps) + positions = make([]*core.Array, steps+1) + momenta = make([]*core.Array, steps+1) + positions[0] = arrayFromVector(q) + momenta[0] = arrayFromVector(p) + for s := 1; s <= steps; s++ { + // Kick, drift, kick: two half kicks bracket the drift, so the + // force is evaluated once per step. + for i := range n { + p[i] += h / 2 * a[i] + q[i] += h * p[i] + } + if err := eval(q, a); err != nil { + return nil, nil, err + } + for i := range n { + p[i] += h / 2 * a[i] + } + positions[s] = arrayFromVector(q) + momenta[s] = arrayFromVector(p) + } + return positions, momenta, nil +} diff --git a/integrate/symplectic2.go b/integrate/symplectic2.go new file mode 100644 index 0000000..388d101 --- /dev/null +++ b/integrate/symplectic2.go @@ -0,0 +1,284 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/optim" +) + +// Higher-order symplectic integration, extending the velocity-Verlet +// leapfrog of symplectic.go in the two directions it cannot go by +// itself: to fourth order while staying explicit (Yoshida's +// composition of three leapfrog sub-steps), and to non-separable +// Hamiltonians at all (the implicit midpoint rule, whose implicit +// stage is a root find per step). Both keep the leapfrog's reason for +// being: the exact conservation of a shadow Hamiltonian, so the true +// energy oscillates in a bounded band instead of drifting, over +// arbitrarily long runs, at a fixed step by design. + +// yoshidaW1 and yoshidaW0 are Yoshida's triple-jump weights: the +// composition V(w1·h)·V(w0·h)·V(w1·h) of three velocity-Verlet +// sub-steps is fourth order exactly when w0 + 2·w1 = 1, and the +// classic choice w1 = 1/(2 − ∛2), w0 = −∛2/(2 − ∛2) satisfies that +// identity with w1 positive and w0 negative (the middle sub-step runs +// backwards in time, which is what buys the order). +var ( + yoshidaW1 = 1 / (2 - math.Cbrt(2)) + yoshidaW0 = -math.Cbrt(2) * yoshidaW1 +) + +// IntegrateYoshida4 integrates a separable Hamiltonian system with +// unit masses over an even time grid by Yoshida's fourth-order +// composition of the kick-drift-kick leapfrog: three velocity-Verlet +// sub-steps of widths w1·h, w0·h and w1·h per step, with the weights +// above. The contract is IntegrateVerlet's: accel returns the +// acceleration −∂V/∂q at a position, q0 and p0 are the initial +// position and momentum (velocity), steps fixes the number of equal +// steps h = (t1−t0)/steps, positions[s] and momenta[s] sit at +// t0 + s·h, and the state may be float64 or float32. The step size +// stays fixed by design. The error contract is IntegrateVerlet's +// too: a mismatched pair, a complex or int state, an empty state or +// a non-finite acceleration is an error, never a silently corrupted +// trajectory. +func IntegrateYoshida4(accel func(q *core.Array) (*core.Array, error), + t0, t1 float64, q0, p0 *core.Array, steps int) (positions, momenta []*core.Array, err error) { + const name = "IntegrateYoshida4" + if steps <= 0 { + return nil, nil, base.Errf("%s: steps must be ≥ 1, got %d", name, steps) + } + q, p, n, verr := verletPrelude(name, q0, p0) + if verr != nil { + return nil, nil, verr + } + eval := verletForce(name, accel, n) + a := make([]float64, n) + if err := eval(q, a); err != nil { + return nil, nil, err + } + h := (t1 - t0) / float64(steps) + positions = make([]*core.Array, steps+1) + momenta = make([]*core.Array, steps+1) + positions[0] = arrayFromVector(q) + momenta[0] = arrayFromVector(p) + weights := [3]float64{yoshidaW1, yoshidaW0, yoshidaW1} + for s := 1; s <= steps; s++ { + // Three kick-drift-kick sub-steps. The acceleration that ends + // one sub-step is exactly the one the next sub-step's first + // kick needs, so the whole step costs three force evaluations. + for _, w := range weights { + tau := w * h + for i := range n { + p[i] += tau / 2 * a[i] + q[i] += tau * p[i] + } + if err := eval(q, a); err != nil { + return nil, nil, err + } + for i := range n { + p[i] += tau / 2 * a[i] + } + } + positions[s] = arrayFromVector(q) + momenta[s] = arrayFromVector(p) + } + return positions, momenta, nil +} + +// verletPrelude is the validation the Verlet family shares: the +// dtype, shape and finiteness gates of IntegrateVerlet and the +// element-wise read of the initial state (RawFloats backs float64 +// payloads only, so a float32 state would come through as nil). It +// returns the working position and momentum and the state length. +func verletPrelude(name string, q0, p0 *core.Array) (q, p []float64, n int, err error) { + if q0.Dtype() == core.Complex || p0.Dtype() == core.Complex { + return nil, nil, 0, base.Errf("%s: complex states are not supported", name) + } + // The whole integer class follows Int into the standing refusal: + // bool and the narrow widths carry a discrete state, which has no + // place in a continuous integrator, and the wording is Int's own. + if integerState(q0.Dtype()) || integerState(p0.Dtype()) { + return nil, nil, 0, base.Errf("%s: int states cannot integrate, use float or float32 states", name) + } + if q0.NDim() != 1 || p0.NDim() != 1 || q0.Len() != p0.Len() { + return nil, nil, 0, base.Errf("%s: position and momentum must be vectors of equal length, got %s and %s", + name, base.ShapeText(q0.Shape()), base.ShapeText(p0.Shape())) + } + n = q0.Len() + if n == 0 { + return nil, nil, 0, base.Errf("%s: the state must not be empty", name) + } + q = make([]float64, n) + p = make([]float64, n) + for i := range n { + q[i] = q0.FloatAt(i) + p[i] = p0.FloatAt(i) + if math.IsNaN(q[i]) || math.IsInf(q[i], 0) { + return nil, nil, 0, base.Errf("%s: q0 holds the non-finite value %g at %d", name, q[i], i) + } + if math.IsNaN(p[i]) || math.IsInf(p[i], 0) { + return nil, nil, 0, base.Errf("%s: p0 holds the non-finite value %g at %d", name, p[i], i) + } + } + return q, p, n, nil +} + +// integerState reports whether dt is one of the integer-class state +// dtypes the symplectic family refuses: bool and the narrow integer +// widths follow Int into the standing "int states cannot integrate" +// refusal, exactly as the round's follow-Int rule requires. The float +// dtypes, float16 included, keep the treatment they carry today. +func integerState(dt core.Dtype) bool { + switch dt { + case core.Bool, core.Int, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32: + return true + } + return false +} + +// verletForce wraps the acceleration callback the way the Verlet +// family evaluates it: shape-checked, non-finite-refused, and read +// into a plain float64 buffer. One cached view per wrapper serves +// every call: the family evaluates the force at the same position +// buffer throughout a run, so the wrapper is built once. +func verletForce(name string, accel func(q *core.Array) (*core.Array, error), n int) func(x []float64, out []float64) error { + var views odeViews + return func(x []float64, out []float64) error { + v, err := accel(views.of(x)) + if err != nil { + return base.Errf("%s: %w", name, err) + } + if v.NDim() != 1 || v.Len() != n { + return base.Errf("%s: accel returned shape %s, want a vector of length %d", + name, base.ShapeText(v.Shape()), n) + } + readVector(out, v) + // A non-finite acceleration would flow through the kicks + // silently, and the published trajectory would be NaN with a + // nil error. + for i := range n { + if math.IsNaN(out[i]) || math.IsInf(out[i], 0) { + return base.Errf("%s: accel returned the non-finite value %g at coordinate %d", name, out[i], i) + } + } + return nil + } +} + +// MidpointOptions tunes the per-step implicit stage of +// IntegrateMidpoint. Tolerance ≤ 0 means 1e-13 (the stage solve is a +// root find, and the energy band the rule is famous for wants it +// tight); MaxIterations ≤ 0 means 100. +type MidpointOptions struct { + Tolerance float64 + MaxIterations int +} + +// IntegrateMidpoint integrates the general Hamiltonian flow +// dz/dt = J·∇H(z), z = (q, p), by the implicit midpoint rule +// z_{n+1} = z_n + h·J·∇H((z_n + z_{n+1})/2) over an even time grid. +// gradH returns the gradient of H as the stacked vector +// (∂H/∂q, ∂H/∂p); H itself never has to be separable, which is the +// rule's claim over the leapfrog family. Each step's implicit stage +// is solved by the library's damped Newton root find for systems +// (optim.FindRootSystem under the given options), seeded with the +// current state, so the per-step work is a handful of gradient +// evaluations. Equilibria map to themselves exactly, and every +// quadratic invariant of the flow, H included when H is quadratic, +// is conserved to rounding. The contract otherwise mirrors +// IntegrateVerlet: q0 and p0 are the initial position and momentum, +// steps fixes the number of equal steps h = (t1−t0)/steps, +// positions[s] and momenta[s] sit at t0 + s·h, the step size stays +// fixed by design, and a mismatched pair, a complex or int state, an +// empty state, a non-finite gradient or a root find that cannot +// converge is an error, never a silent answer. +func IntegrateMidpoint(gradH func(z *core.Array) (*core.Array, error), + t0, t1 float64, q0, p0 *core.Array, steps int, opts MidpointOptions) (positions, momenta []*core.Array, err error) { + const name = "IntegrateMidpoint" + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-13 + } + if opts.MaxIterations <= 0 { + opts.MaxIterations = 100 + } + if steps <= 0 { + return nil, nil, base.Errf("%s: steps must be ≥ 1, got %d", name, steps) + } + q, p, n, verr := verletPrelude(name, q0, p0) + if verr != nil { + return nil, nil, verr + } + m := 2 * n + z := make([]float64, m) + copy(z, q) + copy(z[n:], p) + // One cached read-only view serves every gradient call: mid is the + // run's own stable buffer, so the wrapper is built once. + var views odeViews + // The gradient evaluator with the midpoint rule's finiteness gate: + // a non-finite gradient would flow through the stage equation + // silently. + gradient := func(gz []float64, out []float64) error { + v, err := gradH(views.of(gz)) + if err != nil { + return base.Errf("%s: %w", name, err) + } + if v.NDim() != 1 || v.Len() != m { + return base.Errf("%s: gradH returned shape %s, want a vector of length %d", + name, base.ShapeText(v.Shape()), m) + } + readVector(out, v) + for i := range m { + if math.IsNaN(out[i]) || math.IsInf(out[i], 0) { + return base.Errf("%s: gradH returned the non-finite value %g at coordinate %d", name, out[i], i) + } + } + return nil + } + mid := make([]float64, m) + grad := make([]float64, m) + h := (t1 - t0) / float64(steps) + // The stage residual: x is the candidate z_{n+1}, and the flow + // J·∇H flips the two halves with a sign: q' = ∂H/∂p, p' = −∂H/∂q. + // FindRootSystem copies its input before every residual call and + // reads the returned residual before the next one, so reusing mid, + // grad and out across calls is the same arithmetic the fresh + // buffers would give. + out := make([]float64, m) + stage := func(x *core.Array) (*core.Array, error) { + for i := range m { + mid[i] = (z[i] + x.FloatAt(i)) / 2 + } + if err := gradient(mid, grad); err != nil { + return nil, err + } + for i := range n { + out[i] = x.FloatAt(i) - z[i] - h*grad[n+i] + out[n+i] = x.FloatAt(n+i) - z[n+i] + h*grad[i] + } + return wrapVector(out), nil + } + positions = make([]*core.Array, steps+1) + momenta = make([]*core.Array, steps+1) + positions[0] = arrayFromVector(z[:n]) + momenta[0] = arrayFromVector(z[n:]) + for s := 1; s <= steps; s++ { + solution, _, rerr := optim.FindRootSystem(stage, wrapVector(z), optim.RootSystemOptions{ + Tolerance: opts.Tolerance, + MaxIterations: opts.MaxIterations, + }) + if rerr != nil { + return nil, nil, base.Errf("%s: %w", name, rerr) + } + for i := range m { + z[i] = solution.FloatAt(i) + } + positions[s] = arrayFromVector(z[:n]) + momenta[s] = arrayFromVector(z[n:]) + } + return positions, momenta, nil +} diff --git a/integrate/symplectic2_test.go b/integrate/symplectic2_test.go new file mode 100644 index 0000000..79a0152 --- /dev/null +++ b/integrate/symplectic2_test.go @@ -0,0 +1,370 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestYoshidaWeightsIdentity pins the triple-jump weights: the +// composition V(w1·h)·V(w0·h)·V(w1·h) is fourth order exactly when +// w0 + 2·w1 = 1, with w1 positive and w0 negative. +func TestYoshidaWeightsIdentity(t *testing.T) { + if math.Abs(yoshidaW0+2*yoshidaW1-1) > 1e-15 { + t.Fatalf("w0 + 2·w1 = %.17g, want 1", yoshidaW0+2*yoshidaW1) + } + if yoshidaW1 <= 0 || yoshidaW0 >= 0 { + t.Fatalf("weights %g and %g, want w1 positive and w0 negative", yoshidaW0, yoshidaW1) + } + c := math.Cbrt(2) + if math.Abs(yoshidaW1-1/(2-c)) > 1e-15 || math.Abs(yoshidaW0+c/(2-c)) > 1e-15 { + t.Fatalf("weights %.17g and %.17g are not the triple-jump choice", yoshidaW0, yoshidaW1) + } +} + +// harmonicEnergy returns the energy of a one-dimensional harmonic +// state. +func harmonicEnergy(q, p []float64, omega float64) float64 { + return 0.5 * (p[0]*p[0] + omega*omega*q[0]*q[0]) +} + +// yoshidaEnergyBand integrates the harmonic oscillator over periods +// and returns the largest energy deviation from the initial one. +func yoshidaEnergyBand(t *testing.T, perPeriod, periods int, verlet bool) float64 { + t.Helper() + const omega = 1.0 + accel := func(q *core.Array) (*core.Array, error) { + return core.MulF(q, -omega*omega), nil + } + q0 := mustFloats(t, []float64{1}) + p0 := mustFloats(t, []float64{0}) + steps := perPeriod * periods + var positions, momenta []*core.Array + var err error + if verlet { + positions, momenta, err = IntegrateVerlet(accel, 0, float64(periods)*2*math.Pi, q0, p0, steps) + } else { + positions, momenta, err = IntegrateYoshida4(accel, 0, float64(periods)*2*math.Pi, q0, p0, steps) + } + if err != nil { + t.Fatalf("integrator: %v", err) + } + e0 := harmonicEnergy([]float64{positions[0].FloatAt(0)}, []float64{momenta[0].FloatAt(0)}, omega) + worst := 0.0 + q := make([]float64, 1) + p := make([]float64, 1) + for s := range positions { + q[0] = positions[s].FloatAt(0) + p[0] = momenta[s].FloatAt(0) + if d := math.Abs(harmonicEnergy(q, p, omega) - e0); d > worst { + worst = d + } + } + return worst +} + +// TestYoshidaHarmonicFourthOrder pins the order: the harmonic +// oscillator's energy-band error scales as h⁴ for Yoshida 4 (ratios +// near 16 on halved steps) where Verlet's scales as h² (ratios near +// 4), measured on the same problem. +func TestYoshidaHarmonicFourthOrder(t *testing.T) { + for _, tc := range []struct { + name string + verlet bool + low float64 + high float64 + }{{"Verlet", true, 3, 5.5}, {"Yoshida4", false, 12, 20}} { + previous := 0.0 + for _, perPeriod := range []int{20, 40, 80} { + band := yoshidaEnergyBand(t, perPeriod, 4, tc.verlet) + t.Logf("%s steps/period %d: energy band %.3g", tc.name, perPeriod, band) + if previous > 0 { + if r := previous / band; r < tc.low || r > tc.high { + t.Fatalf("%s: error ratio %.2f at %d steps/period, want the band [%g, %g]", + tc.name, r, perPeriod, tc.low, tc.high) + } + } + previous = band + } + } +} + +// TestYoshidaKeplerEnergyBand pins the long-time fidelity: the +// Kepler two-body problem on an eccentric orbit keeps its energy in +// a narrow band over twenty periods, and the orbit closes. +func TestYoshidaKeplerEnergyBand(t *testing.T) { + const eccentricity = 0.5 + accel := func(q *core.Array) (*core.Array, error) { + x, y := q.FloatAt(0), q.FloatAt(1) + r3 := math.Pow(x*x+y*y, 1.5) + return core.FromFloats([]float64{-x / r3, -y / r3}, 2) + } + q0 := mustFloats(t, []float64{1 + eccentricity, 0}) + p0 := mustFloats(t, []float64{0, math.Sqrt((1 - eccentricity) / (1 + eccentricity))}) + periods := 20 + stepsPerPeriod := 200 + positions, momenta, err := IntegrateYoshida4(accel, 0, float64(periods)*2*math.Pi, q0, p0, periods*stepsPerPeriod) + if err != nil { + t.Fatalf("IntegrateYoshida4: %v", err) + } + energy := func(s int) float64 { + vx := momenta[s].FloatAt(0) + vy := momenta[s].FloatAt(1) + x := positions[s].FloatAt(0) + y := positions[s].FloatAt(1) + return 0.5*(vx*vx+vy*vy) - 1/math.Hypot(x, y) + } + e0 := energy(0) + worst := 0.0 + for s := range positions { + if d := math.Abs(energy(s) - e0); d > worst { + worst = d + } + } + t.Logf("Kepler e=0.5 over %d periods: energy band %.3g (E0 = %.6g)", periods, worst, e0) + if worst > 1e-4*math.Abs(e0) { + t.Fatalf("energy drifted by %.3g over %d periods", worst, periods) + } + // The closing error is the accumulated per-period phase error, + // fourth order in h, not an energy drift: at 200 steps per period + // it stays a few thousandths of the orbit radius. + last := len(positions) - 1 + if math.Hypot(positions[last].FloatAt(0)-(1+eccentricity), positions[last].FloatAt(1)) > 5e-3 { + t.Fatalf("the orbit did not close: q = (%.8g, %.8g)", + positions[last].FloatAt(0), positions[last].FloatAt(1)) + } +} + +func TestYoshidaErrors(t *testing.T) { + accel := func(q *core.Array) (*core.Array, error) { + return core.MulF(q, -1), nil + } + q0 := mustFloats(t, []float64{1}) + p0 := mustFloats(t, []float64{0}) + if _, _, err := IntegrateYoshida4(accel, 0, 1, q0, p0, 0); err == nil { + t.Fatal("zero steps: want an error") + } + if _, _, err := IntegrateYoshida4(accel, 0, 1, q0, mustFloats(t, []float64{0, 1}), 5); err == nil { + t.Fatal("mismatched vectors: want an error") + } + if _, _, err := IntegrateYoshida4(accel, 0, 1, mustFloats(t, []float64{}), mustFloats(t, []float64{}), 5); err == nil { + t.Fatal("empty state: want an error") + } + boom := func(*core.Array) (*core.Array, error) { return nil, base.Errf("accel failed") } + if _, _, err := IntegrateYoshida4(boom, 0, 1, q0, p0, 5); err == nil || !stringsContains(err, "accel failed") { + t.Fatal("a failing accel: want the error propagated") + } + wrong := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{1, 1}), nil } + if _, _, err := IntegrateYoshida4(wrong, 0, 1, q0, p0, 5); err == nil { + t.Fatal("wrong accel shape: want an error") + } + // A non-finite acceleration is refused instead of publishing NaNs. + nan := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{math.NaN()}), nil } + if _, _, err := IntegrateYoshida4(nan, 0, 1, q0, p0, 5); err == nil || !stringsContains(err, "non-finite") { + t.Fatal("a NaN accel: want an error") + } + // Int and complex states are refused, float32 states integrate. + pi0, err := core.FromInts([]int64{1}, 1) + if err != nil { + t.Fatal(err) + } + if _, _, err := IntegrateYoshida4(accel, 0, 1, pi0, p0, 5); err == nil { + t.Fatal("an int state: want an error") + } + q32, err := core.FromFloat32s([]float32{1}, 1) + if err != nil { + t.Fatal(err) + } + p32, err := core.FromFloat32s([]float32{0}, 1) + if err != nil { + t.Fatal(err) + } + if _, _, err := IntegrateYoshida4(accel, 0, 1, q32, p32, 5); err != nil { + t.Fatalf("a float32 state: %v", err) + } +} + +// midpointQP is the gradient of H(q, p) = q·p: the flow is +// q' = q, p' = −p, whose midpoint solution is the Cayley transform +// q_n = q0·((1+h/2)/(1−h/2))ⁿ, p_n = p0·((1−h/2)/(1+h/2))ⁿ, and +// H = q·p is conserved exactly. +func midpointQP(z *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{z.FloatAt(1), z.FloatAt(0)}, 2) +} + +// TestMidpointCayleyExact pins the analytic solution: for the +// nonseparable H = q·p the midpoint iterates must land on the Cayley +// transform and conserve H to rounding. +func TestMidpointCayleyExact(t *testing.T) { + const ( + q0 = 1.2 + p0 = 0.7 + steps = 50 + t1 = 1.0 + ) + h := t1 / steps + positions, momenta, err := IntegrateMidpoint(midpointQP, 0, t1, + mustFloats(t, []float64{q0}), mustFloats(t, []float64{p0}), steps, MidpointOptions{}) + if err != nil { + t.Fatalf("IntegrateMidpoint: %v", err) + } + up := (1 + h/2) / (1 - h/2) + down := (1 - h/2) / (1 + h/2) + for s := range positions { + q := positions[s].FloatAt(0) + p := momenta[s].FloatAt(0) + wantQ := q0 * math.Pow(up, float64(s)) + wantP := p0 * math.Pow(down, float64(s)) + if math.Abs(q-wantQ) > 1e-12 || math.Abs(p-wantP) > 1e-12 { + t.Fatalf("step %d: (q, p) = (%.14g, %.14g), want (%.14g, %.14g)", s, q, p, wantQ, wantP) + } + if math.Abs(q*p-q0*p0) > 1e-12 { + t.Fatalf("step %d: H = %.17g, want %.17g", s, q*p, q0*p0) + } + } +} + +// TestMidpointFixedPointExact pins the equilibrium property: an +// equilibrium of the flow is a fixed point of the midpoint map, +// exactly, with no roundoff drift. +func TestMidpointFixedPointExact(t *testing.T) { + positions, momenta, err := IntegrateMidpoint(midpointQP, 0, 0.5, + mustFloats(t, []float64{0}), mustFloats(t, []float64{0}), 3, MidpointOptions{}) + if err != nil { + t.Fatalf("IntegrateMidpoint: %v", err) + } + for s := range positions { + if positions[s].FloatAt(0) != 0 || momenta[s].FloatAt(0) != 0 { + t.Fatalf("step %d: the fixed point moved to (%g, %g)", + s, positions[s].FloatAt(0), momenta[s].FloatAt(0)) + } + } +} + +// TestMidpointNonseparableBand pins the energy band on a genuinely +// nonseparable two-degree system: H = ½(q₁²+1)(p₁²+1) + +// ½(q₂²+1)(p₂²+1) cannot be written as T(p) + V(q), and the midpoint +// rule must keep H inside a bounded band over the whole run. +func TestMidpointNonseparableBand(t *testing.T) { + gradH := func(z *core.Array) (*core.Array, error) { + q1, q2 := z.FloatAt(0), z.FloatAt(1) + p1, p2 := z.FloatAt(2), z.FloatAt(3) + return core.FromFloats([]float64{ + q1 * (p1*p1 + 1), q2 * (p2*p2 + 1), + p1 * (q1*q1 + 1), p2 * (q2*q2 + 1), + }, 4) + } + hamiltonian := func(q1, q2, p1, p2 float64) float64 { + return 0.5*(q1*q1+1)*(p1*p1+1) + 0.5*(q2*q2+1)*(p2*p2+1) + } + q0 := mustFloats(t, []float64{0.5, -0.4}) + p0 := mustFloats(t, []float64{0.3, 0.8}) + h0 := hamiltonian(0.5, -0.4, 0.3, 0.8) + positions, momenta, err := IntegrateMidpoint(gradH, 0, 10, q0, p0, 1000, MidpointOptions{}) + if err != nil { + t.Fatalf("IntegrateMidpoint: %v", err) + } + worst := 0.0 + for s := range positions { + h := hamiltonian(positions[s].FloatAt(0), positions[s].FloatAt(1), + momenta[s].FloatAt(0), momenta[s].FloatAt(1)) + if d := math.Abs(h - h0); d > worst { + worst = d + } + } + t.Logf("nonseparable two-degree band over t = 10 at h = 0.01: %.3g (H0 = %.6g)", worst, h0) + if worst > 1e-3*h0 { + t.Fatalf("energy drifted by %.3g (H0 = %.6g)", worst, h0) + } +} + +func TestMidpointErrors(t *testing.T) { + q0 := mustFloats(t, []float64{1}) + p0 := mustFloats(t, []float64{0}) + if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, q0, p0, 0, MidpointOptions{}); err == nil { + t.Fatal("zero steps: want an error") + } + if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, q0, mustFloats(t, []float64{0, 1}), 5, MidpointOptions{}); err == nil { + t.Fatal("mismatched vectors: want an error") + } + if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, mustFloats(t, []float64{}), mustFloats(t, []float64{}), 5, MidpointOptions{}); err == nil { + t.Fatal("empty state: want an error") + } + pi0, err := core.FromInts([]int64{1}, 1) + if err != nil { + t.Fatal(err) + } + if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, pi0, p0, 5, MidpointOptions{}); err == nil { + t.Fatal("an int state: want an error") + } + boom := func(*core.Array) (*core.Array, error) { return nil, base.Errf("gradH failed") } + if _, _, err := IntegrateMidpoint(boom, 0, 1, q0, p0, 5, MidpointOptions{}); err == nil || !stringsContains(err, "gradH failed") { + t.Fatal("a failing gradient: want the error propagated") + } + wrong := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{1}), nil } + if _, _, err := IntegrateMidpoint(wrong, 0, 1, q0, p0, 5, MidpointOptions{}); err == nil { + t.Fatal("wrong gradient shape: want an error") + } + nan := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{math.NaN(), 1}), nil } + if _, _, err := IntegrateMidpoint(nan, 0, 1, q0, p0, 5, MidpointOptions{}); err == nil || !stringsContains(err, "non-finite") { + t.Fatal("a NaN gradient: want an error") + } +} + +// TestVerletFamilyComplexRefusal pins the dtype contract across the +// whole family: complex states are refused everywhere. +func TestVerletFamilyComplexRefusal(t *testing.T) { + accel := func(q *core.Array) (*core.Array, error) { return core.MulF(q, -1), nil } + gradH := func(z *core.Array) (*core.Array, error) { return core.FromFloats([]float64{0, 0}, 2) } + qi, err := core.FromFloat32s([]float32{1}, 1) + if err != nil { + t.Fatal(err) + } + pi, err := core.FromFloat32s([]float32{0}, 1) + if err != nil { + t.Fatal(err) + } + if _, _, err := IntegrateYoshida4(accel, 0, 1, qi, pi, 5); err != nil { + t.Fatalf("a float32 Yoshida state: %v", err) + } + if _, _, err := IntegrateMidpoint(gradH, 0, 1, qi, pi, 5, MidpointOptions{}); err != nil { + t.Fatalf("a float32 midpoint state: %v", err) + } + // A complex state is refused on both integrators. + complexArr := core.New(core.Complex, 1) + zeros := mustFloats(t, []float64{0}) + if _, _, err := IntegrateYoshida4(accel, 0, 1, complexArr, zeros, 5); err == nil || !stringsContains(err, "complex") { + t.Fatal("a complex Yoshida state: want an error") + } + if _, _, err := IntegrateMidpoint(gradH, 0, 1, complexArr, zeros, 5, MidpointOptions{}); err == nil || !stringsContains(err, "complex") { + t.Fatal("a complex midpoint state: want an error") + } +} + +// TestVerletFamilyNonFiniteRefusal pins the up-front refusal of +// non-finite initial states on both new integrators. +func TestVerletFamilyNonFiniteRefusal(t *testing.T) { + accel := func(q *core.Array) (*core.Array, error) { return core.MulF(q, -1), nil } + gradH := func(z *core.Array) (*core.Array, error) { return core.FromFloats([]float64{0, 0}, 2) } + nanQ := mustFloats(t, []float64{math.NaN()}) + p0 := mustFloats(t, []float64{0}) + if _, _, err := IntegrateYoshida4(accel, 0, 1, nanQ, p0, 5); err == nil || !stringsContains(err, "non-finite") { + t.Fatal("a NaN position: want an error") + } + if _, _, err := IntegrateMidpoint(gradH, 0, 1, nanQ, p0, 5, MidpointOptions{}); err == nil || !stringsContains(err, "non-finite") { + t.Fatal("a NaN position: want an error") + } + q0 := mustFloats(t, []float64{1}) + nanP := mustFloats(t, []float64{math.Inf(-1)}) + if _, _, err := IntegrateYoshida4(accel, 0, 1, q0, nanP, 5); err == nil || !stringsContains(err, "non-finite") { + t.Fatal("an infinite momentum: want an error") + } + if _, _, err := IntegrateMidpoint(gradH, 0, 1, q0, nanP, 5, MidpointOptions{}); err == nil || !stringsContains(err, "non-finite") { + t.Fatal("an infinite momentum: want an error") + } +} diff --git a/integrate/symplectic_test.go b/integrate/symplectic_test.go new file mode 100644 index 0000000..c66bee2 --- /dev/null +++ b/integrate/symplectic_test.go @@ -0,0 +1,178 @@ +// Copyright (c) 2026 Petr Balvín (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" +) + +// TestIntegrateVerletHarmonicEnergy runs the harmonic oscillator over +// fifty periods: the true energy must stay inside a bounded band the +// whole time (the symplectic property) and the orbit must close on +// itself to second-order accuracy. +func TestIntegrateVerletHarmonicEnergy(t *testing.T) { + const omega = 1.0 + accel := func(q *core.Array) (*core.Array, error) { + return core.MulF(q, -omega*omega), nil + } + const periods = 50 + // With h per period the energy band sits at O((ωh)²/8)·E; 300 + // steps per period put it under the 1e-4 bar. + stepsPerPeriod := 400 + q0 := mustFloats(t, []float64{1}) + p0 := mustFloats(t, []float64{0}) + positions, momenta, err := IntegrateVerlet(accel, 0, periods*2*math.Pi, q0, p0, periods*stepsPerPeriod) + if err != nil { + t.Fatalf("IntegrateVerlet: %v", err) + } + energy := func(q, p float64) float64 { + return 0.5 * (p*p + omega*omega*q*q) + } + e0 := energy(positions[0].FloatAt(0), momenta[0].FloatAt(0)) + worst := 0.0 + for s := range positions { + e := energy(positions[s].FloatAt(0), momenta[s].FloatAt(0)) + if d := math.Abs(e - e0); d > worst { + worst = d + } + } + // The band must stay narrow relative to the energy itself, over + // fifty periods, where a non-symplectic scheme drifts away. + if worst > 1e-4*e0 { + t.Fatalf("energy drifted by %g (energy %g) over %d periods", worst, e0, periods) + } + last := positions[len(positions)-1].FloatAt(0) + if math.Abs(last-1) > 1e-3 { + t.Fatalf("after a whole number of periods q = %.10g, want ≈ 1", last) + } + // Second order: halving the step must cut the closing error by + // roughly four. + q1, _, err := IntegrateVerlet(accel, 0, 2*math.Pi, q0, p0, 20) + if err != nil { + t.Fatalf("IntegrateVerlet coarse: %v", err) + } + q2, _, err := IntegrateVerlet(accel, 0, 2*math.Pi, q0, p0, 40) + if err != nil { + t.Fatalf("IntegrateVerlet fine: %v", err) + } + eCoarse := math.Abs(q1[len(q1)-1].FloatAt(0) - 1) + eFine := math.Abs(q2[len(q2)-1].FloatAt(0) - 1) + if eFine > 0.35*eCoarse { + t.Fatalf("convergence order looks wrong: coarse %g, fine %g", eCoarse, eFine) + } +} + +// TestIntegrateVerletPendulum checks a nonlinear system: the pendulum +// with E = p²/2 − cos q must keep its energy bounded, including near +// the separatrix where the force is far from linear. +func TestIntegrateVerletPendulum(t *testing.T) { + accel := func(q *core.Array) (*core.Array, error) { + return mustFloats(t, []float64{-math.Sin(q.FloatAt(0))}), nil + } + q0 := mustFloats(t, []float64{2.5}) + p0 := mustFloats(t, []float64{0}) + positions, momenta, err := IntegrateVerlet(accel, 0, 200, q0, p0, 20000) + if err != nil { + t.Fatalf("IntegrateVerlet: %v", err) + } + energyOf := func(q, p float64) float64 { + return 0.5*p*p - math.Cos(q) + } + e0 := energyOf(positions[0].FloatAt(0), momenta[0].FloatAt(0)) + worst := 0.0 + for s := range positions { + e := energyOf(positions[s].FloatAt(0), momenta[s].FloatAt(0)) + if d := math.Abs(e - e0); d > worst { + worst = d + } + } + if worst > 1e-3*math.Abs(e0) { + t.Fatalf("pendulum energy drifted by %g (energy %g)", worst, e0) + } +} + +func TestIntegrateVerletErrors(t *testing.T) { + accel := func(q *core.Array) (*core.Array, error) { + return core.MulF(q, -1), nil + } + q0 := mustFloats(t, []float64{1}) + p0 := mustFloats(t, []float64{0}) + if _, _, err := IntegrateVerlet(accel, 0, 1, q0, p0, 0); err == nil { + t.Fatal("zero steps: want an error") + } + if _, _, err := IntegrateVerlet(accel, 0, 1, q0, mustFloats(t, []float64{0, 1}), 5); err == nil { + t.Fatal("mismatched vectors: want an error") + } + if _, _, err := IntegrateVerlet(accel, 0, 1, mustFloats(t, []float64{}), mustFloats(t, []float64{}), 5); err == nil { + t.Fatal("empty state: want an error") + } + boom := func(*core.Array) (*core.Array, error) { return nil, base.Errf("accel failed") } + if _, _, err := IntegrateVerlet(boom, 0, 1, q0, p0, 5); err == nil { + t.Fatal("accel error: want an error") + } + wrong := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{1, 1}), nil } + if _, _, err := IntegrateVerlet(wrong, 0, 1, q0, p0, 5); err == nil { + t.Fatal("wrong accel shape: want an error") + } +} + +// TestIntegrateVerletFloat32State pins the dtype contract: a float32 +// state integrates (read per element, widened exactly) and lands where +// the float64 run lands; RawFloats is nil for float32, so the old +// code silently integrated zeros. +func TestIntegrateVerletFloat32State(t *testing.T) { + accel := func(q *core.Array) (*core.Array, error) { + return core.MulF(q, -1), nil + } + q0, err := core.FromFloat32s([]float32{1}, 1) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + p0, err := core.FromFloat32s([]float32{0}, 1) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + positions, momenta, err := IntegrateVerlet(accel, 0, 2*math.Pi, q0, p0, 400) + if err != nil { + t.Fatalf("IntegrateVerlet: %v", err) + } + if got := positions[400].FloatAt(0); math.Abs(got-1) > 1e-3 { + t.Fatalf("float32 state closed at %.6g, want ≈ 1", got) + } + if got := momenta[0].FloatAt(0); got != 0 { + t.Fatalf("initial momentum = %g, want 0", got) + } +} + +// TestIntegrateVerletIntStateErrors pins the refusal of int states: +// they have no place in a continuous integrator and must error rather +// than panic or read as zeros. +func TestIntegrateVerletIntStateErrors(t *testing.T) { + accel := func(q *core.Array) (*core.Array, error) { + return core.MulF(q, -1), nil + } + qi, err := core.FromInts([]int64{1}, 1) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + pi, err := core.FromInts([]int64{0}, 1) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + if _, _, err := IntegrateVerlet(accel, 0, 1, qi, pi, 5); err == nil { + t.Fatal("expected an error for an int state") + } + if _, _, err := IntegrateVerlet(accel, 0, 1, mustFloats(t, []float64{1}), pi, 5); err == nil { + t.Fatal("expected an error for an int momentum with a float position") + } + if _, _, err := IntegrateVerlet(accel, 0, 1, qi, mustFloats(t, []float64{0}), 5); err == nil { + t.Fatal("expected an error for an int position with a float momentum") + } +} diff --git a/integrate/tolerance_pins_test.go b/integrate/tolerance_pins_test.go new file mode 100644 index 0000000..b9466cd --- /dev/null +++ b/integrate/tolerance_pins_test.go @@ -0,0 +1,196 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression tests for the integrators: the +// absolute cubature tolerance, the collapse floor of the ODE steppers +// and the sample times of the PDE evolutions. + +// TestCubatureScalesWithMagnitude pins the stopping rule: a smooth +// integrand of large magnitude must converge, not exhaust the budget +// chasing an absolute bound below the rounding floor of the sum. +func TestCubatureScalesWithMagnitude(t *testing.T) { + const want = 1e6 // ∫∫ 1e6 over the unit square + got, err := IntegrateND(func([]float64) float64 { return want }, []float64{0, 0}, []float64{1, 1}, CubatureOptions{}) + if err != nil { + t.Fatalf("IntegrateND of a constant: %v", err) + } + if math.Abs(got-want) > 1e-6*want { + t.Fatalf("IntegrateND = %v, want %v", got, want) + } + // A small integral keeps its absolute accuracy. + small, err := IntegrateND(func([]float64) float64 { return 1e-8 }, []float64{0, 0}, []float64{1, 1}, CubatureOptions{}) + if err != nil { + t.Fatalf("IntegrateND of a small constant: %v", err) + } + if math.Abs(small-1e-8) > 1e-12 { + t.Fatalf("IntegrateND = %v, want 1e-8", small) + } +} + +// TestODETinySpan pins the collapse rule: a span far below the absolute +// time scale is integrable, and the stepper must not refuse it. +func TestODETinySpan(t *testing.T) { + zero := func(float64, *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{0}, 1) + } + y0, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + for _, span := range []float64{1e-13, 1e-15, 1e-20} { + end, err := IntegrateODE(zero, 0, span, y0, ODEOptions{MaxSteps: 100000}) + if err != nil { + t.Fatalf("span %g: %v", span, err) + } + if got := end.FloatAt(0); got != 1 { + t.Fatalf("span %g: y = %v, want 1", span, got) + } + } + // A real decay over a tiny span: y' = −y, y(1e-12) = exp(−1e-12). + const span = 1e-12 + decay := func(_ float64, y *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-y.FloatAt(0)}, 1) + } + end, err := IntegrateODE(decay, 0, span, y0, ODEOptions{MaxSteps: 100000}) + if err != nil { + t.Fatalf("decay over %g: %v", span, err) + } + if got, want := end.FloatAt(0), math.Exp(-span); math.Abs(got-want) > 1e-12 { + t.Fatalf("y = %.17g, want %.17g", got, want) + } +} + +// sineMode returns the interior grid of sin(π·x) on [0, 1] with n +// points, the eigenmode of the Dirichlet Laplacian. +func sineMode(t *testing.T, n int) *core.Array { + t.Helper() + u := make([]float64, n) + for i := range n { + x := float64(i+1) / float64(n+1) + u[i] = math.Sin(math.Pi * x) + } + a, err := core.FromFloats(u, n) + if err != nil { + t.Fatal(err) + } + return a +} + +// TestPDESchedule pins the step schedule directly, because the +// physics tests below pass with either schedule for a fine enough +// step: the published times must be j·tFinal/(samples−1) exactly, so +// the step count is a multiple of samples−1 and the last step lands on +// tFinal. `dt` is an upper bound, never a divisor to be honoured +// blindly. +func TestPDESchedule(t *testing.T) { + cases := []struct { + tFinal, dt float64 + samples int + }{ + {1, 0.3, 3}, {1, 0.3, 5}, {1, 0.7, 4}, {2.5, 0.4, 3}, {0.1, 0.3, 2}, {1, 1, 3}, {1, 0.01, 8}, + } + for _, tc := range cases { + steps, h := pdeSchedule(tc.tFinal, tc.dt, tc.samples) + if steps <= 0 || h <= 0 { + t.Fatalf("pdeSchedule(%v, %v, %d) = %d steps of %v", tc.tFinal, tc.dt, tc.samples, steps, h) + } + if float64(steps)*h != tc.tFinal { + t.Errorf("pdeSchedule(%v, %v, %d): %d steps of %v reach %v", + tc.tFinal, tc.dt, tc.samples, steps, h, float64(steps)*h) + } + if steps%(tc.samples-1) != 0 { + t.Errorf("pdeSchedule(%v, %v, %d): %d steps do not divide by %d", + tc.tFinal, tc.dt, tc.samples, steps, tc.samples-1) + } + if h > tc.dt { + t.Errorf("pdeSchedule(%v, %v, %d): the step %v exceeds the bound %v", + tc.tFinal, tc.dt, tc.samples, h, tc.dt) + } + every := steps / (tc.samples - 1) + for j := range tc.samples { + got := float64(j*every) * h + want := tc.tFinal * float64(j) / float64(tc.samples-1) + if math.Abs(got-want) > 1e-12*tc.tFinal { + t.Errorf("sample %d at %v, want %v", j, got, want) + } + } + } +} + +// TestHeatSamplesLandOnTheirTimes checks the values at the published +// times; the schedule itself is pinned above. +// published states are the states at t = 0, tFinal/2 and tFinal, which +// the single sine mode turns into an exact decay ratio. The step is +// chosen for accuracy (r = κ·h/dx² ≈ 0.45) and does not divide tFinal, +// so the schedule has to round the count up. +func TestHeatSamplesLandOnTheirTimes(t *testing.T) { + const ( + kappa = 1.0 + n = 128 + ) + u0 := sineMode(t, n) + dx := 1.0 / float64(n+1) + dt := 0.9 * 0.5 * dx * dx + got, err := IntegrateHeat1D(u0, kappa, dx, 1, dt, 3, 0, 0) + if err != nil { + t.Fatalf("IntegrateHeat1D: %v", err) + } + if s := got.Shape(); s[0] != 3 || s[1] != n { + t.Fatalf("shape %v, want [3 %d]", s, n) + } + decay := func(tm float64) float64 { return math.Exp(-kappa * math.Pi * math.Pi * tm) } + for j := range n { + first := got.FloatAt(j) + if math.Abs(first-u0.FloatAt(j)) > 1e-12 { + t.Fatalf("sample 0 is not the initial state at %d", j) + } + mid := got.FloatAt(n + j) + wantMid := first * decay(0.5) + if rel := math.Abs(mid-wantMid) / wantMid; rel > 1e-3 { + t.Fatalf("sample 1 at %d: %v, want %v (relative %.2g): the time is not 0.5", j, mid, wantMid, rel) + } + last := got.FloatAt(2*n + j) + wantLast := first * decay(1.0) + if rel := math.Abs(last-wantLast) / wantLast; rel > 1e-3 { + t.Fatalf("sample 2 at %d: %v, want %v (relative %.2g): the time is not 1", j, last, wantLast, rel) + } + } +} + +// TestWaveSamplesLandOnTheirTimes does the same for the wave equation +// (the schedule is pinned above): +// the standing mode is cos(π·c·t)·sin(π·x), so the sample at tFinal = 1 +// with c = 1 is the initial state negated and the one at 0.5 is zero. +func TestWaveSamplesLandOnTheirTimes(t *testing.T) { + const n = 128 + u0 := sineMode(t, n) + v0, err := core.FromFloats(make([]float64, n), n) + if err != nil { + t.Fatal(err) + } + dx := 1.0 / float64(n+1) + got, err := IntegrateWave1D(u0, v0, 1, dx, 1, 0.9*dx, 3) + if err != nil { + t.Fatalf("IntegrateWave1D: %v", err) + } + for j := range n { + last := got.FloatAt(2*n + j) + want := -u0.FloatAt(j) // cos(π·1) = −1 + if math.Abs(last-want) > 5e-3*math.Abs(want) { + t.Fatalf("sample 2 at %d: %v, want %v: the last sample is not at t = 1", j, last, want) + } + mid := got.FloatAt(n + j) + if math.Abs(mid) > 1e-2*math.Abs(u0.FloatAt(j)) { + t.Fatalf("sample 1 at %d: %v, want near zero at t = 0.5", j, mid) + } + } +} diff --git a/integrate/tridiag.go b/integrate/tridiag.go new file mode 100644 index 0000000..40526a0 --- /dev/null +++ b/integrate/tridiag.go @@ -0,0 +1,39 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +// The tridiagonal scratch the PDE step sweeps reuse. The Crank- +// Nicolson and Peaceman-Rachford schemes solve one tridiagonal system +// per line per step against matrices that are constants of the scheme, +// so the elimination scratch belongs to the solve, not to the line: +// the sweeps below allocate it once and reuse it across every line and +// every step, where the library's array-level SolveTridiagonal +// allocated its working vectors per call. The elimination itself is +// internal/base's TriSolve, shared with the public solver, so the +// arithmetic exists once. + +// triScratch holds one tridiagonal solve's reusable buffers: the +// right-hand side, the eliminated superdiagonal and right side of the +// Thomas algorithm, and the solution buffer for callers whose output +// is not written straight into their state. aux is a second right-side +// lane for the sweeps that build two neighbouring systems in one pass. +// Every buffer is fully overwritten before the kernel reads it, except +// cp, whose prefix is written and read in the same sweep order the +// fresh buffers saw. +type triScratch struct { + rhs, cp, dp, dst, aux []float64 +} + +// triSized sizes s for a system of sys unknowns whose right side is a +// line of lineLen entries, reusing storage that is already large +// enough. The solution buffer dst is sized for callers that scatter +// the result elsewhere; a caller writing straight into its state leaves +// it unread. +func triSized(s *triScratch, sys, lineLen int) { + s.rhs = sizedBuf(s.rhs, lineLen) + s.aux = sizedBuf(s.aux, lineLen) + s.cp = sizedBuf(s.cp, sys) + s.dp = sizedBuf(s.dp, sys) + s.dst = sizedBuf(s.dst, sys) +} diff --git a/integrate/wave2d_pins_test.go b/integrate/wave2d_pins_test.go new file mode 100644 index 0000000..cfd6ff2 --- /dev/null +++ b/integrate/wave2d_pins_test.go @@ -0,0 +1,391 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Regression pins for integrate/pde2d.go: the three-buffer leapfrog +// rotation, the pristine read of the Taylor start, the published-sample +// floor and schedule, and the complex-velocity guard. The oracles here +// are derived from the stencil itself, not from the library's own +// paths. + +package integrate + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// pde2dSteps re-derives the documented schedule contract independently +// of the implementation's pdeSchedule: ceil(tFinal/dt) steps, rounded up +// to a multiple of the sampling interval samples-1, all of size +// tFinal/steps, so every published time j·tFinal/(samples−1) is a step +// boundary. The returned every = steps/(samples−1) steps sit between two +// published samples. +func pde2dSteps(tFinal, dt float64, samples int) (steps, every int, h float64) { + steps = max(int(math.Ceil(tFinal/dt)), samples-1) + if rem := steps % (samples - 1); rem != 0 { + steps += samples - 1 - rem + } + return steps, steps / (samples - 1), tFinal / float64(steps) +} + +// pinSineMode builds sin(kx·π·x)·sin(ky·π·y) on a rows×cols grid whose +// ring is zero, sampled with x = c·dx and y = r·dy. +func pinSineMode(t *testing.T, rows, cols, kx, ky int) *core.Array { + t.Helper() + vals := make([]float64, rows*cols) + for r := range rows { + for c := range cols { + vals[r*cols+c] = math.Sin(float64(kx)*math.Pi*float64(c)/float64(cols-1)) * + math.Sin(float64(ky)*math.Pi*float64(r)/float64(rows-1)) + } + } + a, err := core.FromFloats(vals, rows, cols) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// modeMu is the eigenvalue the five-point Laplacian gives that +// mode: Δmode = −μ·mode with μ = 4/dx²·sin²(kxπ/(2(cols−1))) + +// 4/dy²·sin²(kyπ/(2(rows−1))). +func modeMu(rows, cols, kx, ky int, dx, dy float64) float64 { + sx := math.Sin(float64(kx) * math.Pi / (2 * float64(cols-1))) + sy := math.Sin(float64(ky) * math.Pi / (2 * float64(rows-1))) + return 4/(dx*dx)*sx*sx + 4/(dy*dy)*sy*sy +} + +// point3x3 is the exact trajectory of the single interior point of +// a 3x3 grid with the zero ring: its four neighbours are all boundary +// values, so the five-point Laplacian is exactly −μ·u with +// μ = 2/dx² + 2/dy² and the leapfrog collapses to the scalar recurrence +// u_{s+1} = 2u_s − u_{s−1} + (c·h)²·(−μ·u_s), started by the Taylor step +// u¹ = u⁰ + h·v⁰ + (c·h)²/2·(−μ·u⁰). Its closed form is cos(s·θ) with +// cos θ = 1 − (c·h)²·μ/2. +func point3x3(steps int, h, c, dx, dy, u0, v0 float64) (hist []float64, theta float64) { + mu := 2/(dx*dx) + 2/(dy*dy) + lapW := (c * h) * (c * h) + theta = math.Acos(1 - 0.5*lapW*mu) + hist = make([]float64, steps+1) + hist[0] = u0 + hist[1] = u0 + h*v0 + 0.5*lapW*(-mu*u0) + for s := 2; s <= steps; s++ { + hist[s] = 2*hist[s-1] - hist[s-2] + lapW*(-mu*hist[s-1]) + } + return hist, theta +} + +// singlePoint3x3 builds the 3x3 initial state whose one interior +// point carries u = 1 and whose ring is zero. +func singlePoint3x3(t *testing.T) *core.Array { + t.Helper() + vals := make([]float64, 9) + vals[4] = 1 + a, err := core.FromFloats(vals, 3, 3) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestWave2DLeapfrogScalarTrajectory pins the three-buffer rotation. On +// the 3x3 grid the interior point has no interior neighbour, so the +// whole trajectory IntegrateWave2D returns must equal the scalar +// leapfrog step by step. The old rotation (prev, cur = cur, next) left +// next aliased to cur after two swaps, which replaces the second-order +// recurrence with the first-order map u next = u + lapW*lap(u): step 201 came +// back as +0.7253864115 where the leapfrog gives −0.1460225271. +func TestWave2DLeapfrogScalarTrajectory(t *testing.T) { + const dx, dy, c = 0.5, 0.5, 1.0 + const tFinal, dt, samples = 2.0, 0.01, 202 + steps, every, h := pde2dSteps(tFinal, dt, samples) + if steps != 201 || every != 1 { + t.Fatalf("schedule %d steps of %v, every %d, want 201 steps every 1", steps, h, every) + } + u0 := singlePoint3x3(t) + v0 := core.New(core.Float, 3, 3) + hist, err := IntegrateWave2D(u0, v0, c, dx, dy, tFinal, dt, samples) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + ref, theta := point3x3(steps, h, c, dx, dy, 1, 0) + worst, worstAt := 0.0, 0 + for j := range samples { + got := hist.FloatAt(j*9 + 4) + if e := math.Abs(got - ref[j*every]); e > worst { + worst, worstAt = e, j + } + } + if worst > 1e-12 { + t.Fatalf("sample %d = %.12f, want the scalar leapfrog %.12f (worst deviation %.3e over the whole history)", + worstAt, hist.FloatAt(worstAt*9+4), ref[worstAt*every], worst) + } + // The recurrence itself is the closed form cos(s·θ), so the oracle is + // pinned to the discrete characteristic and not to the implementation. + worst, worstAt = 0.0, 0 + for s := range steps + 1 { + want := math.Cos(float64(s) * theta) + if e := math.Abs(ref[s] - want); e > worst { + worst, worstAt = e, s + } + } + if worst > 1e-11 { + t.Fatalf("the reference recurrence deviates from cos(s·θ) by %.3e at step %d", worst, worstAt) + } +} + +// TestWave2DTaylorStartReadsUntouchedState pins the start-up read. The +// first step is the Taylor start, and it must read the untouched u⁰: +// computed into the same slice it reads (the old form wrote cur in place), +// the Laplacian at an interior point picks up already-updated left and +// upper neighbours. On this exact eigenmode with v⁰ = 0 the first +// published sample is cos θ·u⁰ to rounding, and the in-place start moved +// it by 1.4e-7 (relative to the unit amplitude at the centre). +func TestWave2DTaylorStartReadsUntouchedState(t *testing.T) { + const n = 9 + const c = 1.0 + dx := 1.0 / float64(n-1) + const tFinal, dt, samples = 0.008, 0.004, 3 + steps, every, h := pde2dSteps(tFinal, dt, samples) + if steps != 2 || every != 1 { + t.Fatalf("schedule %d steps of %v, every %d, want 2 steps every 1", steps, h, every) + } + u0 := pinSineMode(t, n, n, 1, 1) + v0 := core.New(core.Float, n, n) + hist, err := IntegrateWave2D(u0, v0, c, dx, dx, tFinal, dt, samples) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + mu := modeMu(n, n, 1, 1, dx, dx) + cosTheta := 1 - 0.5*(c*h)*(c*h)*mu + worst, worstAt := 0.0, 0 + for r := range n { + for cc := range n { + i := r*n + cc + want := cosTheta * u0.FloatAt(i) + if e := math.Abs(hist.FloatAt(n*n+i) - want); e > worst { + worst, worstAt = e, i + } + } + } + if worst > 1e-13 { + t.Fatalf("the Taylor start deviates from cos θ·u⁰ by %.3e at point %d (row %d, column %d): %.12f, want %.12f", + worst, worstAt, worstAt/n, worstAt%n, hist.FloatAt(n*n+worstAt), cosTheta*u0.FloatAt(worstAt)) + } +} + +// TestWave2DSamplesFloorTwo pins the floors of the published-sample +// argument. One sample leaves steps/(samples−1) with a zero divisor, so +// both 2-D entry points must refuse it the way the 1-D solvers do +// ("at least two samples are needed"), not panic with an integer divide +// by zero as the earlier revision did. +func TestWave2DSamplesFloorTwo(t *testing.T) { + u0 := singlePoint3x3(t) + v0 := core.New(core.Float, 3, 3) + for _, samples := range []int{1, 0} { + _, err := IntegrateHeat2D(u0, 1, 0.25, 0.25, 0.1, 0.01, samples, 0, 0, 0, 0) + if err == nil { + t.Fatalf("IntegrateHeat2D accepted samples = %d", samples) + } + if !strings.Contains(err.Error(), "at least two samples are needed") { + t.Fatalf("IntegrateHeat2D samples = %d: error %q, want the at-least-two wording", samples, err) + } + _, err = IntegrateWave2D(u0, v0, 1, 0.25, 0.25, 0.1, 0.01, samples) + if err == nil { + t.Fatalf("IntegrateWave2D accepted samples = %d", samples) + } + if !strings.Contains(err.Error(), "at least two samples are needed") { + t.Fatalf("IntegrateWave2D samples = %d: error %q, want the at-least-two wording", samples, err) + } + } + // Two samples are still a valid call: the endpoints alone. + if _, err := IntegrateWave2D(u0, v0, 1, 0.25, 0.25, 0.1, 0.01, 2); err != nil { + t.Fatalf("IntegrateWave2D with samples = 2: %v", err) + } + if _, err := IntegrateHeat2D(u0, 1, 0.25, 0.25, 0.1, 0.01, 2, 0, 0, 0, 0); err != nil { + t.Fatalf("IntegrateHeat2D with samples = 2: %v", err) + } +} + +// TestWave2DComplexVelocityRefusal pins the dtype guard on the initial +// velocity. The earlier revision checked only the rank and shape of v0, +// so a complex array reached v0.FloatAt and panicked inside the Taylor +// start ("index out of range [6] with length 0"); the 1-D solver refuses +// the same input with a message, which is the behaviour mirrored here. +func TestWave2DComplexVelocityRefusal(t *testing.T) { + u0, err := core.FromFloats(make([]float64, 25), 5, 5) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + v0 := core.New(core.Complex, 5, 5) + _, err = IntegrateWave2D(u0, v0, 1, 0.25, 0.25, 0.1, 0.01, 3) + if err == nil { + t.Fatal("IntegrateWave2D accepted a complex velocity array") + } + if !strings.Contains(err.Error(), "complex velocities are not supported") { + t.Fatalf("IntegrateWave2D: error %q, want the unsupported-velocity wording", err) + } + // The same wording the 1-D wave solver uses for the same input. + u1 := core.New(core.Float, 5) + v1 := core.New(core.Complex, 5) + _, err = IntegrateWave1D(u1, v1, 1, 0.25, 0.1, 0.01, 3) + if err == nil || !strings.Contains(err.Error(), "complex velocities are not supported") { + t.Fatalf("IntegrateWave1D: error %v, want the same refusal", err) + } +} + +// TestWave2DSlotsLandOnTheirTimes pins the published-sample schedule with +// an interval wider than one step. With samples = 5 over 8 steps the +// interval is every = 2, so slot j must hold the state after 2j steps, at +// the time 2j·h = j·tFinal/(samples−1). The earlier revision published +// the Taylor start (step 1) in slot 1 instead of step 2, shifting every +// slot from 1 to samples−2 off its published time. +func TestWave2DSlotsLandOnTheirTimes(t *testing.T) { + const dx, dy, c = 0.5, 0.5, 1.0 + const tFinal, dt, samples = 0.08, 0.01, 5 + steps, every, h := pde2dSteps(tFinal, dt, samples) + if steps != 8 || every != 2 { + t.Fatalf("schedule %d steps of %v, every %d, want 8 steps every 2", steps, h, every) + } + u0 := singlePoint3x3(t) + v0 := core.New(core.Float, 3, 3) + hist, err := IntegrateWave2D(u0, v0, c, dx, dy, tFinal, dt, samples) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + ref, _ := point3x3(steps, h, c, dx, dy, 1, 0) + for j := range samples { + // The published time must be the step boundary j·every, which is + // what makes the slots a uniform time grid with both endpoints. + time, boundary := float64(j)*tFinal/float64(samples-1), float64(j*every)*h + if e := math.Abs(time - boundary); e > 1e-15 { + t.Fatalf("slot %d: published time %v is %d steps (%v) into the run, off by %.3e", + j, time, j*every, boundary, e) + } + got, want := hist.FloatAt(j*9+4), ref[j*every] + if e := math.Abs(got - want); e > 1e-12 { + t.Fatalf("slot %d (t = %v) = %.12f, want the state after %d steps %.12f (off by %.3e)", + j, time, got, j*every, want, e) + } + } +} + +// TestWave2DRingHeldAtZero pins the documented ring contract against the +// buffer recycling of the rotation fix: the solver holds the boundary +// ring at zero, so a caller's own ring values may not reach the stencil +// at any published sample. A run whose state is all ones (ring 1) must +// therefore agree, from slot 1 on, with the run whose ring is zero, and +// the returned ring must be zero there. The earlier revision read the +// caller's ring into the first two Laplacians, and the three-buffer +// rotation makes that slip possible every third step unless the recycled +// buffer is cleared. +func TestWave2DRingHeldAtZero(t *testing.T) { + const n = 5 + const c = 1.0 + dx := 1.0 / float64(n-1) + const tFinal, dt, samples = 0.04, 0.01, 5 + steps, every, _ := pde2dSteps(tFinal, dt, samples) + if steps != 4 || every != 1 { + t.Fatalf("schedule %d steps, every %d, want 4 steps every 1", steps, every) + } + ones := make([]float64, n*n) + for i := range ones { + ones[i] = 1 + } + cleared := make([]float64, n*n) + copy(cleared, ones) + for r := range n { + cleared[r*n], cleared[r*n+n-1] = 0, 0 + } + for cc := range n { + cleared[cc], cleared[(n-1)*n+cc] = 0, 0 + } + ringed, err := core.FromFloats(ones, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + bare, err := core.FromFloats(cleared, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + v0 := core.New(core.Float, n, n) + got, err := IntegrateWave2D(ringed, v0, c, dx, dx, tFinal, dt, samples) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + want, err := IntegrateWave2D(bare, v0, c, dx, dx, tFinal, dt, samples) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + for j := 1; j < samples; j++ { + for i := range n * n { + r, cc := i/n, i%n + g := got.FloatAt(j*n*n + i) + if e := math.Abs(g - want.FloatAt(j*n*n+i)); e > 1e-15 { + t.Fatalf("slot %d point (row %d, column %d) = %.16f, want %.16f: the caller's ring reached the stencil (off by %.3e)", + j, r, cc, g, want.FloatAt(j*n*n+i), e) + } + if (r == 0 || r == n-1 || cc == 0 || cc == n-1) && g != 0 { + t.Fatalf("slot %d point (row %d, column %d) = %v, want the ring held at zero", j, r, cc, g) + } + } + } +} + +// TestWave2DStandingModeMatchesDiscreteEigenvalue is the acceptance test +// for the leapfrog. The (1,1) sine mode of the 9x9 grid is an exact +// eigenfunction of the five-point Laplacian, so with zero initial +// velocity every published sample must be cos(ω·t_j)·u⁰, where the +// discrete characteristic of the scheme is +// cos(ω·h) = 1 − (c·h)²·μ/2 and μ is the eigenvalue derived above. The +// earlier revision returned +0.9250145127 at t = 1 where the discrete +// solution is −0.2935539474 (monotone decay instead of oscillation), and +// its in-place Taylor start missed the first sample by 1.4e-7. +func TestWave2DStandingModeMatchesDiscreteEigenvalue(t *testing.T) { + const n = 9 + const c = 1.0 + dx := 1.0 / float64(n-1) + const tFinal, dt, samples = 1.0, 0.004, 253 + steps, every, h := pde2dSteps(tFinal, dt, samples) + if steps != 252 || every != 1 { + t.Fatalf("schedule %d steps of %v, every %d, want 252 steps every 1", steps, h, every) + } + u0 := pinSineMode(t, n, n, 1, 1) + v0 := core.New(core.Float, n, n) + hist, err := IntegrateWave2D(u0, v0, c, dx, dx, tFinal, dt, samples) + if err != nil { + t.Fatalf("IntegrateWave2D: %v", err) + } + mu := modeMu(n, n, 1, 1, dx, dx) + theta := math.Acos(1 - 0.5*(c*h)*(c*h)*mu) // ω = θ/h, and t_j = j·every·h + worst, worstAt, worstSlot := 0.0, 0, 0 + for j := range samples { + amp := math.Cos(float64(j*every) * theta) + for i := range n * n { + got := hist.FloatAt(j*n*n + i) + if e := math.Abs(got - amp*u0.FloatAt(i)); e > worst { + worst, worstAt, worstSlot = e, i, j + } + } + } + if worst > 1e-11 { + t.Fatalf("sample %d (t = %v) point %d = %.12f, want cos(ω·t)·u⁰ = %.12f (worst %.3e over every published sample)", + worstSlot, float64(worstSlot)*tFinal/float64(samples-1), worstAt, + hist.FloatAt(worstSlot*n*n+worstAt), + math.Cos(float64(worstSlot*every)*theta)*u0.FloatAt(worstAt), worst) + } +} + +// TestIntegrateRefusesInfiniteIntegrand pins the non-finite gate: an +// integrand returning +Inf used to poison the error sum with NaN, +// whose comparisons are false, and the integral came back as Inf with +// a nil error. +func TestIntegrateRefusesInfiniteIntegrand(t *testing.T) { + inf := func(x float64) (float64, error) { return math.Inf(1), nil } + if _, _, err := IntegrateFunction(inf, 0, 1, QuadratureOptions{}); err == nil { + t.Fatal("IntegrateFunction with an infinite integrand returned no error") + } +} diff --git a/integrate/wave_bench_test.go b/integrate/wave_bench_test.go new file mode 100644 index 0000000..72f8aff --- /dev/null +++ b/integrate/wave_bench_test.go @@ -0,0 +1,152 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package integrate + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the library-side cost of the solver drivers: the +// same workloads the package benchmarks run, with the fixture's own +// per-call allocation removed, so the allocation count a run reports +// is the library's alone. Every fixture derivative refills one output +// array the driver copies out of immediately, which the f contract +// allows: no solver retains the returned array. + +// wavePooledDecay builds the diagonal stiff system y' = rate·(i+1)·y_i +// with an f that refills one shared output array per call. +func wavePooledDecay(n int, rate float64) func(t float64, y *core.Array) (*core.Array, error) { + rates := make([]float64, n) + for i := range rates { + rates[i] = rate * float64(i+1) + } + out := core.New(core.Float, n) + vals := out.RawFloats() + return func(t float64, y *core.Array) (*core.Array, error) { + ys := y.RawFloats() + for i := range n { + vals[i] = rates[i] * ys[i] + } + return out, nil + } +} + +// waveVector wraps a fixed literal as a rank-1 array. +func waveVector(b *testing.B, vals []float64) *core.Array { + b.Helper() + a, err := core.FromFloats(vals, len(vals)) + if err != nil { + b.Fatal(err) + } + return a +} + +// waveOnes returns a vector of n ones. +func waveOnes(n int) []float64 { + vals := make([]float64, n) + for i := range vals { + vals[i] = 1 + } + return vals +} + +// BenchmarkWaveLibraryRK4 runs the fixed-step RK4 loop over the linear +// system of BenchmarkIntegrateRK4, with the derivative evaluation +// allocation-free: the reported allocations are the driver's own. +func BenchmarkWaveLibraryRK4(b *testing.B) { + const n = 16 + f := wavePooledDecay(n, -0.25) + start := waveVector(b, waveOnes(n)) + b.ReportAllocs() + for b.Loop() { + if _, err := IntegrateRK4(f, 0, 10, start, 2000); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkWaveLibraryBDF2 runs the variable-step BDF2 loop over the +// stiff diagonal system of BenchmarkIntegrateBDF2, allocation-free on +// the fixture side. +func BenchmarkWaveLibraryBDF2(b *testing.B) { + const n = 4 + f := wavePooledDecay(n, -100) + start := waveVector(b, waveOnes(n)) + b.ReportAllocs() + for b.Loop() { + if _, err := IntegrateBDF2(f, 0, 1, start, ODEOptions{}); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkWaveLibraryBDFVar runs the variable-order BDF loop over the +// stiff diagonal system of BenchmarkBDFVarStiff, allocation-free on the +// fixture side. +func BenchmarkWaveLibraryBDFVar(b *testing.B) { + const n = 32 + f := wavePooledDecay(n, -100) + start := waveVector(b, waveOnes(n)) + opts := BDFVarOptions{RelTol: 1e-6, AbsTol: 1e-9} + b.ReportAllocs() + for b.Loop() { + if _, err := IntegrateBDFVar(f, 0, 1, start, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkWaveLibraryROS4 runs the ROS4 loop over the stiff diagonal +// system of BenchmarkROS4Stiff, allocation-free on the fixture side. +func BenchmarkWaveLibraryROS4(b *testing.B) { + const n = 32 + f := wavePooledDecay(n, -100) + start := waveVector(b, waveOnes(n)) + opts := ODEOptions{RelTol: 1e-6, AbsTol: 1e-9} + b.ReportAllocs() + for b.Loop() { + if _, err := IntegrateROS4(f, 0, 1, start, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkWaveLibraryVerlet runs the velocity Verlet loop over the +// coupled oscillators of BenchmarkIntegrateVerlet, with the +// acceleration evaluation allocation-free: the reported allocations +// are the driver's own. +func BenchmarkWaveLibraryVerlet(b *testing.B) { + n := 32 + q0 := waveOnes(n) + for i := range q0 { + q0[i] = math.Sin(float64(i)) + } + out := core.New(core.Float, n) + vals := out.RawFloats() + accel := func(q *core.Array) (*core.Array, error) { + qs := q.RawFloats() + for i := range n { + l, r := 0.0, 0.0 + if i > 0 { + l = qs[i-1] + } + if i < n-1 { + r = qs[i+1] + } + vals[i] = l - 2*qs[i] + r + } + return out, nil + } + qs := waveVector(b, q0) + ps := waveVector(b, waveOnes(n)) + b.ReportAllocs() + for b.Loop() { + if _, _, err := IntegrateVerlet(accel, 0, 10, qs, ps, 500); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/base/base.go b/internal/base/base.go new file mode 100644 index 0000000..af7d353 --- /dev/null +++ b/internal/base/base.go @@ -0,0 +1,428 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package base holds the low-level primitives the library's domain +// packages share: error construction, shape formatting, the machine +// epsilon and the generic LU factorisation the solvers build on. It +// touches no arrays, so every package in the module can depend on it +// without a cycle. +package base + +import ( + "fmt" + "math" + "math/cmplx" + "strconv" + "strings" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// EpsF is the float64 machine epsilon. +const EpsF = 2.220446049250313e-16 + +// Errf builds a package-prefixed error: every error the library +// returns starts with "tensor: ", whatever package raised it. +func Errf(format string, args ...any) error { + return fmt.Errorf("tensor: "+format, args...) +} + +// WrapErr wraps an error the library already prefixed, naming the +// operation that surfaces it. The chain stays intact for errors.Is and +// errors.As, and the message carries the tensor tag and the operation +// once each instead of the doubled prefix a plain Errf wrap of a +// prefixed error prints. +func WrapErr(operation string, err error) error { + if err == nil { + return nil + } + return &wrappedErr{op: operation, err: err} +} + +type wrappedErr struct { + op string + err error +} + +func (e *wrappedErr) Error() string { + return "tensor: " + e.op + ": " + strings.TrimPrefix(e.err.Error(), "tensor: ") +} + +func (e *wrappedErr) Unwrap() error { return e.err } + +// ShapeText renders a shape as (n, m, ...). +func ShapeText(shape []int) string { + if len(shape) == 0 { + return "scalar" + } + parts := make([]string, len(shape)) + for i, d := range shape { + parts[i] = strconv.Itoa(d) + } + return "(" + strings.Join(parts, ", ") + ")" +} + +// AbsComplex returns |z|, the magnitude of a complex number. It is +// math.Hypot rather than sqrt(r² + i²), which overflows above about +// 1.34e154 and underflows below about 1.5e-162 even though the true +// magnitude is an ordinary number there. +func AbsComplex(z complex128) float64 { + return math.Hypot(real(z), imag(z)) +} + +// RangeN returns 0..n-1. +func RangeN(n int) []int { + out := make([]int, n) + for i := range out { + out[i] = i + } + return out +} + +// CmplxPolar builds a complex number from magnitude and angle. +func CmplxPolar(r, theta float64) complex128 { + return complex(r*math.Cos(theta), r*math.Sin(theta)) +} + +// Real2 returns |z|², the squared magnitude of a complex number. +func Real2(z complex128) float64 { + return real(z)*real(z) + imag(z)*imag(z) +} + +// transposeBlock is the tile side of TransposeFlat: the destination +// column segment a tile writes spans transposeBlock consecutive +// entries of the tile's own rows, which keeps the strided side of the +// copy inside the cache instead of paying a miss per element. +const transposeBlock = 16 + +// TransposeFlat transposes a row-major m×n matrix held flat. The copy +// is tiled, so the strided side walks transposeBlock rows at a time +// rather than the whole matrix. A pure permutation of values, so the +// tiling cannot move a bit. +func TransposeFlat(a []float64, m, n int) []float64 { + out := make([]float64, m*n) + for i0 := 0; i0 < m; i0 += transposeBlock { + i1 := min(i0+transposeBlock, m) + for j0 := 0; j0 < n; j0 += transposeBlock { + j1 := min(j0+transposeBlock, n) + for i := i0; i < i1; i++ { + row := a[i*n+j0 : i*n+j1] + dst := j0*m + i + _ = out[dst] // bounds-check proof: the first tile entry + for _, v := range row { + out[dst] = v + dst += m + } + } + } + } + return out +} + +// Scalar is the element type the generic linear algebra serves. +type Scalar interface { + float64 | float32 | complex128 +} + +// factorWorkQuantum is the element-update budget one worker receives in +// a rank-1 update dispatch, counted as rows remaining × columns +// remaining. A pivot whose update does not fill one quantum runs on the +// calling goroutine: the dispatch would cost more than the update. The +// crew is sized by the work, never by the machine, because the updates +// shrink every pivot and a wide dispatch of a small update costs more +// than it saves. Measured over the rank-1 update alone at the sizes the +// solvers reach, the best crew shrinks with the quantum and the curve is +// flat between 8 192 and 49 152, one size preferring a smaller crew and +// the next a larger one; this value keeps both within a few percent of +// their own optimum. The pivot SEARCH and the row SWAP stay strictly +// serial: they pick the pivot every update reads, and their order +// defines the factorisation. +const factorWorkQuantum = 24576 + +// factorSpawnFloor is the update size from which any crew may spawn. +// Crew SIZE above it stays tied to factorWorkQuantum; below it the whole +// update runs on the calling goroutine. The paired crossover probe +// (BenchmarkFactorCrossoverAB) measures the lone walk ahead of every +// crew through n = 384, whose largest update is 147 072, and the shipped +// crew ahead of the lone walk from n = 448, whose largest update is +// 200 256; the floor sits between the two measured updates. The crew +// size never moves a bit: every row keeps the serial per-row sequence +// whatever dispatch carries it (TestFactorDispatchBitIdentical). +const factorSpawnFloor = 147456 + +// factorMaxCrew caps the crew a single dispatch may use. Past this the +// per-goroutine spawn and join cost outgrows the update it spreads: an +// empty fork-join of w goroutines costs about 0.2 µs per goroutine, so a +// 512² pivot on a 16-way crew already pays more in synchronisation than +// the last worker's chunk earns. +const factorMaxCrew = 16 + +// factorMinRows is the floor on rows per worker for the rank-1 update; +// fewer rows per goroutine than this only adds scheduling latency. +const factorMinRows = 4 + +// factorJob is one crew member's slice of a pivot's rank-1 update: the +// operands every member shares and the row range this member owns. The +// jobs are allocated once per Factor call and rewritten in place per +// pivot, so a dispatch allocates the goroutine's argument frame only, +// never a closure. +type factorJob[T Scalar] struct { + rows [][]T + pivotRow []T + k, n int + start int + end int + wg *sync.WaitGroup +} + +// factorUpdate runs one job through the serial per-row sequence, the +// same call the inline path makes, so the crew size never moves a bit. +func factorUpdate[T Scalar](job *factorJob[T]) { + factorRows(job.rows[job.start:job.end], job.k, job.n, job.pivotRow) + job.wg.Done() +} + +// factorRows subtracts f·pivotRow from every row of rows, the rank-1 +// update of one pivot: f is the row's multiplier, then the products are +// subtracted along j ascending. The whole family of entry points shares +// this one function, which is what makes the parallel result +// bit-identical to the serial one. +func factorRows[T Scalar](rows [][]T, k, n int, pivotRow []T) { + pk := pivotRow[k] + _ = pivotRow[n-1] // bounds-check proof: every pivot row spans n entries + for _, row := range rows { + f := row[k] / pk + row[k] = f + _ = row[n-1] // bounds-check proof: every row spans n entries + for j := k + 1; j < n; j++ { + row[j] -= f * pivotRow[j] + } + } +} + +// Factor factors a row-major square matrix in place with partial +// pivoting, returning the row permutation and the parity of the +// permutation. The stored subdiagonal holds the elimination factors, +// so a zero pivot column is skipped and reported through the zero +// diagonal the singular check reads. +// +// The elimination is parallelised per pivot over the rows below it. +// Every row's update reads only the pivot row, which is final once the +// search and swap have run, and writes only its own row: each element +// is updated exactly once per pivot with the same operands and in the +// same order as the serial loop (f, then f·pivot[j] subtracted along +// j ascending), so the result is bit-identical for either element +// type, complex128 included. +func Factor[T Scalar](m [][]T) ([]int, int) { + n := len(m) + perm := make([]int, n) + for i := range perm { + perm[i] = i + } + parity := 1 + var wg sync.WaitGroup + var jobs []factorJob[T] + for k := range n { + pivot := k + for i := k + 1; i < n; i++ { + if absOf(m[i][k]) > absOf(m[pivot][k]) { + pivot = i + } + } + if m[pivot][k] == 0 { + continue // the whole column is zero: singular from here on + } + if pivot != k { + m[pivot], m[k] = m[k], m[pivot] + perm[pivot], perm[k] = perm[k], perm[pivot] + parity = -parity + } + pivotRow := m[k] + rows := m[k+1:] + // Crew sized by the update's size, not by the machine: the + // updates shrink every pivot, and a 32-goroutine dispatch of a + // small update costs more than it saves. Below one quantum the + // update cannot fund even a two-way split, and below the spawn + // floor no crew at all earns its fork-join, so the update runs + // inline. + work := len(rows) * (n - k) + w := min(work/factorWorkQuantum+1, engine.WorkersFor(len(rows))) + w = min(w, factorMaxCrew) + if work < factorSpawnFloor { + w = 1 + } + if w < 2 { + factorRows(rows, k, n, pivotRow) + continue + } + chunk := (len(rows) + w - 1) / w + if chunk < factorMinRows { + w = max(len(rows)/factorMinRows, 1) + chunk = (len(rows) + w - 1) / w + if w < 2 { + factorRows(rows, k, n, pivotRow) + continue + } + } + // Each row's update reads only the pivot row, which is final + // once the search and swap have run, and writes only its own + // row: the per-row sequence (f, then f·pivotRow[j] subtracted + // along j ascending) is the serial one, so the crew size and + // the chunk boundaries never move a bit. The crew runs over + // preallocated slots, so a dispatch allocates the goroutine's + // argument frame and nothing else. + if jobs == nil { + jobs = make([]factorJob[T], factorMaxCrew) + } + spawned := 0 + for start := 0; start < len(rows); start += chunk { + job := &jobs[spawned] + job.rows, job.pivotRow, job.k, job.n = rows, pivotRow, k, n + job.start, job.end, job.wg = start, min(start+chunk, len(rows)), &wg + spawned++ + } + wg.Add(spawned) + for i := range spawned { + go factorUpdate(&jobs[i]) + } + wg.Wait() + } + return perm, parity +} + +// CheckSingular reports an error when a factored matrix carries a +// zero pivot. +func CheckSingular[T Scalar](name string, m [][]T) error { + for i := range m { + if m[i][i] == 0 { + return Errf("%s: matrix is singular", name) + } + } + return nil +} + +// SolveColumn back-substitutes one right-hand column through the +// factored matrix. It is a recurrence along the column (each col[i] +// reads every earlier entry) and is intentionally serial; the hoisted +// row slices and the index proofs only strip bounds checks, the +// per-element operation sequence is untouched. +func SolveColumn[T Scalar](m [][]T, col []T) { + n := len(m) + if n == 0 { + return + } + _ = col[n-1] // bounds-check proof: every index below stays in range + for i := 1; i < n; i++ { + row := m[i] + _ = row[i-1] + ci := col[i] + for j := range i { + ci -= row[j] * col[j] + } + col[i] = ci + } + for i := n - 1; i >= 0; i-- { + row := m[i] + _ = row[n-1] + ci := col[i] + for j := i + 1; j < n; j++ { + ci -= row[j] * col[j] + } + col[i] = ci / m[i][i] + } +} + +// PermuteColumn reorders col in place so that col[i] takes the value +// that sat at perm[i]. The permutation is applied by walking the cycles +// of perm: every value moves exactly once with plain assignment, bit +// for bit, and perm itself is never written (callers reuse it across +// columns). The only allocation is the visit bitmap, not a copy of the +// column. +func PermuteColumn[T Scalar](col []T, perm []int) { + permuteColumn(col, perm, make([]bool, len(col))) +} + +// permuteColumn is PermuteColumn with the caller supplying the visit +// bitmap, so a caller solving many columns can reuse one bitmap instead +// of allocating one per column. The bitmap is cleared on entry, so a +// dirty buffer behaves exactly like a fresh one. +func permuteColumn[T Scalar](col []T, perm []int, visited []bool) { + clear(visited[:len(col)]) + for i := range col { + if visited[i] || perm[i] == i { + visited[i] = true + continue + } + // Carry the displaced value around the cycle. + tmp := col[i] + j := i + for { + visited[j] = true + k := perm[j] + if k == i { + break + } + col[j] = col[k] + j = k + } + col[j] = tmp + } +} + +// SolveSystem factors a fresh copy-safe matrix, rejects singularity +// and solves every right-hand column in one pass: each column is +// permuted the way factoring moved the rows, then substitutes through +// L and U. m is consumed in place; callers hand over freshly built +// matrices. +// +// With several right-hand columns the substitution runs in parallel +// over the columns: a column's permutation and back-substitution read +// only the factored matrix and write only that column, so each column +// runs the identical serial sequence on identical inputs, and the +// results are bit-identical to solving the columns one after another. +// SolveColumn itself is a recurrence along the column and stays +// serial. +func SolveSystem[T Scalar](name string, m [][]T, rhs [][]T) ([][]T, error) { + perm, _ := Factor(m) + if err := CheckSingular(name, m); err != nil { + return nil, err + } + if len(rhs) > 1 { + // Each crew member owns one visit bitmap reused across its + // columns: permuteColumn clears it on entry, so the walk sees + // a fresh bitmap without an allocation per column. Bitmaps are + // never shared between crew members, so there is no race. + engine.ParallelMin(len(rhs), 1, func(start, end int) { + visited := make([]bool, len(m)) + for _, col := range rhs[start:end] { + permuteColumn(col, perm, visited) + SolveColumn(m, col) + } + }) + return rhs, nil + } + // The serial path reuses one bitmap across the columns for the + // same reason the parallel path does: permuteColumn clears it on + // entry, so a dirty buffer behaves exactly like a fresh one. + visited := make([]bool, len(m)) + for _, col := range rhs { + permuteColumn(col, perm, visited) + SolveColumn(m, col) + } + return rhs, nil +} + +func absOf[T Scalar](v T) float64 { + switch x := any(v).(type) { + case float64: + return math.Abs(x) + case float32: + return math.Abs(float64(x)) + case complex128: + return cmplx.Abs(x) + } + // Every type the Scalar constraint admits is handled above; the + // clause is what tells the compiler so. + return 0 +} diff --git a/internal/base/factor_bench_test.go b/internal/base/factor_bench_test.go new file mode 100644 index 0000000..3cb54c4 --- /dev/null +++ b/internal/base/factor_bench_test.go @@ -0,0 +1,431 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package base + +import ( + "fmt" + "sync" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The dispatch of Factor is the one place where a crew size, a chunk +// boundary and a spawn style can silently change the arithmetic, so the +// candidates live here beside a bit-identity test and the measurement +// that picks the constants. Every candidate must reproduce factorSerial +// byte for byte, and the production Factor must reproduce the +// parameterised mirror at the shipped constants. + +func benchFactorRows(n int) (rows [][]float64, pristine []float64) { + flat := make([]float64, n*n) + for i := range n { + for j := range n { + flat[i*n+j] = float64((i*7+j*13)%11) - 5 + } + flat[i*n+i] += float64(n) + } + rows = make([][]float64, n) + for i := range n { + rows[i] = flat[i*n : (i+1)*n] + } + pristine = make([]float64, len(flat)) + copy(pristine, flat) + return rows, pristine +} + +func resetFactorRows(rows [][]float64, pristine []float64) { + off := 0 + for _, row := range rows { + copy(row, pristine[off:off+len(row)]) + off += len(row) + } +} + +// factorSerial is Factor with the dispatch removed: the reference every +// candidate must match byte for byte. +func factorSerial(m [][]float64) ([]int, int) { + n := len(m) + perm := make([]int, n) + for i := range perm { + perm[i] = i + } + parity := 1 + for k := range n { + pivot := k + for i := k + 1; i < n; i++ { + if absOf(m[i][k]) > absOf(m[pivot][k]) { + pivot = i + } + } + if m[pivot][k] == 0 { + continue + } + if pivot != k { + m[pivot], m[k] = m[k], m[pivot] + perm[pivot], perm[k] = perm[k], perm[pivot] + parity = -parity + } + factorRows(m[k+1:], k, n, m[k]) + } + return perm, parity +} + +// factorRetired is the dispatch Factor used before the crew moved onto +// preallocated slots: one goroutine per chunk through +// sync.WaitGroup.Go, so every chunk allocates a closure. It is the +// reference the current dispatch is measured against. +func factorRetired(m [][]float64, quantum int) ([]int, int) { + n := len(m) + perm := make([]int, n) + for i := range perm { + perm[i] = i + } + parity := 1 + for k := range n { + pivot := k + for i := k + 1; i < n; i++ { + if absOf(m[i][k]) > absOf(m[pivot][k]) { + pivot = i + } + } + if m[pivot][k] == 0 { + continue + } + if pivot != k { + m[pivot], m[k] = m[k], m[pivot] + perm[pivot], perm[k] = perm[k], perm[pivot] + parity = -parity + } + pivotRow := m[k] + rows := m[k+1:] + work := len(rows) * (n - k) + w := min(work/quantum+1, engine.WorkersFor(len(rows))) + if w < 2 { + factorRows(rows, k, n, pivotRow) + continue + } + chunk := (len(rows) + w - 1) / w + if chunk < factorMinRows { + w = max(len(rows)/factorMinRows, 1) + chunk = (len(rows) + w - 1) / w + } + var wg sync.WaitGroup + for start := 0; start < len(rows); start += chunk { + end := min(start+chunk, len(rows)) + wg.Go(func() { + factorRows(rows[start:end], k, n, pivotRow) + }) + } + wg.Wait() + } + return perm, parity +} + +// factorJobSlot is one crew member's preallocated work slot, mirroring +// factorJob for the float64 candidates. +type factorJobSlot struct { + rows [][]float64 + pivotRow []float64 + k, n int + start int + end int + wg *sync.WaitGroup +} + +func factorSlotWorker(job *factorJobSlot) { + factorRows(job.rows[job.start:job.end], job.k, job.n, job.pivotRow) + job.wg.Done() +} + +// factorCrew is the shipped dispatch with both tuning constants +// exposed. +func factorCrew(m [][]float64, quantum, maxCrew int) ([]int, int) { + n := len(m) + perm := make([]int, n) + for i := range perm { + perm[i] = i + } + parity := 1 + var wg sync.WaitGroup + var jobs []factorJobSlot + for k := range n { + pivot := k + for i := k + 1; i < n; i++ { + if absOf(m[i][k]) > absOf(m[pivot][k]) { + pivot = i + } + } + if m[pivot][k] == 0 { + continue + } + if pivot != k { + m[pivot], m[k] = m[k], m[pivot] + perm[pivot], perm[k] = perm[k], perm[pivot] + parity = -parity + } + pivotRow := m[k] + rows := m[k+1:] + work := len(rows) * (n - k) + w := min(work/quantum+1, engine.WorkersFor(len(rows))) + w = min(w, maxCrew) + if w < 2 { + factorRows(rows, k, n, pivotRow) + continue + } + chunk := (len(rows) + w - 1) / w + if chunk < factorMinRows { + w = max(len(rows)/factorMinRows, 1) + chunk = (len(rows) + w - 1) / w + if w < 2 { + factorRows(rows, k, n, pivotRow) + continue + } + } + if jobs == nil { + jobs = make([]factorJobSlot, maxCrew) + } + spawned := 0 + for start := 0; start < len(rows); start += chunk { + job := &jobs[spawned] + job.rows, job.pivotRow, job.k, job.n = rows, pivotRow, k, n + job.start, job.end, job.wg = start, min(start+chunk, len(rows)), &wg + spawned++ + } + wg.Add(spawned) + for i := range spawned { + go factorSlotWorker(&jobs[i]) + } + wg.Wait() + } + return perm, parity +} + +// TestFactorDispatchBitIdentical proves every crew size, chunk boundary +// and spawn style leaves the factor and the permutation byte for byte +// as the serial reference, and that the shipped Factor matches the +// parameterised mirror the tuning benchmark measures. +func TestFactorDispatchBitIdentical(t *testing.T) { + for _, n := range []int{1, 2, 5, 16, 64, 129, 260} { + base, _ := benchFactorRows(n) + want := make([][]float64, n) + for i := range n { + want[i] = append([]float64(nil), base[i]...) + } + wp, wpar := factorSerial(want) + check := func(name string, got [][]float64, gp []int, gpar int) { + t.Helper() + if len(gp) != len(wp) { + t.Fatalf("%s: permutation length %d, want %d", name, len(gp), len(wp)) + } + for i := range gp { + if gp[i] != wp[i] { + t.Fatalf("%s: permutation[%d] = %d, want %d", name, i, gp[i], wp[i]) + } + } + for i := range got { + for j := range got[i] { + if got[i][j] != want[i][j] { + t.Fatalf("%s: factor[%d][%d] = %v, want %v", name, i, j, got[i][j], want[i][j]) + } + } + } + if gpar != wpar { + t.Fatalf("%s: parity %d, want %d", name, gpar, wpar) + } + } + for _, q := range []int{1024, 8192, 65536} { + for _, w0 := range []int{1, 2, 4, 8, 16, 32} { + rows, _ := benchFactorRows(n) + gp, gpar := factorCrew(rows, q, w0) + check(fmt.Sprintf("crew q=%d w=%d n=%d", q, w0, n), rows, gp, gpar) + } + rows, _ := benchFactorRows(n) + gp, gpar := factorRetired(rows, q) + check(fmt.Sprintf("retired q=%d n=%d", q, n), rows, gp, gpar) + } + rows, _ := benchFactorRows(n) + gp, gpar := Factor(rows) + check(fmt.Sprintf("Factor n=%d", n), rows, gp, gpar) + mirror, _ := benchFactorRows(n) + mirrorGP, mirrorPar := factorCrew(mirror, factorWorkQuantum, factorMaxCrew) + check(fmt.Sprintf("mirror n=%d", n), mirror, mirrorGP, mirrorPar) + } +} + +// BenchmarkForkJoinFloor measures an empty fork-join: the spawn, wake +// and join of w goroutines that touch nothing. It is the lower bound on +// what one dispatch costs, and the reason a pivot whose update is +// smaller than this cannot be worth splitting. +func BenchmarkForkJoinFloor(b *testing.B) { + for _, w := range []int{1, 2, 4, 8, 16, 32} { + b.Run(fmt.Sprintf("w=%d", w), func(b *testing.B) { + for b.Loop() { + var wg sync.WaitGroup + for range w { + wg.Go(func() {}) + } + wg.Wait() + } + }) + } +} + +// BenchmarkFactorStylesAB interleaves the dispatch styles iteration by +// iteration on separate matrices, so a drift in the machine's speed +// lands on every style equally: the reported ns/op per style is the +// paired comparison, and the ns/op of the group as a whole is not +// comparable across groups. +func BenchmarkFactorStylesAB(b *testing.B) { + styles := []struct { + name string + run func(m [][]float64) + }{ + {"serial", func(m [][]float64) { factorSerial(m) }}, + {"retired/q=16384", func(m [][]float64) { factorRetired(m, 16384) }}, + {"slots/q=8192/cap=16", func(m [][]float64) { factorCrew(m, 8192, 16) }}, + {"slots/q=16384/cap=16", func(m [][]float64) { factorCrew(m, 16384, 16) }}, + {"slots/q=24576/cap=16", func(m [][]float64) { factorCrew(m, 24576, 16) }}, + {"slots/q=32768/cap=16", func(m [][]float64) { factorCrew(m, 32768, 16) }}, + } + for _, n := range []int{256, 512} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + rows := make([][][]float64, len(styles)) + pristine := make([][]float64, len(styles)) + for i := range styles { + rows[i], pristine[i] = benchFactorRows(n) + } + elapsed := make([]time.Duration, len(styles)) + for b.Loop() { + for i, s := range styles { + resetFactorRows(rows[i], pristine[i]) + start := time.Now() + s.run(rows[i]) + elapsed[i] += time.Since(start) + } + } + for i, s := range styles { + b.ReportMetric(float64(elapsed[i].Nanoseconds())/float64(b.N), "ns/op-"+s.name) + } + }) + } +} + +// transposeFlatPlain is the untiled copy the tiled TransposeFlat +// replaces: the read runs along the rows, the write strides by the row +// length. +func transposeFlatPlain(a []float64, m, n int) []float64 { + out := make([]float64, m*n) + for i := range m { + for j := range n { + out[j*m+i] = a[i*n+j] + } + } + return out +} + +// TestTransposeFlatTiledBitIdentical pins the tiled copy against the +// untiled one: a transpose is a permutation, so it must match value for +// value, and the tile boundary must not lose a corner of a ragged +// matrix. +func TestTransposeFlatTiledBitIdentical(t *testing.T) { + for _, dim := range [][2]int{{1, 1}, {1, 7}, {7, 1}, {16, 16}, {17, 16}, {16, 17}, {33, 5}, {5, 33}, {64, 64}} { + m, n := dim[0], dim[1] + a := make([]float64, m*n) + for i := range a { + a[i] = float64(i%13) - 6 + } + got := TransposeFlat(a, m, n) + want := transposeFlatPlain(a, m, n) + if len(got) != len(want) { + t.Fatalf("%dx%d: length %d, want %d", m, n, len(got), len(want)) + } + for i := range got { + if got[i] != want[i] { + t.Fatalf("%dx%d: entry %d = %v, want %v", m, n, i, got[i], want[i]) + } + } + } +} + +// BenchmarkTransposeFlatAB interleaves the tiled copy with the untiled +// one on separate buffers, so a drift in the machine's speed lands on +// both: the reported ns/op per style is the paired comparison. +func BenchmarkTransposeFlatAB(b *testing.B) { + for _, n := range []int{64, 256} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + flat := make([]float64, n*n) + for i := range flat { + flat[i] = float64(i%17) - 8 + } + var tiled, plain time.Duration + for b.Loop() { + start := time.Now() + out := TransposeFlat(flat, n, n) + tiled += time.Since(start) + if len(out) != len(flat) { + b.Fatal("transpose lost entries") + } + start = time.Now() + out = transposeFlatPlain(flat, n, n) + plain += time.Since(start) + if len(out) != len(flat) { + b.Fatal("transpose lost entries") + } + } + b.ReportMetric(float64(tiled.Nanoseconds())/float64(b.N), "ns/op-tiled") + b.ReportMetric(float64(plain.Nanoseconds())/float64(b.N), "ns/op-plain") + }) + } +} + +// BenchmarkTransposeFlat measures the tiled transpose alone. +func BenchmarkTransposeFlat(b *testing.B) { + for _, n := range []int{64, 256} { + flat := make([]float64, n*n) + for i := range flat { + flat[i] = float64(i%17) - 8 + } + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + out := TransposeFlat(flat, n, n) + if len(out) != len(flat) { + b.Fatal("transpose lost entries") + } + } + }) + } +} + +// BenchmarkSolveSystemColumns measures the column loop of SolveSystem +// with one column and with several: the several-column path permutes +// each column through a bitmap the crew reuses, which is what the +// one-column path does too. +func BenchmarkSolveSystemColumns(b *testing.B) { + const n = 256 + for _, cols := range []int{1, 8} { + b.Run(fmt.Sprintf("cols=%d", cols), func(b *testing.B) { + rhs := make([][]float64, cols) + for c := range rhs { + rhs[c] = make([]float64, n) + for i := range rhs[c] { + rhs[c][i] = float64((i+c)%9) - 4 + } + } + b.ReportAllocs() + for b.Loop() { + m, _ := benchFactorRows(n) + work := make([][]float64, cols) + for c := range work { + work[c] = append([]float64(nil), rhs[c]...) + } + if _, err := SolveSystem("Solve", m, work); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/internal/base/factor_crossover_test.go b/internal/base/factor_crossover_test.go new file mode 100644 index 0000000..95defac --- /dev/null +++ b/internal/base/factor_crossover_test.go @@ -0,0 +1,189 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package base + +import ( + "fmt" + "sync" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Paired crossover probe for Factor's dispatch: the serial elimination +// against the shipped crew constants on the same deterministic matrix, +// interleaved iteration by iteration so a drift in the machine's speed +// lands on both styles equally. The reported ns/op per style is the +// paired comparison; the ns/op of the group as a whole is not +// comparable across groups. The spawn floor factorSpawnFloor sits where +// the crew starts winning. + +// probeFactorRows builds the deterministic n×n fixture: a diagonally +// dominant integer pattern every style factors identically. +func probeFactorRows(n int) (rows [][]float64, pristine []float64) { + flat := make([]float64, n*n) + for i := range n { + for j := range n { + flat[i*n+j] = float64((i*7+j*13)%11) - 5 + } + flat[i*n+i] += float64(n) + } + rows = make([][]float64, n) + for i := range n { + rows[i] = flat[i*n : (i+1)*n] + } + pristine = make([]float64, len(flat)) + copy(pristine, flat) + return rows, pristine +} + +func probeFactorReset(rows [][]float64, pristine []float64) { + off := 0 + for _, row := range rows { + copy(row, pristine[off:off+len(row)]) + off += len(row) + } +} + +// probeFactorSerial is Factor with the dispatch removed: the reference +// every candidate must match byte for byte. +func probeFactorSerial(m [][]float64) { + n := len(m) + for k := range n { + pivot := k + for i := k + 1; i < n; i++ { + if absOf(m[i][k]) > absOf(m[pivot][k]) { + pivot = i + } + } + if m[pivot][k] == 0 { + continue + } + if pivot != k { + m[pivot], m[k] = m[k], m[pivot] + } + factorRows(m[k+1:], k, n, m[k]) + } +} + +// probeFactorJob mirrors factorJob for the probe's own dispatch, so the +// probe keeps compiling whatever the production type later gains. +type probeFactorJob struct { + rows [][]float64 + pivotRow []float64 + k, n int + start int + end int + wg *sync.WaitGroup +} + +func probeFactorWorker(job *probeFactorJob) { + factorRows(job.rows[job.start:job.end], job.k, job.n, job.pivotRow) + job.wg.Done() +} + +// probeFactorCrew mirrors Factor's dispatch at the production constants, +// which it reads directly: the quantum sizes the crew, the spawn floor +// decides whether any crew runs, and the per-row update is the same +// factorRows call the serial reference makes. +func probeFactorCrew(m [][]float64) { + n := len(m) + var wg sync.WaitGroup + var jobs []probeFactorJob + for k := range n { + pivot := k + for i := k + 1; i < n; i++ { + if absOf(m[i][k]) > absOf(m[pivot][k]) { + pivot = i + } + } + if m[pivot][k] == 0 { + continue + } + if pivot != k { + m[pivot], m[k] = m[k], m[pivot] + } + pivotRow := m[k] + rows := m[k+1:] + work := len(rows) * (n - k) + w := min(work/factorWorkQuantum+1, engine.WorkersFor(len(rows))) + w = min(w, factorMaxCrew) + if work < factorSpawnFloor { + w = 1 + } + if w < 2 { + factorRows(rows, k, n, pivotRow) + continue + } + chunk := (len(rows) + w - 1) / w + if chunk < factorMinRows { + w = max(len(rows)/factorMinRows, 1) + chunk = (len(rows) + w - 1) / w + if w < 2 { + factorRows(rows, k, n, pivotRow) + continue + } + } + if jobs == nil { + jobs = make([]probeFactorJob, factorMaxCrew) + } + spawned := 0 + for start := 0; start < len(rows); start += chunk { + job := &jobs[spawned] + job.rows, job.pivotRow, job.k, job.n = rows, pivotRow, k, n + job.start, job.end, job.wg = start, min(start+chunk, len(rows)), &wg + spawned++ + } + wg.Add(spawned) + for i := range spawned { + go probeFactorWorker(&jobs[i]) + } + wg.Wait() + } +} + +// TestFactorSpawnFloorBitIdentical pins the crew mirror against the +// serial reference at the sizes the crossover probe sweeps, whatever +// the dispatch decides: crew size never moves a bit. +func TestFactorSpawnFloorBitIdentical(t *testing.T) { + for _, n := range []int{2, 3, 16, 64, 256, 320, 384} { + want, _ := probeFactorRows(n) + probeFactorSerial(want) + got, _ := probeFactorRows(n) + probeFactorCrew(got) + for i := range n { + for j := range n { + if got[i][j] != want[i][j] { + t.Fatalf("n=%d factor[%d][%d] = %v, want %v", n, i, j, got[i][j], want[i][j]) + } + } + } + } +} + +// BenchmarkFactorCrossoverAB interleaves the serial elimination with the +// shipped dispatch iteration by iteration across the sizes around the +// spawn floor. +func BenchmarkFactorCrossoverAB(b *testing.B) { + for _, n := range []int{256, 320, 384, 448, 512} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + rowsS, prS := probeFactorRows(n) + rowsC, prC := probeFactorRows(n) + var tSerial, tCrew time.Duration + for b.Loop() { + probeFactorReset(rowsS, prS) + probeFactorReset(rowsC, prC) + start := time.Now() + probeFactorSerial(rowsS) + tSerial += time.Since(start) + start = time.Now() + probeFactorCrew(rowsC) + tCrew += time.Since(start) + } + b.ReportMetric(float64(tSerial.Nanoseconds())/float64(b.N), "ns/op-serial") + b.ReportMetric(float64(tCrew.Nanoseconds())/float64(b.N), "ns/op-crew") + }) + } +} diff --git a/internal/base/scalar_abs_pin_test.go b/internal/base/scalar_abs_pin_test.go new file mode 100644 index 0000000..a51ac6b --- /dev/null +++ b/internal/base/scalar_abs_pin_test.go @@ -0,0 +1,129 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package base + +import ( + "math" + "testing" +) + +// TestAbsOfEveryScalarType pins absOf for every element type the Scalar +// constraint admits. The float32 case used to fall through to the 0 the +// switch returns for an unhandled type, so a Factor[float32] pivot +// search compared every element against that zero and lost its partial +// pivoting silently. +func TestAbsOfEveryScalarType(t *testing.T) { + // float64: the plain magnitude, including the negative zero and the + // signed extremes. + for _, c := range []struct { + in float64 + want float64 + }{ + {3.5, 3.5}, + {-3.5, 3.5}, + {0, 0}, + {math.Copysign(0, -1), 0}, + {-math.MaxFloat64, math.MaxFloat64}, + {math.Inf(-1), math.Inf(1)}, + } { + if got := absOf(c.in); got != c.want { + t.Errorf("absOf(float64(%v)) = %v, want %v", c.in, got, c.want) + } + } + // float32: what Factor[float32] reads. The value must be the exact + // magnitude widened to float64, not a zero. + for _, c := range []struct { + in float32 + want float64 + }{ + {3.5, 3.5}, + {-3.5, 3.5}, + {0, 0}, + {float32(math.Copysign(0, -1)), 0}, + {-math.MaxFloat32, float64(math.MaxFloat32)}, + {1e-30, float64(float32(1e-30))}, + } { + got := absOf(c.in) + if got != c.want { + t.Errorf("absOf(float32(%v)) = %v, want %v", c.in, got, c.want) + } + if got == 0 && c.in != 0 { + t.Errorf("absOf(float32(%v)) = 0: the float32 case is missing again", c.in) + } + } + // complex128: the modulus. + for _, c := range []struct { + in complex128 + want float64 + }{ + {complex(3, 4), 5}, + {complex(-3, -4), 5}, + {complex(0, 0), 0}, + {complex(-0.0, 0), 0}, + {complex(1e200, 1e200), math.Sqrt2 * 1e200}, + } { + if got := absOf(c.in); got != c.want { + t.Errorf("absOf(complex128(%v)) = %v, want %v", c.in, got, c.want) + } + } + // The pivot search is the caller that matters: a float32 matrix + // whose largest element is negative must pick that pivot. + m := [][]float32{ + {0.05, 0}, + {-0.1, 0}, + } + perm, _ := Factor(m) + if perm[0] != 1 { + t.Errorf("Factor[float32] picked row %d as the first pivot, want the larger magnitude at row 1", perm[0]) + } +} + +// TestAbsComplexExtremeScale pins AbsComplex beyond the range of the +// sqrt(r² + i²) it used to be. The squared form overflows above about +// 1.34e154 and underflows below about 1.5e-162 while the magnitude is an +// ordinary number in both directions. +func TestAbsComplexExtremeScale(t *testing.T) { + const big = 1e200 + z := complex(big, big) + naive := math.Sqrt(real(z)*real(z) + imag(z)*imag(z)) + if !math.IsInf(naive, 1) { + t.Fatalf("sqrt(r² + i²) at 1e200 = %v, the overflow this test relies on is gone", naive) + } + if got, want := AbsComplex(z), math.Sqrt2*big; math.Abs(got-want) > 1e-15*want { + t.Errorf("AbsComplex(1e200 + 1e200i) = %v, want %v", got, want) + } + // The same pair with one component zero: the magnitude is finite and + // the squared form would still overflow. + if got := AbsComplex(complex(big, 0)); got != big { + t.Errorf("AbsComplex(1e200) = %v, want %v", got, big) + } + // The underflow side. + const small = 1e-200 + if naive := math.Sqrt(small*small + small*small); naive != 0 { + t.Fatalf("sqrt(r² + i²) at 1e-200 = %v, the underflow this test relies on is gone", naive) + } + if got, want := AbsComplex(complex(small, small)), math.Sqrt2*small; math.Abs(got-want) > 1e-15*want { + t.Errorf("AbsComplex(1e-200 + 1e-200i) = %v, want %v", got, want) + } + // Ordinary magnitudes keep their exact values: the classics and a + // Pythagorean pair whose components are exactly representable. + for _, c := range []struct { + in complex128 + want float64 + }{ + {complex(3, 4), 5}, + {complex(5, 12), 13}, + {complex(1, 0), 1}, + {complex(0, -1), 1}, + {complex(0, 0), 0}, + } { + if got := AbsComplex(c.in); got != c.want { + t.Errorf("AbsComplex(%v) = %v, want %v", c.in, got, c.want) + } + } + // A mixed pair: the large component does not swallow the small one. + if got := AbsComplex(complex(1e200, 1e-200)); got != 1e200 { + t.Errorf("AbsComplex(1e200 + 1e-200i) = %v, want 1e200", got) + } +} diff --git a/internal/base/tridiag.go b/internal/base/tridiag.go new file mode 100644 index 0000000..331a412 --- /dev/null +++ b/internal/base/tridiag.go @@ -0,0 +1,51 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package base + +// The tridiagonal Thomas elimination, one kernel for every caller. The +// linalg package's public SolveTridiagonal and the integrate package's +// PDE step sweeps solve the same systems; both call TriSolve, so the +// arithmetic and the refusal texts exist once. The messages name +// SolveTridiagonal because both public surfaces publish that text +// today, and a message a caller matches is part of the contract. + +// TriSolve solves the tridiagonal system with lower diagonal a +// (length n-1), main diagonal b (length n), upper diagonal c (length +// n-1) and right side d (length n), writing the solution into dst. The +// scratch cp and dp must hold at least n elements. A zero pivot is +// refused. Every buffer is fully overwritten before the kernel reads +// it, except cp, whose prefix is written and read in the same sweep +// order a fresh buffer saw, so reused scratch and fresh allocations +// solve bit-identically. dst must not alias any diagonal or d. +func TriSolve(dst, cp, dp, a, b, c, d []float64) error { + n := len(b) + if n == 0 { + return Errf("SolveTridiagonal: empty system") + } + b0 := b[0] + if b0 == 0 { + return Errf("SolveTridiagonal: zero pivot at row 0") + } + // For n = 1 the c diagonal is empty per the length contract, so the + // seed must not read it; the single unknown falls out of dp[0]. + if n > 1 { + cp[0] = c[0] / b0 + } + dp[0] = d[0] / b0 + for i := 1; i < n; i++ { + den := b[i] - a[i-1]*cp[i-1] + if den == 0 { + return Errf("SolveTridiagonal: zero pivot at row %d", i) + } + if i < n-1 { + cp[i] = c[i] / den + } + dp[i] = (d[i] - a[i-1]*dp[i-1]) / den + } + dst[n-1] = dp[n-1] + for i := n - 2; i >= 0; i-- { + dst[i] = dp[i] - cp[i]*dst[i+1] + } + return nil +} diff --git a/internal/base/wrap_pin_test.go b/internal/base/wrap_pin_test.go new file mode 100644 index 0000000..30ee81b --- /dev/null +++ b/internal/base/wrap_pin_test.go @@ -0,0 +1,57 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package base + +import ( + "errors" + "strings" + "testing" +) + +// TestWrapErrCarriesTagOnce pins the WrapErr contract the ten wrap +// points rely on: the message carries the tensor tag and the operation +// once each, whatever the inner error already carries, and the chain +// stays open for errors.Is. The wrap these properties answer to used to +// be fmt.Errorf("op: %w", err) over an Errf-built cause, which printed +// the tag twice. +func TestWrapErrCarriesTagOnce(t *testing.T) { + if got := WrapErr("Anything", nil); got != nil { + t.Fatalf("WrapErr of nil = %v, want nil", got) + } + + // A bare cause the library did not build: the tag and the operation + // are added, once each. + bare := errors.New("the caller's own failure") + got := WrapErr("Jacobian", bare) + if want := "tensor: Jacobian: the caller's own failure"; got.Error() != want { + t.Errorf("WrapErr of a bare cause = %q, want %q", got.Error(), want) + } + if !errors.Is(got, bare) { + t.Errorf("WrapErr of a bare cause does not unwrap to it") + } + + // A cause the library already tagged: the tag is not doubled. + tagged := Errf("the shape holds more elements than fit in an index") + got = WrapErr("Tile", tagged) + if want := "tensor: Tile: the shape holds more elements than fit in an index"; got.Error() != want { + t.Errorf("WrapErr of a tagged cause = %q, want %q", got.Error(), want) + } + if !errors.Is(got, tagged) { + t.Errorf("WrapErr of a tagged cause does not unwrap to it") + } + if n := strings.Count(got.Error(), "tensor: "); n != 1 { + t.Errorf("WrapErr message carries the tag %d times, want exactly one", n) + } + + // A nested WrapErr composes the same way: each layer names its own + // operation and the tag still appears once. + inner := WrapErr("Concat", tagged) + got = WrapErr("Repeat", inner) + if want := "tensor: Repeat: Concat: the shape holds more elements than fit in an index"; got.Error() != want { + t.Errorf("nested WrapErr = %q, want %q", got.Error(), want) + } + if !errors.Is(got, tagged) { + t.Errorf("nested WrapErr does not unwrap down to the cause") + } +} diff --git a/internal/core/argmax_topk_test.go b/internal/core/argmax_topk_test.go new file mode 100644 index 0000000..37e0172 --- /dev/null +++ b/internal/core/argmax_topk_test.go @@ -0,0 +1,215 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +func TestArgMaxAxis(t *testing.T) { + // (3, 4) input, argmax along dim 1 drops the dim: shape (3,) with + // each row holding the column index of its maximum. + a := mustFromFloats(t, []float64{ + 1, 5, 3, 4, // row 0: max at index 1 + 9, 2, 7, 6, // row 1: max at index 0 + 8, 1, 4, 2, // row 2: max at index 0 + }, 3, 4) + got, err := ArgMaxAxis(a, 1) + if err != nil { + t.Fatal(err) + } + if want := []int{3}; !sameShape(got.Shape(), want) { + t.Fatalf("ArgMaxAxis shape: got %v, want %v", got.Shape(), want) + } + for i, w := range []int64{1, 0, 0} { + v, _ := IntAt(got, i) + if v != w { + t.Errorf("ArgMaxAxis dim=1 [%d]: got %d, want %d", i, v, w) + } + } + // argmax along dim 0: each column holds the row index of its max. + // col 0: rows [1, 9, 8] give idx 1; col 1: [5, 2, 1] give 0; + // col 2: [3, 7, 4] give 1; col 3: [4, 6, 2] give 1. + got, err = ArgMaxAxis(a, 0) + if err != nil { + t.Fatal(err) + } + if want := []int{4}; !sameShape(got.Shape(), want) { + t.Fatalf("ArgMaxAxis shape: got %v, want %v", got.Shape(), want) + } + want := []int64{1, 0, 1, 1} + for i, w := range want { + v, _ := IntAt(got, i) + if v != w { + t.Errorf("ArgMaxAxis dim=0 [%d]: got %d, want %d", i, v, w) + } + } +} + +func TestArgMinAxis(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 5, 3, 4, // row 0: min at index 0 (1) + 9, 2, 7, 6, // row 1: min at index 1 (2) + 8, 1, 4, 2, // row 2: min at index 1 (1) + }, 3, 4) + got, err := ArgMinAxis(a, 1) + if err != nil { + t.Fatal(err) + } + if want := []int{3}; !sameShape(got.Shape(), want) { + t.Fatalf("ArgMinAxis shape: got %v, want %v", got.Shape(), want) + } + for i, w := range []int64{0, 1, 1} { + v, _ := IntAt(got, i) + if v != w { + t.Errorf("ArgMinAxis dim=1 [%d]: got %d, want %d", i, v, w) + } + } +} + +func TestArgMaxAxisErrors(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 3) + if _, err := ArgMaxAxis(a, 0); err == nil { + t.Error("ArgMaxAxis: expected error for 1-D input") + } + // Complex rejected. + c, _ := FromComplexes([]complex128{complex(1, 0)}, 1) + if _, err := ArgMaxAxis(c, 0); err == nil { + t.Error("ArgMaxAxis: expected error for complex input") + } +} + +func TestTopK(t *testing.T) { + a := mustFromFloats(t, []float64{3, 1, 4, 1, 5, 9, 2, 6}, 8) + vals, idxs, err := TopK(a, 3, 0) + if err != nil { + t.Fatal(err) + } + if vals.Len() != 3 || idxs.Len() != 3 { + t.Errorf("TopK: wrong length (%d, %d)", vals.Len(), idxs.Len()) + } + // Top 3: 9 (idx 5), 6 (idx 7), 5 (idx 4). + expectVals := []float64{9, 6, 5} + expectIdx := []int64{5, 7, 4} + for i := range 3 { + v, _ := FloatAt(vals, i) + j, _ := IntAt(idxs, i) + if v != expectVals[i] { + t.Errorf("TopK vals [%d]: got %v, want %v", i, v, expectVals[i]) + } + if j != expectIdx[i] { + t.Errorf("TopK idxs [%d]: got %d, want %d", i, j, expectIdx[i]) + } + } +} + +func TestTopK2D(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 5, 3, 4, + 9, 2, 7, 6, + }, 2, 4) + vals, idxs, err := TopK(a, 2, 1) + if err != nil { + t.Fatal(err) + } + if vals.Shape()[0] != 2 || vals.Shape()[1] != 2 { + t.Errorf("TopK2D shape: %v", vals.Shape()) + } + // Row 0 top-2: 5 (idx 1), 4 (idx 3). Row 1 top-2: 9 (idx 0), 7 (idx 2). + wantVals := []float64{5, 4, 9, 7} + wantIdx := []int64{1, 3, 0, 2} + for i := range 4 { + v, _ := FloatAt(vals, i/2, i%2) + j, _ := IntAt(idxs, i/2, i%2) + if v != wantVals[i] { + t.Errorf("TopK2D vals [%d]: got %v, want %v", i, v, wantVals[i]) + } + if j != wantIdx[i] { + t.Errorf("TopK2D idxs [%d]: got %d, want %d", i, j, wantIdx[i]) + } + } +} + +// TestTopK2DDim0 pins the non-trailing-dimension layout: the output +// keeps (k, suffix) row-major order, which the write offset used to +// transpose into (suffix, k). +func TestTopK2DDim0(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 4, + 3, 2, + 5, 0, + }, 3, 2) + vals, idxs, err := TopK(a, 2, 0) + if err != nil { + t.Fatal(err) + } + if vals.Shape()[0] != 2 || vals.Shape()[1] != 2 { + t.Errorf("TopK2DDim0 shape: %v", vals.Shape()) + } + // Column 0 top-2: 5 (idx 2), 3 (idx 1). Column 1 top-2: 4 (idx 0), 2 (idx 1). + wantVals := []float64{5, 4, 3, 2} + wantIdx := []int64{2, 0, 1, 1} + for i := range 4 { + v, _ := FloatAt(vals, i/2, i%2) + j, _ := IntAt(idxs, i/2, i%2) + if v != wantVals[i] { + t.Errorf("TopK2DDim0 vals [%d]: got %v, want %v", i, v, wantVals[i]) + } + if j != wantIdx[i] { + t.Errorf("TopK2DDim0 idxs [%d]: got %d, want %d", i, j, wantIdx[i]) + } + } +} + +func TestTopKErrors(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 3) + if _, _, err := TopK(a, 5, 0); err == nil { + t.Error("TopK: expected error for k > dim size") + } + if _, _, err := TopK(a, -1, 0); err == nil { + t.Error("TopK: expected error for negative k") + } +} + +func TestTopKNaN(t *testing.T) { + // NaN elements never rank: the finite values win and NaN fills the + // remaining slots of the reduced dimension when there are fewer than + // k finite values. + a := mustFromFloats(t, []float64{1, math.NaN(), 3, 2}, 4) + vals, idxs, err := TopK(a, 3, 0) + if err != nil { + t.Fatal(err) + } + // Three finite values exist, so all three slots are finite, in + // descending order. + wantVals := []float64{3, 2, 1} + wantIdx := []int64{2, 3, 0} + for i := range 3 { + v, _ := FloatAt(vals, i) + if v != wantVals[i] { + t.Errorf("TopKNaN vals [%d]: got %v, want %v", i, v, wantVals[i]) + } + j, _ := IntAt(idxs, i) + if j != wantIdx[i] { + t.Errorf("TopKNaN idxs [%d]: got %d, want %d", i, j, wantIdx[i]) + } + } + // With more NaN than the slots allow, NaN fills the tail: no finite + // candidate remains, so the slot reports NaN at the fill index 0. + b := mustFromFloats(t, []float64{1, math.NaN(), math.NaN()}, 3) + vb, ib, err := TopK(b, 2, 0) + if err != nil { + t.Fatal(err) + } + if v, _ := FloatAt(vb, 0); v != 1 { + t.Errorf("TopKNaN tail vals[0]: got %v, want 1", v) + } + if v, _ := FloatAt(vb, 1); !math.IsNaN(v) { + t.Errorf("TopKNaN tail vals[1]: got %v, want NaN", v) + } + if j, _ := IntAt(ib, 1); j != 0 { + t.Errorf("TopKNaN tail idxs[1]: got %d, want 0", j) + } +} diff --git a/internal/core/arrayutil.go b/internal/core/arrayutil.go new file mode 100644 index 0000000..e9ecb5d --- /dev/null +++ b/internal/core/arrayutil.go @@ -0,0 +1,2046 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +import ( + "math" + "slices" + "sync/atomic" +) + +// Array utilities, array creation, manipulation and conversion +// helpers: Linspace, +// Repeat, Tile, Flip, Roll, Unique, Argwhere, Astype, Item and Diag. + +// copyMinPerWorker is the per-worker chunk floor for the kernels whose +// per-element work is a load and a store: the mirror, gather, repeated +// block and payload-conversion copies. Their chunks move bytes, so a +// spawned worker needs a few thousand elements before the chunk +// outweighs its own start-up; below that floor the copy runs on the +// calling goroutine. +const copyMinPerWorker = 1 << 12 + +// Linspace returns n evenly spaced values from start to stop inclusive. +// n = 1 yields the single value start; n = 0 yields an empty array. +func Linspace(start, stop float64, n int) (*Array, error) { + if n < 0 { + return nil, errf("Linspace: n must be zero or greater, got %d", n) + } + if n == 0 { + return FromFloats(nil, 0) + } + if n == 1 { + return FromFloats([]float64{start}, 1) + } + h := (stop - start) / float64(n-1) + vals := make([]float64, n) + for i := range n { + vals[i] = start + float64(i)*h + } + // Pin the endpoints exactly; the loop can drift on awkward ranges. + vals[0], vals[n-1] = start, stop + return FromFloats(vals, n) +} + +// Repeat repeats each element of a `repeats` times along dim. +func Repeat(a *Array, repeats, dim int) (*Array, error) { + if repeats < 0 { + return nil, errf("Repeat: repeats must be zero or greater, got %d", repeats) + } + if dim < 0 || dim >= a.NDim() { + return nil, errf("Repeat: dimension %d out of range for shape %s", dim, shapeText(a.shape)) + } + // Bound the extent before multiplying and the product before + // allocating: hostile repeats must be an error, not a wrapped + // extent or a tiny allocation paired with a huge shape. + if a.shape[dim] != 0 && repeats != 0 && a.shape[dim] > math.MaxInt/repeats { + return nil, errf("Repeat: repeats %d overflows dimension %d of length %d", repeats, dim, a.shape[dim]) + } + newShape := a.Shape() + newShape[dim] *= repeats + total, _, terr := checkedDims(newShape) + if terr != nil { + return nil, base.WrapErr("Repeat", terr) + } + out := &Array{shape: newShape, dt: a.dt} + out.alloc(total) + // An empty result has nothing to fill: with repeats = 0 the extent + // collapses to zero while the row walk below would still visit the + // source rows and slice the empty destination payload, which + // panicked instead of answering the empty array (the repeat-zero + // test pins it). + if total == 0 { + return out, nil + } + src := a + if !src.isContiguous() { + src = src.materialise() + } + n := a.shape[dim] + outer, inner := 1, 1 + for d := range dim { + outer *= a.shape[d] + } + for d := dim + 1; d < a.NDim(); d++ { + inner *= a.shape[d] + } + if inner == 1 { + // The repeated axis is innermost, so result element i is source + // element i/repeats: one flat pass per dtype, no block walk. + repeatEach(out, src, repeats) + return out, nil + } + // Every output row of `inner` elements is a whole copy of one source + // row; the repeat factor only decides which destination rows share it. + for o := range outer { + srcRow := o * n * inner + dstRow := o * n * repeats * inner + for j := range n { + dst := dstRow + j*repeats*inner + copyRun(out, src, dst, srcRow+j*inner, inner) + for k := 1; k < repeats; k++ { + copyRun(out, out, dst+k*inner, dst, inner) + } + } + } + return out, nil +} + +// repeatEach fills dst with each element of src repeated r times, the +// mapping of a repeat along the innermost axis: dst[i] = src[i/r]. +// repeatEach fills dst with each element of src repeated r times, the +// mapping of a repeat along the innermost axis: dst[i] = src[i/r]. The +// dispatch carries every element type on its own payload. +func repeatEach(dst, src *Array, r int) { + switch dst.dt { + case Int: + repeatInto(dst.ints, src.ints, r) + case Float16: + repeatInto(dst.halves, src.halves, r) + case Float32: + repeatInto(dst.floats32, src.floats32, r) + case Float: + repeatInto(dst.floats, src.floats, r) + case Complex: + repeatInto(dst.complexes, src.complexes, r) + case Bool: + repeatInto(dst.bools, src.bools, r) + case Int8: + repeatInto(dst.i8s, src.i8s, r) + case Uint8: + repeatInto(dst.u8s, src.u8s, r) + case Int16: + repeatInto(dst.i16s, src.i16s, r) + case Uint16: + repeatInto(dst.u16s, src.u16s, r) + case Int32: + repeatInto(dst.i32s, src.i32s, r) + case Uint32: + repeatInto(dst.u32s, src.u32s, r) + } +} + +func repeatInto[T any](dst, src []T, r int) { + for i := range dst { + dst[i] = src[i/r] + } +} + +// Tile repeats a whole array reps times per dimension. The repetition +// vector is right-aligned to the shape, and a shorter shape prepends +// size-1 dimensions. The result never aliases a. +func Tile(a *Array, reps ...int) (*Array, error) { + if len(reps) == 0 { + return Copy(a), nil + } + ndim := max(a.NDim(), len(reps)) + shape := make([]int, ndim) + for d := range ndim { + ad := d - (ndim - a.NDim()) + rd := d - (ndim - len(reps)) + as := 1 + if ad >= 0 { + as = a.shape[ad] + } + r := 1 + if rd >= 0 { + r = reps[rd] + } + if r < 0 { + return nil, errf("Tile: negative repetition %d", r) + } + // Bound the extent before multiplying, exactly as Repeat does: + // a wrapped extent must be an error, not a negative allocation. + if as != 0 && r != 0 && as > math.MaxInt/r { + return nil, errf("Tile: repetition %d overflows a dimension of length %d", r, as) + } + shape[d] = as * r + } + total, _, terr := checkedDims(shape) + if terr != nil { + return nil, base.WrapErr("Tile", terr) + } + src := a + if !src.isContiguous() { + src = src.materialise() + } + // The source shape right-aligned to the result's rank; the leading + // size-1 dimensions the alignment adds do not move a flat index. + srcShape := make([]int, ndim) + for d := range ndim { + srcShape[d] = 1 + if ad := d - (ndim - a.NDim()); ad >= 0 { + srcShape[d] = a.shape[ad] + } + } + // rep[d] is how many times dimension d repeats, right-aligned like + // the shape. + rep := make([]int, ndim) + last := -1 + for d := range ndim { + rep[d] = 1 + if rd := d - (ndim - len(reps)); rd >= 0 { + rep[d] = reps[rd] + } + if rep[d] != 1 { + last = d + } + } + out := &Array{shape: shape, dt: a.dt} + out.alloc(total) + if last == ndim-1 && ndim > 0 { + // The innermost dimension repeats, so a block of the result is no + // longer one run of the source; the coordinate walk fills it. + coord := make([]int, ndim) + for i := range out.Len() { + off := 0 + for d := range ndim { + if srcShape[d] == 0 { + continue + } + off = off*srcShape[d] + coord[d]%srcShape[d] + } + out.setFrom(i, src, off) + advanceOdometer(coord, shape) + } + return out, nil + } + // Every dimension after the last repeat copies whole, so one output + // block of `run` elements is one contiguous run of the source. + run := 1 + for d := last + 1; d < ndim; d++ { + run *= srcShape[d] + } + coord := make([]int, last+1) + blocks := 1 + for d := 0; d <= last; d++ { + blocks *= shape[d] + } + for range blocks { + dstOff, srcOff := 0, 0 + for d := 0; d <= last; d++ { + dstOff = dstOff*shape[d] + coord[d] + srcOff = srcOff*srcShape[d] + coord[d]%srcShape[d] + } + copyRun(out, src, dstOff*run, srcOff*run, run) + advanceOdometer(coord, shape) + } + return out, nil +} + +// Flip reverses a along the given dimensions (all of them by default). +func Flip(a *Array, dims ...int) (*Array, error) { + if len(dims) == 0 { + dims = rangeN(a.NDim()) + } + flipSet := make([]bool, a.NDim()) + for _, d := range dims { + if d < 0 || d >= a.NDim() { + return nil, errf("Flip: dimension %d out of range for shape %s", d, shapeText(a.shape)) + } + flipSet[d] = true + } + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(a.Len()) + src := a + if !src.isContiguous() { + src = src.materialise() + } + // Every dimension after the last flipped one keeps its order, so one + // run of `run` elements is copied whole and only the run's position + // reflects the reversal. The prefix order is walked with an odometer + // over the flipped coordinate, which folds into the destination run + // index. + last := -1 + for _, d := range dims { + last = max(last, d) + } + run := 1 + for d := last + 1; d < a.NDim(); d++ { + run *= a.shape[d] + } + if run == 1 { + // Nothing past the last flipped dimension survives the reversal, + // so the flipped axis itself becomes the reversed run: a row of + // shape[last] elements mirrors whole and the prefix odometer + // only decides where the row lands. + n := a.shape[last] + head := 1 + for d := range last { + head *= a.shape[d] + } + row := make([]int, last) + for range head { + dstRow, srcRow := 0, 0 + for d := range last { + c := row[d] + dst := c + if flipSet[d] { + dst = a.shape[d] - 1 - c + } + dstRow = dstRow*a.shape[d] + dst + srcRow = srcRow*a.shape[d] + c + } + reverseRows(out, src, dstRow*n, srcRow*n, n) + advanceOdometer(row, a.shape) + } + return out, nil + } + // Every dimension after the last flipped one keeps its order, so one + // run of `run` elements moves whole and only the run's position + // reflects the reversal. + coord := make([]int, last+1) + blocks := 1 + for d := 0; d <= last; d++ { + blocks *= a.shape[d] + } + for range blocks { + dstOff := 0 + for d := 0; d <= last; d++ { + c := coord[d] + if flipSet[d] { + c = a.shape[d] - 1 - c + } + dstOff = dstOff*a.shape[d] + c + } + copyRun(out, src, dstOff*run, blockFlat(coord, a.shape)*run, run) + advanceOdometer(coord, a.shape) + } + return out, nil +} + +// reverseRows mirrors the run of src at srcOff into dst at dstOff, one +// row of a flip: the serial form the per-row walk needs, since spawning +// workers for every row would cost more than the row. The dispatch +// carries every element type on its own payload. +func reverseRows(dst, src *Array, dstOff, srcOff, run int) { + switch dst.dt { + case Int: + reverseIntoSerial(dst.ints[dstOff:dstOff+run], src.ints[srcOff:srcOff+run]) + case Float16: + reverseIntoSerial(dst.halves[dstOff:dstOff+run], src.halves[srcOff:srcOff+run]) + case Float32: + reverseIntoSerial(dst.floats32[dstOff:dstOff+run], src.floats32[srcOff:srcOff+run]) + case Float: + reverseIntoSerial(dst.floats[dstOff:dstOff+run], src.floats[srcOff:srcOff+run]) + case Complex: + reverseIntoSerial(dst.complexes[dstOff:dstOff+run], src.complexes[srcOff:srcOff+run]) + case Bool: + reverseIntoSerial(dst.bools[dstOff:dstOff+run], src.bools[srcOff:srcOff+run]) + case Int8: + reverseIntoSerial(dst.i8s[dstOff:dstOff+run], src.i8s[srcOff:srcOff+run]) + case Uint8: + reverseIntoSerial(dst.u8s[dstOff:dstOff+run], src.u8s[srcOff:srcOff+run]) + case Int16: + reverseIntoSerial(dst.i16s[dstOff:dstOff+run], src.i16s[srcOff:srcOff+run]) + case Uint16: + reverseIntoSerial(dst.u16s[dstOff:dstOff+run], src.u16s[srcOff:srcOff+run]) + case Int32: + reverseIntoSerial(dst.i32s[dstOff:dstOff+run], src.i32s[srcOff:srcOff+run]) + case Uint32: + reverseIntoSerial(dst.u32s[dstOff:dstOff+run], src.u32s[srcOff:srcOff+run]) + } +} + +// blockFlat folds a full prefix of coordinates into its row-major flat +// index. +func blockFlat(coord, shape []int) int { + off := 0 + for d, c := range coord { + off = off*shape[d] + c + } + return off +} + +// Roll shifts a along dim by shift; values wrap around and a negative +// shift moves elements backward. +func Roll(a *Array, shift, dim int) (*Array, error) { + if dim < 0 || dim >= a.NDim() { + return nil, errf("Roll: dimension %d out of range for shape %s", dim, shapeText(a.shape)) + } + n := a.shape[dim] + if n == 0 { + return Copy(a), nil + } + shift %= n + if shift < 0 { + shift += n + } + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(a.Len()) + src := a + if !src.isContiguous() { + src = src.materialise() + } + // out[c] = a[(c − shift) mod n] along the axis, so one outer block + // rotates as two runs: the axis tail after the shift, then the head. + // Every position off the axis keeps its order, so each run of `tail` + // elements moves whole. + tail := 1 + for d := dim + 1; d < a.NDim(); d++ { + tail *= a.shape[d] + } + outer := 1 + for d := range dim { + outer *= a.shape[d] + } + split := (n - shift) * tail + for o := range outer { + base := o * n * tail + copyRun(out, src, base+shift*tail, base, split) + copyRun(out, src, base, base+split, n*tail-split) + } + return out, nil +} + +// Unique returns the sorted unique values of a (real arrays only); NaN +// counts once. The ordering rule is the same for every dtype, the one +// the Int branch has always read off Sort: values ascend in the dtype's +// natural order, equal adjacent values collapse to the first, and the +// output keeps the input dtype. The narrow integer class sorts in exact +// int64 space, where that rule is the Int rule verbatim; bool orders +// false before true. The float dtypes keep the Sort-backed walk, with +// NaN counted once through sameFloat. +func Unique(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("Unique: complex arrays have no ordering") + } + switch a.dt { + case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32: + // Sort carries no narrow payload dispatch, so the narrow class + // walks its own exact int64 widening: ascending order and the + // adjacent dedupe are exactly what the Int branch reads off + // Sort, and the result rebuilds in the source dtype. + n := a.Len() + vals := make([]int64, n) + for i := range n { + vals[i] = a.intAt(i) + } + slices.Sort(vals) + uniq := make([]int64, 0, len(vals)) + for i, v := range vals { + if i == 0 || v != vals[i-1] { + uniq = append(uniq, v) + } + } + switch a.dt { + case Bool: + out := make([]bool, len(uniq)) + for i, v := range uniq { + out[i] = v != 0 + } + return FromBools(out, len(out)) + case Int8: + out := make([]int8, len(uniq)) + for i, v := range uniq { + out[i] = int8(v) + } + return FromInt8s(out, len(out)) + case Uint8: + out := make([]uint8, len(uniq)) + for i, v := range uniq { + out[i] = uint8(v) + } + return FromUint8s(out, len(out)) + case Int16: + out := make([]int16, len(uniq)) + for i, v := range uniq { + out[i] = int16(v) + } + return FromInt16s(out, len(out)) + case Uint16: + out := make([]uint16, len(uniq)) + for i, v := range uniq { + out[i] = uint16(v) + } + return FromUint16s(out, len(out)) + case Int32: + out := make([]int32, len(uniq)) + for i, v := range uniq { + out[i] = int32(v) + } + return FromInt32s(out, len(out)) + default: + out := make([]uint32, len(uniq)) + for i, v := range uniq { + out[i] = uint32(v) + } + return FromUint32s(out, len(out)) + } + } + sorted, err := Sort(a) + if err != nil { + return nil, err + } + switch a.dt { + case Int: + out := make([]int64, 0, len(sorted.ints)) + for i, v := range sorted.ints { + if i == 0 || v != sorted.ints[i-1] { + out = append(out, v) + } + } + return FromInts(out, len(out)) + case Float16: + out := make([]uint16, 0, len(sorted.halves)) + for i, v := range sorted.halves { + if i == 0 || !sameFloat(HalfToFloat64(v), HalfToFloat64(sorted.halves[i-1])) { + out = append(out, v) + } + } + return HalvesFromArray(out, len(out)) + case Float32: + out := make([]float32, 0, len(sorted.floats32)) + for i, v := range sorted.floats32 { + if i == 0 || !sameFloat(float64(v), float64(sorted.floats32[i-1])) { + out = append(out, v) + } + } + return FromFloat32s(out, len(out)) + default: + out := make([]float64, 0, len(sorted.floats)) + for i, v := range sorted.floats { + if i == 0 || !sameFloat(v, sorted.floats[i-1]) { + out = append(out, v) + } + } + return FromFloats(out, len(out)) + } +} + +// sameFloat reports equality with NaN treated as equal to NaN: the +// dedupe rule for Unique. +func sameFloat(x, y float64) bool { + return x == y || (x != x && y != y) +} + +// argwhereChunkBufMin is the first capacity of a chunk's coordinate +// buffer in Argwhere: large enough that a sparse chunk grows a handful +// of times, and capped by what the chunk could ever hold, so a small +// array pays a single exact allocation. It never guesses density: +// growth past it doubles, and the merge copies each chunk's coordinates +// into the exact output either way. +const argwhereChunkBufMin = 1024 + +// Argwhere returns the coordinates of every non-zero element as an +// (nnz, ndim) int array of non-zero coordinates. +func Argwhere(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("Argwhere: complex arrays have no notion of zero") + } + n := a.Len() + ndim := a.NDim() + chunk := n + chunks := 1 + if n >= copyMinPerWorker { + w := workersFor(n) + chunk = (n + w - 1) / w + chunks = (n + chunk - 1) / chunk + } + // One input pass: every chunk walks its own range once, with its own + // odometer, appending its coordinates to a chunk-local buffer, so + // the parts concatenate into exactly the row-major order one walk + // produced. The merge then copies them in chunk order into the exact + // output, which is fully overwritten before anyone reads it. + parts := make([][]int64, chunks) + parallelMin(chunks, 1, func(s, e int) { + for c := s; c < e; c++ { + start := c * chunk + end := min(start+chunk, n) + buf := make([]int64, 0, min(argwhereChunkBufMin, (end-start)*ndim)) + parts[c] = argwhereAppend(a, buf, start, end) + } + }) + total := 0 + for _, p := range parts { + total += len(p) + } + // Each part carries ndim values per non-zero element, so the row + // count is the value count over ndim. + out := &Array{shape: []int{total / ndim, ndim}, dt: Int} + out.alloc(total) + off := 0 + for _, p := range parts { + copy(out.ints[off:off+len(p)], p) + off += len(p) + } + return out, nil +} + +// argwhereAppend walks [start, end) with the chunk's own odometer and +// appends the coordinates of every non-zero element to buf, in +// row-major order. The seeding of the odometer from the flat start is +// what lets each chunk of a split walk produce its own part in order. +// The zero test is dispatched on the dtype once, and the loops are +// bounded by the array's own extent, never by a payload length: a +// rebased view carries a payload longer than its own Len. +func argwhereAppend(a *Array, buf []int64, start, end int) []int64 { + ndim := a.NDim() + coord := make([]int, ndim) + if start > 0 { + rest := start + for d := ndim - 1; d >= 0; d-- { + coord[d] = rest % a.shape[d] + rest /= a.shape[d] + } + } + appendCoord := func() { + for d := range ndim { + buf = append(buf, int64(coord[d])) + } + } + if !a.isContiguous() { + for i := start; i < end; i++ { + if !isZero(a, i) { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + return buf + } + switch a.dt { + case Int: + p := a.ints + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Float16: + p := a.halves + for i := start; i < end; i++ { + if p[i]&0x7FFF != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Float32: + p := a.floats32 + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Float: + p := a.floats + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Bool: + // The mask reading: a true element is the non-zero one. + p := a.bools + for i := start; i < end; i++ { + if p[i] { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Int8: + p := a.i8s + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Uint8: + p := a.u8s + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Int16: + p := a.i16s + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Uint16: + p := a.u16s + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Int32: + p := a.i32s + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Uint32: + p := a.u32s + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + default: + // Complex: Argwhere's entry gate rejects complex arrays before + // the walk, so the arm exists for dispatch completeness; the + // value test matches the one the float arms carry. + p := a.complexes + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + } + return buf +} + +// Astype returns a copy of a converted to dt. The split is deliberate: +// the legacy conversions keep their historical cast semantics, the new +// narrow targets check their range. int to float and float to int +// convert like Go's casts; a float destination rounds through float64, +// so a float16 or float32 target narrows with HalfFromFloat64 and +// float32(v); complex to float keeps the real part, complex to int and +// complex to float16 are errors; real to complex adds a zero imaginary +// part. Narrowing into bool, int8, uint8, int16, uint16, int32 or +// uint32 is range-checked in the source's value space, and the first +// element the target cannot represent fails the call loudly: integer +// sources check exactly in int64; a float source must be finite, +// integral and inside the target's range, so a NaN, an infinity or a +// fraction is an error rather than a silent cast. Bool is the one +// target with no range to check: every source reaches it by the test +// against zero (a NaN reads true), and bool converts out as 0/1, exact +// everywhere including complex. Converting to the array's own dtype +// copies it unchanged, which the int route cannot express as a cast: +// float64 rounds above 2^53. +func Astype(a *Array, dt Dtype) (*Array, error) { + if a.dt == dt { + // Same dtype: a straight copy through cloneArray, fast and + // exact for every element type; the float64 detour a numeric + // route would take rounds above 2^53. + return a.cloneArray(), nil + } + out := &Array{shape: a.Shape(), dt: dt} + out.alloc(a.Len()) + if dt == Float && a.dt == Complex { + // Complex to float keeps the real part, as documented. + parallelMin(a.Len(), copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.floats[i] = real(a.complexes[i]) + } + }) + return out, nil + } + if dt == Bool { + // Every source reaches bool by the test against zero, through + // the widening reader that resolves a view's strides; a NaN + // compares unequal to zero and reads true. No value can fail + // this conversion, so the walk cannot error. + n := a.Len() + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = a.boolAt(i) + } + }) + return out, nil + } + if a.dt == Complex { + // Every remaining target sits below complex on the ladder and + // has no defined complex value space, the narrow numeric + // targets included: the historical loud refusal. + return nil, errf("Astype: cannot narrow complex to %s", dt) + } + src := a + if !src.isContiguous() { + // A strided source has no payload run to read; reduce it to a + // dense copy through the same accessor the walk used. + src = src.materialise() + } + // The source dispatch sits outside the element loop, so each + // destination runs one monomorphised conversion whose arithmetic is + // the widening floatAt performed followed by the destination's own + // cast, bit for bit. + n := src.Len() + switch src.dt { + case Int: + p := src.ints[:n] + switch dt { + case Float16: + d := out.halves + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = HalfFromFloat64(float64(p[i])) + } + }) + case Float32: + d := out.floats32 + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = float32(float64(p[i])) + } + }) + case Float: + d := out.floats + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = float64(p[i]) + } + }) + case Int8, Uint8, Int16, Uint16, Int32, Uint32: + // Narrowing into a new narrow target: range-checked in the + // source's exact int64 value space. + if err := astypeNarrowFromInt(out, p); err != nil { + return nil, err + } + return out, nil + default: + d := out.complexes + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = complex(float64(p[i]), 0) + } + }) + } + case Float16: + p := src.halves[:n] + switch dt { + case Int: + d := out.ints + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = int64(HalfToFloat64(p[i])) + } + }) + case Float32: + d := out.floats32 + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = float32(HalfToFloat64(p[i])) + } + }) + case Float: + d := out.floats + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = HalfToFloat64(p[i]) + } + }) + case Int8, Uint8, Int16, Uint16, Int32, Uint32: + // Range-checked in the source's float64 value space: the + // half widens exactly, then the float rule applies. + if err := astypeNarrowFromFloat(out, floatPayload(src)); err != nil { + return nil, err + } + return out, nil + default: + d := out.complexes + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = complex(HalfToFloat64(p[i]), 0) + } + }) + } + case Float32: + p := src.floats32[:n] + switch dt { + case Int: + d := out.ints + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = int64(float64(p[i])) + } + }) + case Float16: + d := out.halves + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = HalfFromFloat64(float64(p[i])) + } + }) + case Float: + d := out.floats + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = float64(p[i]) + } + }) + case Int8, Uint8, Int16, Uint16, Int32, Uint32: + // Range-checked in the source's float64 value space: the + // float32 widens exactly, then the float rule applies. + if err := astypeNarrowFromFloat(out, floatPayload(src)); err != nil { + return nil, err + } + return out, nil + default: + d := out.complexes + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = complex(float64(p[i]), 0) + } + }) + } + case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32: + // The narrow sources widen exactly into int64, the value space + // their conversions and range checks work in; bool reads 0/1. + vals := make([]int64, n) + for i := range n { + vals[i] = src.intAt(i) + } + switch dt { + case Int: + d := out.ints + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = vals[i] + } + }) + case Float16: + d := out.halves + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + // The exact float64 widening followed by the + // destination's own nearest-narrowing cast; only the + // bool, int8 and uint8 value ranges stay exact in + // every half, wider integers round above 2048. + d[i] = HalfFromFloat64(float64(vals[i])) + } + }) + case Float32: + d := out.floats32 + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + // Exact through float64; float32 rounds above 2^24, + // which reaches the int32 and uint32 sources. + d[i] = float32(float64(vals[i])) + } + }) + case Float: + d := out.floats + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + // Exact: every narrow source value fits float64. + d[i] = float64(vals[i]) + } + }) + case Complex: + d := out.complexes + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + // Exact: the element's exact float64 value with a + // zero imaginary part. + d[i] = complex(float64(vals[i]), 0) + } + }) + case Int8, Uint8, Int16, Uint16, Int32, Uint32: + // Narrow to narrow: exact where the target contains the + // source's value set, a loud range error where it does not + // (int8 to uint8 of a negative value, say). + if err := astypeNarrowFromInt(out, vals); err != nil { + return nil, err + } + return out, nil + } + default: + // Float, the one source still unnamed. + p := src.floats[:n] + switch dt { + case Int: + d := out.ints + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = int64(p[i]) + } + }) + case Float16: + d := out.halves + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = HalfFromFloat64(p[i]) + } + }) + case Float32: + d := out.floats32 + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = float32(p[i]) + } + }) + case Int8, Uint8, Int16, Uint16, Int32, Uint32: + // Range-checked in the source's float64 value space. + if err := astypeNarrowFromFloat(out, p); err != nil { + return nil, err + } + return out, nil + default: + d := out.complexes + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + d[i] = complex(p[i], 0) + } + }) + } + } + return out, nil +} + +// astypeNarrowFromInt writes vals, the source values in exact int64 +// space, into out's narrow integer payload. Representability is checked +// in that space, and the lowest index whose value the target cannot +// hold fails the call with the Astype range error. +func astypeNarrowFromInt(out *Array, vals []int64) error { + n := len(vals) + switch out.dt { + case Int8: + return narrowWrite(out.i8s[:n], vals, math.MinInt8, math.MaxInt8, out.dt) + case Uint8: + return narrowWrite(out.u8s[:n], vals, 0, math.MaxUint8, out.dt) + case Int16: + return narrowWrite(out.i16s[:n], vals, math.MinInt16, math.MaxInt16, out.dt) + case Uint16: + return narrowWrite(out.u16s[:n], vals, 0, math.MaxUint16, out.dt) + case Int32: + return narrowWrite(out.i32s[:n], vals, math.MinInt32, math.MaxInt32, out.dt) + default: + return narrowWrite(out.u32s[:n], vals, 0, math.MaxUint32, out.dt) + } +} + +// astypeNarrowFromFloat is astypeNarrowFromInt for float sources, whose +// value space is float64. A representable element is finite, integral +// and inside the target's range; a NaN, an infinity or a fraction fails +// the call instead of casting silently, and the error names the value +// in the source's own float64 space. +func astypeNarrowFromFloat(out *Array, vals []float64) error { + n := len(vals) + switch out.dt { + case Int8: + return narrowWriteF(out.i8s[:n], vals, math.MinInt8, math.MaxInt8, out.dt) + case Uint8: + return narrowWriteF(out.u8s[:n], vals, 0, math.MaxUint8, out.dt) + case Int16: + return narrowWriteF(out.i16s[:n], vals, math.MinInt16, math.MaxInt16, out.dt) + case Uint16: + return narrowWriteF(out.u16s[:n], vals, 0, math.MaxUint16, out.dt) + case Int32: + return narrowWriteF(out.i32s[:n], vals, math.MinInt32, math.MaxInt32, out.dt) + default: + return narrowWriteF(out.u32s[:n], vals, 0, math.MaxUint32, out.dt) + } +} + +// narrowWrite stores vals in dst, failing the call when a value falls +// outside [lo, hi]. The walk parallelises like every conversion kernel, +// so the failure published is the lowest failing index: the first +// element in row-major order, whichever worker sees it. The stored cast +// is exact once the range holds. +func narrowWrite[T int8 | uint8 | int16 | uint16 | int32 | uint32](dst []T, vals []int64, lo, hi int64, dt Dtype) error { + n := len(vals) + var bad atomic.Int64 + bad.Store(int64(n)) + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + v := vals[i] + if v < lo || v > hi { + for { + cur := bad.Load() + if int64(i) >= cur || bad.CompareAndSwap(cur, int64(i)) { + break + } + } + continue + } + dst[i] = T(v) + } + }) + if idx := int(bad.Load()); idx < n { + return errf("Astype: value %v at index %d does not fit %s", vals[idx], idx, dt) + } + return nil +} + +// narrowWriteF is narrowWrite for float64 source values: representable +// means finite, integral and inside [lo, hi], and the stored cast is +// exact once those hold. +func narrowWriteF[T int8 | uint8 | int16 | uint16 | int32 | uint32](dst []T, vals []float64, lo, hi float64, dt Dtype) error { + n := len(vals) + var bad atomic.Int64 + bad.Store(int64(n)) + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + v := vals[i] + if math.IsNaN(v) || math.IsInf(v, 0) || math.Trunc(v) != v || v < lo || v > hi { + for { + cur := bad.Load() + if int64(i) >= cur || bad.CompareAndSwap(cur, int64(i)) { + break + } + } + continue + } + dst[i] = T(v) + } + }) + if idx := int(bad.Load()); idx < n { + return errf("Astype: value %v at index %d does not fit %s", vals[idx], idx, dt) + } + return nil +} + +// astypeArray is Astype for the autograd engine; the float and +// float32 conversions it needs never error and return the input +// unchanged when the dtype already matches. +func astypeArray(a *Array, dt Dtype) (*Array, error) { + if a.dt == dt { + return a, nil + } + return Astype(a, dt) +} + +// Item returns the single element of a 1-element array as a float64 +// (the real part for complex arrays). +func Item(a *Array) (float64, error) { + if a.Len() != 1 { + return 0, errf("Item: needs a 1-element array, got shape %s", shapeText(a.shape)) + } + if a.dt == Complex { + return real(a.complexes[0]), nil + } + return a.floatAt(0), nil +} + +// Diag extracts the main diagonal of a 2-D array, or builds a diagonal +// matrix from a 1-D array. +func Diag(a *Array) (*Array, error) { + switch a.NDim() { + case 1: + n := a.Len() + out, err := Zeros(a.dt, n, n) + if err != nil { + return nil, err + } + for i := range n { + out.setFrom(i*n+i, a, i) + } + return out, nil + case 2: + return Diagonal(a, 0) + } + return nil, errf("Diag: needs a 1-D or 2-D array, got shape %s", shapeText(a.shape)) +} + +// All reports whether every element is non-zero (real arrays only; the +// mask semantics match Select: any non-zero value counts as true). +func All(a *Array) (bool, error) { + if a.dt == Complex { + return false, errf("All: complex arrays have no notion of zero") + } + return zeroScan(a, true, 1) == 0, nil +} + +// Any reports whether at least one element is non-zero. +func Any(a *Array) (bool, error) { + if a.dt == Complex { + return false, errf("Any: complex arrays have no notion of zero") + } + return zeroScan(a, false, 1) > 0, nil +} + +// CountNonzero returns the number of non-zero elements. +func CountNonzero(a *Array) (int, error) { + if a.dt == Complex { + return 0, errf("CountNonzero: complex arrays have no notion of zero") + } + return zeroScan(a, false, 0), nil +} + +// zeroScan counts the elements of a that test as countZeros wants, +// stopping once it has seen stop of them (stop 0 counts them all). The +// zero test is dispatched on the dtype once rather than per element, +// and the walk stops as early as the old per-element loop did. +func zeroScan(a *Array, countZeros bool, stop int) int { + n := a.Len() + count := 0 + hit := func() bool { + count++ + return stop > 0 && count >= stop + } + if !a.isContiguous() { + for i := range n { + if (a.floatAt(i) == 0) == countZeros && hit() { + return count + } + } + return count + } + switch a.dt { + case Int: + for _, v := range a.ints[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Float16: + // The value test in bit space: clearing the sign bit folds -0.0 + // onto +0.0 exactly as the float comparisons do, and a half's + // non-zero patterns all widen to a non-zero double. + for _, v := range a.halves[:n] { + if (v&0x7FFF == 0) == countZeros && hit() { + return count + } + } + case Float32: + for _, v := range a.floats32[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Float: + for _, v := range a.floats[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Bool: + // The mask reading: false is the zero element, true the + // non-zero one, the same test isZero carries. + for _, v := range a.bools[:n] { + if (!v) == countZeros && hit() { + return count + } + } + case Int8: + for _, v := range a.i8s[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Uint8: + for _, v := range a.u8s[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Int16: + for _, v := range a.i16s[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Uint16: + for _, v := range a.u16s[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Int32: + for _, v := range a.i32s[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + case Uint32: + for _, v := range a.u32s[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + default: + // Complex: All, Any and CountNonzero reject complex arrays at + // entry, so the arm exists for dispatch completeness; the value + // test matches the one the float arms carry. + for _, v := range a.complexes[:n] { + if (v == 0) == countZeros && hit() { + return count + } + } + } + return count +} + +// zeros allocates a zeroed array, treating the constructor's error as +// unreachable for internally derived shapes. +func zeros(dt Dtype, shape []int) *Array { + a, _ := Zeros(dt, shape...) + return a +} + +// Grid builds coordinate matrices from two 1-D vectors: the meshgrid +// pattern for evaluating functions over a 2-D domain. X varies along +// the columns and Y along the rows. +func Grid(a, b *Array) (xGrid, yGrid *Array, err error) { + if a.NDim() != 1 || b.NDim() != 1 { + return nil, nil, errf("Grid: needs two 1-D arrays, got %s and %s", shapeText(a.shape), shapeText(b.shape)) + } + if a.dt == Complex || b.dt == Complex { + return nil, nil, errf("Grid: complex axes are not supported") + } + na, nb := a.Len(), b.Len() + xGrid = &Array{shape: []int{nb, na}, dt: Float} + yGrid = &Array{shape: []int{nb, na}, dt: Float} + xGrid.alloc(nb * na) + yGrid.alloc(nb * na) + // The x row is one copy per row, the y row one fill; the widening of + // the two 1-D sources is the exact one FloatAt reports. + av, bv := floatPayload(a), floatPayload(b) + parallelMin(nb, copyMinPerWorker, func(s, e int) { + for r := s; r < e; r++ { + copy(xGrid.floats[r*na:(r+1)*na], av) + yr := bv[r] + for c := range na { + yGrid.floats[r*na+c] = yr + } + } + }) + return xGrid, yGrid, nil +} + +// CrossProduct computes the vector cross product of two length-3 +// vectors. +func CrossProduct(u, v *Array) (*Array, error) { + if u.dt == Complex || v.dt == Complex { + return nil, errf("CrossProduct: complex vectors are not supported") + } + if u.Len() != 3 || v.Len() != 3 { + return nil, errf("CrossProduct: needs two length-3 vectors, got %d and %d", u.Len(), v.Len()) + } + out := zeros(Float, []int{3}) + u1, u2, u3 := u.FloatAt(0), u.FloatAt(1), u.FloatAt(2) + v1, v2, v3 := v.FloatAt(0), v.FloatAt(1), v.FloatAt(2) + out.SetFloatAt(0, u2*v3-u3*v2) + out.SetFloatAt(1, u3*v1-u1*v3) + out.SetFloatAt(2, u1*v2-u2*v1) + return out, nil +} + +// Integrate computes the definite integral of y over uniform spacing +// dx using the trapezoidal rule. The areas fold through the canonical +// partition, the accuracy the central-sum test measured against the +// exact referent; at and below one block the walk is the plain chain, +// bit for bit. +func Integrate(y *Array, dx float64) (float64, error) { + if y.dt == Complex { + return 0, errf("Integrate: complex samples are not supported") + } + n := y.Len() + if n < 2 { + return 0, errf("Integrate: needs at least two samples, got %d", n) + } + // The widening happens once, outside the accumulation: the addends + // and their order are the ones the accessor loop summed. + yv := floatPayload(y) + return integrateAreas(yv) * dx, nil +} + +// integrateAreas sums the y samples' trapezoid areas through the +// canonical partition: the areas (one fewer than the samples) cut into +// the fixed blocks, one chain partial per block and the balanced tree +// over them. Area i reads samples i and i+1, the arithmetic the plain +// chain kept. +func integrateAreas(y []float64) float64 { + m := len(y) - 1 + parts := foldParts(m) + if parts == 1 { + var total float64 + for i := range m { + total += (y[i] + y[i+1]) / 2 + } + return total + } + partials := make([]float64, parts) + for c := range parts { + lo, hi := c*m/parts, (c+1)*m/parts + var acc float64 + for i := lo; i < hi; i++ { + acc += (y[i] + y[i+1]) / 2 + } + partials[c] = acc + } + return treeSum(partials) +} + +// CumulativeIntegrate returns the running trapezoidal integral of y +// over uniform spacing dx; the first element is zero. +func CumulativeIntegrate(y *Array, dx float64) (*Array, error) { + if y.dt == Complex { + return nil, errf("CumulativeIntegrate: complex samples are not supported") + } + n := y.Len() + yv := floatPayload(y) + out := zeros(Float, []int{n}) + ov := out.floats + for i := 1; i < n; i++ { + ov[i] = ov[i-1] + (yv[i-1]+yv[i])/2*dx + } + return out, nil +} + +// Interpolate evaluates the piecewise-linear interpolation of the +// points (xs[i], ys[i]) at each query position; queries outside the +// range clamp to the boundary values. xs need only be non-decreasing: +// a repeated knot gives a zero-width segment the walk skips, and a +// query landing exactly on it takes the left segment's upper end, so +// it returns the first of the repeated ys. Non-finite knots are +// refused, and so is a NaN query, which has no position to clamp to. +// The segment is located by searching the knots for the first one at +// or above the query, which is the segment the ascending walk stopped +// on; the per-point arithmetic is unchanged. +func Interpolate(xs, ys *Array, query *Array) (*Array, error) { + if xs.dt == Complex || ys.dt == Complex || query.dt == Complex { + return nil, errf("Interpolate: complex samples are not supported") + } + if xs.Len() != ys.Len() || xs.Len() < 2 { + return nil, errf("Interpolate: xs/ys must share length ≥ 2") + } + for k := range xs.Len() { + xk := xs.FloatAt(k) + if math.IsNaN(xk) || math.IsInf(xk, 0) { + return nil, errf("Interpolate: knot %d is not finite", k) + } + } + out := &Array{shape: query.Shape(), dt: Float} + out.alloc(query.Len()) + n := xs.Len() + xv, yv, qv := floatPayload(xs), floatPayload(ys), floatPayload(query) + // The workers write disjoint output slots. A NaN query has no + // segment and aborts the call; the serial walk reported the first + // one in order, so the workers publish the smallest NaN index. + var nanIdx atomic.Int64 + nanIdx.Store(int64(len(qv)) + 1) + parallelMin(len(qv), copyMinPerWorker, func(s, e int) { + ov := out.floats + for i := s; i < e; i++ { + q := qv[i] + if math.IsNaN(q) { + for { + cur := nanIdx.Load() + if int64(i) >= cur || nanIdx.CompareAndSwap(cur, int64(i)) { + break + } + } + continue + } + // The segment is located by bisecting the knots: the first + // knot at or above the query, which is the segment the + // ascending walk stopped on. + var lo int + if q <= xv[0] { + lo = 0 + } else if q >= xv[n-1] { + lo = n - 2 + } else { + lo0, hi0 := 1, n + for lo0 < hi0 { + mid := int(uint(lo0+hi0) >> 1) + if xv[mid] < q { + lo0 = mid + 1 + } else { + hi0 = mid + } + } + lo = lo0 - 1 + } + hi := lo + 1 + x0, x1 := xv[lo], xv[hi] + y0, y1 := yv[lo], yv[hi] + t := 0.0 + if x1 > x0 { + t = (q - x0) / (x1 - x0) + } else if q > x0 { + // A repeated knot at the range edge: the query sits past + // the zero-width segment, so it takes the upper end. + t = 1 + } + if t < 0 { + t = 0 + } + if t > 1 { + t = 1 + } + // The result is float by construction, so the payload write needs + // no dtype dispatch. + ov[i] = y0 + t*(y1-y0) + } + }) + if idx := int(nanIdx.Load()); idx <= len(qv) { + return nil, errf("Interpolate: query %d is NaN, which cannot be clamped", idx) + } + return out, nil +} + +// EvaluatePolynomial evaluates coefficients (lowest power first) at +// the given points. Both operands are widened once and read off the +// payload slices, the same values the accessor calls returned without +// the per-element dispatch; the accumulation and the power walk are +// the ones the accessor loop kept. +func EvaluatePolynomial(coeffs, x *Array) (*Array, error) { + if coeffs.dt == Complex || x.dt == Complex { + return nil, errf("EvaluatePolynomial: complex inputs are not supported") + } + out := &Array{shape: x.Shape(), dt: Float} + out.alloc(x.Len()) + cv, xv := floatPayload(coeffs), floatPayload(x) + for i := range xv { + pow := 1.0 + var sum float64 + for c := range cv { + sum += cv[c] * pow + pow *= xv[i] + } + out.floats[i] = sum + } + return out, nil +} + +// MoveAxis moves an axis to a new position in the shape. +func MoveAxis(a *Array, from, to int) (*Array, error) { + if from < 0 || from >= a.NDim() || to < 0 || to >= a.NDim() { + return nil, errf("MoveAxis: axes %d to %d out of range for rank %d", from, to, a.NDim()) + } + // The destination order: the remaining axes in sequence with from + // reinserted at to. + order := rangeN(a.NDim()) + order = append(order[:from], order[from+1:]...) + order = append(order[:to], append([]int{from}, order[to:]...)...) + newShape := make([]int, a.NDim()) + for k, d := range order { + newShape[k] = a.shape[d] + } + out := &Array{shape: newShape, dt: a.dt} + out.alloc(prodShape(newShape)) + coord := make([]int, a.NDim()) + srcCoord := make([]int, a.NDim()) + for i := range a.Len() { + for d := range a.NDim() { + srcCoord[order[d]] = coord[d] + } + off := 0 + for d := range a.NDim() { + off = off*a.shape[d] + srcCoord[d] + } + out.setFrom(i, a, off) + advanceOdometer(coord, newShape) + } + return out, nil +} + +// floatPayload returns a's elements as a plain float64 slice: the +// payload itself for a contiguous float64 array, and otherwise a dense +// copy through the array's own accessor. Every value is exactly the one +// floatAt reports, so a strided view or a narrower dtype reads the same +// numbers; the searches that follow then index a slice instead of +// paying a dtype dispatch per probed element. +func floatPayload(a *Array) []float64 { + n := a.Len() + if a.dt == Float && a.isContiguous() { + return a.floats[:n] + } + out := make([]float64, n) + for i := range n { + out[i] = a.floatAt(i) + } + return out +} + +// intPayload is floatPayload for the integer-class operands, whose +// comparisons must stay in int64: a contiguous int64 array aliases its +// payload, the fast path the int comparisons have always taken, and +// every other walk copies through intAt, which widens the whole integer +// class exactly (bool reads 0/1) and resolves a view's strides. +func intPayload(a *Array) []int64 { + n := a.Len() + if a.dt == Int && a.isContiguous() { + return a.ints[:n] + } + out := make([]int64, n) + for i := range n { + out[i] = a.intAt(i) + } + return out +} + +// searchMinPerWorker is the per-worker needle floor for SearchSorted: +// a needle costs a bisection whose depth grows with the haystack, tens +// of nanoseconds at the sizes the search benchmarks reach, so a chunk of +// a few hundred needles already outweighs its worker's start-up while +// smaller chunks run on the calling goroutine. +const searchMinPerWorker = 512 + +// SearchSorted finds insertion positions for each needle in the sorted +// haystack so that order is preserved (rightmost rule: positions after +// equal elements). The haystack must be ascending; the position is the +// number of elements at or below the needle, found by bisecting the +// haystack rather than by walking the prefix, which selects the same +// index in a logarithmic number of comparisons. The whole integer class +// bisects its own payload natively in int64, where the widening of +// bool and every narrow integer is exact; float operands widen through +// floatAt, exact for every real dtype, because float64 would round +// int64 values onto one another above 2^53. The workers write disjoint +// output slots, so the split cannot move a single position. +func SearchSorted(haystack, needles *Array) (*Array, error) { + if haystack.dt == Complex || needles.dt == Complex { + return nil, errf("SearchSorted: complex arrays have no ordering") + } + n := haystack.Len() + out := &Array{shape: needles.Shape(), dt: Int} + out.alloc(needles.Len()) + if n == 0 { + // Every position is zero, which is what the fresh payload holds. + return out, nil + } + // Both operands integer-class: compare natively in int64, because + // float64 would round the haystack edges and the needle onto one + // another above 2^53 and return the wrong insertion point. + if intClass(haystack.dt) && intClass(needles.dt) { + h, q := intPayload(haystack), intPayload(needles) + parallelMin(len(q), searchMinPerWorker, func(s, e int) { + oi := out.ints + for i := s; i < e; i++ { + // Bisection for the first element above the needle: lo + // ends at the count of elements at or below it, the + // rightmost position. + lo, hi := 0, n + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if h[mid] <= q[i] { + lo = mid + 1 + } else { + hi = mid + } + } + oi[i] = int64(lo) + } + }) + return out, nil + } + h, q := floatPayload(haystack), floatPayload(needles) + parallelMin(len(q), searchMinPerWorker, func(s, e int) { + oi := out.ints + for i := s; i < e; i++ { + // <= (not <) keeps the position after equal elements, as the + // rightmost rule documented above promises. A NaN needle + // fails every comparison and lands at position zero. + lo, hi := 0, n + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if h[mid] <= q[i] { + lo = mid + 1 + } else { + hi = mid + } + } + oi[i] = int64(lo) + } + }) + return out, nil +} + +// AssignBins maps every value to its bin index given ascending bin +// edges: bin k covers [edges[k], edges[k+1]). Values below the first or +// above the last edge clamp to the outer bins, and a NaN value keeps +// the outermost bin, where the comparisons leave it. The whole integer +// class selects bins natively in int64, where the widening of bool and +// every narrow integer is exact; float operands widen through floatAt, +// because float64 would round an int value onto a bin edge above 2^53. +// The workers write disjoint output slots and every bin index is an +// exact integer selection, so the split cannot move a single bin. +func AssignBins(a *Array, edges *Array) (*Array, error) { + if a.dt == Complex || edges.dt == Complex { + return nil, errf("AssignBins: complex arrays are not supported") + } + m := edges.Len() + if m < 2 { + return nil, errf("AssignBins: edges need at least two values") + } + out := &Array{shape: a.Shape(), dt: Int} + out.alloc(a.Len()) + // Both operands integer-class: select the bin natively in int64, + // because float64 would round the value onto a bin edge above 2^53 + // and pick the wrong one. + if intClass(a.dt) && intClass(edges.dt) { + ev, av := intPayload(edges), intPayload(a) + parallelMin(len(av), copyMinPerWorker, func(s, e int) { + oi := out.ints + for i := s; i < e; i++ { + v := av[i] + bin := 0 + if v >= ev[0] { + // The largest index whose edge sits at or below the + // value; the guard leaves below-range values in bin 0. + lo, hi := 0, m-1 + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if ev[mid] <= v { + lo = mid + 1 + } else { + hi = mid + } + } + bin = lo - 1 + } + oi[i] = int64(bin) + } + }) + return out, nil + } + ev, av := floatPayload(edges), floatPayload(a) + parallelMin(len(av), copyMinPerWorker, func(s, e int) { + oi := out.ints + for i := s; i < e; i++ { + v := av[i] + bin := 0 + // The guard keeps below-range values in bin 0; every value no + // edge compares at or below (a NaN) bisects down to lo zero + // and keeps the outermost bin, which is where the downward + // scan this replaces left it. + if !(v < ev[0]) { + lo, hi := 0, m-1 + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if ev[mid] <= v { + lo = mid + 1 + } else { + hi = mid + } + } + bin = lo - 1 + if lo == 0 { + bin = m - 2 + } + } + oi[i] = int64(bin) + } + }) + return out, nil +} + +// centralSum folds the central cross sum Σ(x−mx)(y−my) through the +// canonical partition: the per-element arithmetic the plain chain +// kept, cut into the fixed blocks the folds use, one chain partial per +// block and the balanced tree over them. At and below one block the +// walk is the plain chain, bit for bit; past it the measured +// deviation from the exact referent drops by orders of magnitude (the +// central-sum test pins the numbers). +func centralSum(x, y []float64, mx, my float64) float64 { + n := len(x) + parts := foldParts(n) + if parts == 1 { + var acc float64 + for i := range n { + acc += (x[i] - mx) * (y[i] - my) + } + return acc + } + partials := make([]float64, parts) + for c := range parts { + // The block's bounds are named locals: an unhoisted divided + // bound in the loop condition kept the walk off the fast path, + // three times the cost of the chain it replaced (measured). + lo, hi := c*n/parts, (c+1)*n/parts + var acc float64 + for i := lo; i < hi; i++ { + acc += (x[i] - mx) * (y[i] - my) + } + partials[c] = acc + } + return treeSum(partials) +} + +// centralMomentSums is centralSum carrying the two squared-deviation +// sums alongside the cross sum, the three Correlation needs. The +// partition and the per-element arithmetic are the same walk's. +func centralMomentSums(x, y []float64, mx, my float64) (num, dx2, dy2 float64) { + n := len(x) + parts := foldParts(n) + if parts == 1 { + for i := range n { + da := x[i] - mx + dbv := y[i] - my + num += da * dbv + dx2 += da * da + dy2 += dbv * dbv + } + return num, dx2, dy2 + } + pn, px, py := make([]float64, parts), make([]float64, parts), make([]float64, parts) + for c := range parts { + lo, hi := c*n/parts, (c+1)*n/parts + var sn, sx, sy float64 + for i := lo; i < hi; i++ { + da := x[i] - mx + dbv := y[i] - my + sn += da * dbv + sx += da * da + sy += dbv * dbv + } + pn[c], px[c], py[c] = sn, sx, sy + } + return treeSum(pn), treeSum(px), treeSum(py) +} + +// Covariance computes the sample covariance of two equally sized +// 1-D samples (denominator n−1). The samples are widened once and the +// central sum folds through the canonical partition, the accuracy the +// central-sum test measured against the exact referent. +func Covariance(a, b *Array) (float64, error) { + if a.dt == Complex || b.dt == Complex { + return 0, errf("Covariance: complex samples are not supported") + } + if a.Len() != b.Len() || a.Len() < 2 { + return 0, errf("Covariance: samples must share length ≥ 2") + } + ma, err := Mean(a) + if err != nil { + return 0, err + } + mb, err := Mean(b) + if err != nil { + return 0, err + } + av, bv := floatPayload(a), floatPayload(b) + return centralSum(av, bv, ma, mb) / float64(a.Len()-1), nil +} + +// Correlation computes the Pearson correlation coefficient of two +// 1-D samples: sum(dx·dy) / sqrt(sum dx² · sum dy²). The samples are +// widened once and the three central sums fold through the canonical +// partition, the accuracy the central-sum test measured against the +// exact referent. +func Correlation(a, b *Array) (float64, error) { + if a.dt == Complex || b.dt == Complex { + return 0, errf("Correlation: complex samples are not supported") + } + if a.Len() != b.Len() || a.Len() < 2 { + return 0, errf("Correlation: samples must share length ≥ 2") + } + ma, err := Mean(a) + if err != nil { + return 0, err + } + mb, err := Mean(b) + if err != nil { + return 0, err + } + av, bv := floatPayload(a), floatPayload(b) + num, da2, db2 := centralMomentSums(av, bv, ma, mb) + return num / math.Sqrt(da2*db2), nil +} + +// Sign returns the sign of every element (-1, 0 or +1) as a float +// array; complex inputs error. A NaN element has no sign and reports 0. +// The elements are widened once and the sign read off the payload +// slice, the same values the accessor walk read without the +// per-element dispatch; the destination starts zeroed, so only the +// non-zero marks are written. +func Sign(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("Sign: complex arrays are not supported") + } + out := zeros(Float, a.Shape()) + src := a + if !src.isContiguous() { + src = src.materialise() + } + vals := floatPayload(src) + d := out.floats + parallelMin(len(vals), elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + v := vals[i] + if v > 0 { + d[i] = 1 + } else if v < 0 { + d[i] = -1 + } + } + }) + return out, nil +} + +// IsNaN returns an int mask marking NaN elements. Complex arrays are +// not supported. +func IsNaN(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("IsNaN: complex arrays are not supported") + } + out := &Array{shape: a.Shape(), dt: Int} + out.alloc(a.Len()) + // The widening happens once and the mask starts zeroed, so only the + // marking slots are written; the chunks are disjoint. + vals := floatPayload(a) + parallelMin(len(vals), copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + if vals[i] != vals[i] { // NaN != NaN + out.ints[i] = 1 + } + } + }) + return out, nil +} + +// IsInf returns an int mask marking positive and negative infinity. +func IsInf(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("IsInf: complex arrays are not supported") + } + pos := math.Inf(1) + neg := math.Inf(-1) + out := &Array{shape: a.Shape(), dt: Int} + out.alloc(a.Len()) + vals := floatPayload(a) + parallelMin(len(vals), copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + if v := vals[i]; v == pos || v == neg { + out.ints[i] = 1 + } + } + }) + return out, nil +} + +// IsFinite returns an int mask marking finite values: neither NaN nor +// infinite. +func IsFinite(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("IsFinite: complex arrays are not supported") + } + pos := math.Inf(1) + neg := math.Inf(-1) + out := &Array{shape: a.Shape(), dt: Int} + out.alloc(a.Len()) + vals := floatPayload(a) + parallelMin(len(vals), copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + if v := vals[i]; v == v && v != pos && v != neg { + out.ints[i] = 1 + } + } + }) + return out, nil +} + +// LowerTriangle returns the lower triangular part of a square matrix as +// a copy with everything above the diagonal zeroed. Each row's kept +// prefix is one contiguous run, so a row copies whole rather than an +// element at a time. +func LowerTriangle(a *Array) (*Array, error) { + if a.NDim() != 2 || a.shape[0] != a.shape[1] { + return nil, errf("LowerTriangle: needs a square matrix, got %s", shapeText(a.shape)) + } + n := a.shape[0] + out := ZerosLike(a) + if a.isContiguous() && out.isContiguous() { + for r := range n { + copyRun(out, a, r*n, r*n, r+1) + } + return out, nil + } + for r := range n { + for c := range n { + if c <= r { + out.setFrom(r*n+c, a, r*n+c) + } + } + } + return out, nil +} + +// UpperTriangle returns the upper triangular part of a square matrix as +// a copy with everything below the diagonal zeroed. As in the lower +// builder, each row's kept suffix copies whole. +func UpperTriangle(a *Array) (*Array, error) { + if a.NDim() != 2 || a.shape[0] != a.shape[1] { + return nil, errf("UpperTriangle: needs a square matrix, got %s", shapeText(a.shape)) + } + n := a.shape[0] + out := ZerosLike(a) + if a.isContiguous() && out.isContiguous() { + for r := range n { + copyRun(out, a, r*n+r, r*n+r, n-r) + } + return out, nil + } + for r := range n { + for c := range n { + if c >= r { + out.setFrom(r*n+c, a, r*n+c) + } + } + } + return out, nil +} + +// copyRun copies run elements of src at srcOff to dst at dstOff. Both +// arrays carry the same dtype and are contiguous; the dispatch happens +// once per run rather than once per element, and the run itself is a +// single block copy. copyRun answers for every element type the package +// stores: bool, int8, uint8, int16, uint16, int32, uint32, int, +// float16, float32, float and complex, so the slice and join kernels +// take their block-copy path for every dtype. +func copyRun(dst, src *Array, dstOff, srcOff, run int) { + switch dst.dt { + case Int: + copy(dst.ints[dstOff:dstOff+run], src.ints[srcOff:srcOff+run]) + case Float16: + copy(dst.halves[dstOff:dstOff+run], src.halves[srcOff:srcOff+run]) + case Float32: + copy(dst.floats32[dstOff:dstOff+run], src.floats32[srcOff:srcOff+run]) + case Float: + copy(dst.floats[dstOff:dstOff+run], src.floats[srcOff:srcOff+run]) + case Complex: + copy(dst.complexes[dstOff:dstOff+run], src.complexes[srcOff:srcOff+run]) + case Bool: + copy(dst.bools[dstOff:dstOff+run], src.bools[srcOff:srcOff+run]) + case Int8: + copy(dst.i8s[dstOff:dstOff+run], src.i8s[srcOff:srcOff+run]) + case Uint8: + copy(dst.u8s[dstOff:dstOff+run], src.u8s[srcOff:srcOff+run]) + case Int16: + copy(dst.i16s[dstOff:dstOff+run], src.i16s[srcOff:srcOff+run]) + case Uint16: + copy(dst.u16s[dstOff:dstOff+run], src.u16s[srcOff:srcOff+run]) + case Int32: + copy(dst.i32s[dstOff:dstOff+run], src.i32s[srcOff:srcOff+run]) + case Uint32: + copy(dst.u32s[dstOff:dstOff+run], src.u32s[srcOff:srcOff+run]) + } +} + +// prodShape returns the product of a shape slice; 1 for an empty slice +// (matches Reshape's behaviour for a 0-D tensor). +func prodShape(s []int) int { + p := 1 + for _, v := range s { + p *= v + } + return p +} diff --git a/internal/core/arrayutil_edge_pins_test.go b/internal/core/arrayutil_edge_pins_test.go new file mode 100644 index 0000000..ed9c8ff --- /dev/null +++ b/internal/core/arrayutil_edge_pins_test.go @@ -0,0 +1,174 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "testing" +) + +// Edge pins for the array utilities: each pin names a behaviour that a +// defect once broke, and fails against the broken form. + +// TestRepeatZeroRepeatsEmptyResult pins Repeat with zero repeats: the +// extent collapses to zero and the answer is the empty array. The row +// walk used to visit the source rows regardless and sliced the empty +// destination payload, which panicked with a slice-bounds error +// whenever the repeated axis was not the innermost one. +func TestRepeatZeroRepeatsEmptyResult(t *testing.T) { + a, err := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + for _, dim := range []int{0, 1} { + got, err := Repeat(a, 0, dim) + if err != nil { + t.Fatalf("Repeat repeats=0 dim=%d: %v", dim, err) + } + if got.Len() != 0 || got.Shape()[dim] != 0 { + t.Errorf("Repeat repeats=0 dim=%d: shape %v, want a zero extent", dim, got.Shape()) + } + } + // The innermost-axis route must agree. + ints, err := FromInts([]int64{1, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatal(err) + } + got, err := Repeat(ints, 0, 1) + if err != nil { + t.Fatalf("Repeat repeats=0 innermost: %v", err) + } + if got.Len() != 0 { + t.Errorf("Repeat repeats=0 innermost: len %d, want 0", got.Len()) + } +} + +// TestSobolMatchesGrayReference pins the incremental Sobol walk against +// a from-scratch Gray-code reference: the XOR accumulation is exact, so +// the incremental advance must reproduce the reference bit for bit at +// any skip. +func TestSobolMatchesGrayReference(t *testing.T) { + for _, tc := range []struct { + n, dim, skip int + }{{257, 4, 0}, {1000, 7, 3}, {4096, 40, 0}, {513, 3, 65534}} { + got, err := SobolPoints(tc.n, tc.dim, tc.skip) + if err != nil { + t.Fatalf("SobolPoints(%d,%d,%d): %v", tc.n, tc.dim, tc.skip, err) + } + for d := range tc.dim { + var v [32]uint32 + sobolDirections(sobolTable[d], v[:]) + for p := range tc.n { + idx := tc.skip + p + 1 + gray := uint32(idx) ^ uint32(idx>>1) + var x uint32 + for b := range 32 { + if gray&(1< (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +func TestGrid(t *testing.T) { + a, _ := FromFloats([]float64{0, 1}, 2) + b, _ := FromFloats([]float64{10, 20, 30}, 3) + xg, yg, err := Grid(a, b) + if err != nil { + t.Fatal(err) + } + if xG := xg.Shape(); xG[0] != 3 || xG[1] != 2 { + t.Fatalf("xGrid shape: %v", xG) + } + // X varies along columns. + if v, _ := FloatAt(xg, 0, 1); v != 1 { + t.Errorf("X grid [0,1]: %v", v) + } + // Y repeats along rows. + if v, _ := FloatAt(yg, 2, 1); v != 30 { + t.Errorf("Y grid [2,1]: %v", v) + } +} + +func TestCrossProduct(t *testing.T) { + u, _ := FromFloats([]float64{1, 0, 0}, 3) + v, _ := FromFloats([]float64{0, 1, 0}, 3) + c, err := CrossProduct(u, v) + if err != nil { + t.Fatal(err) + } + want, _ := FromFloats([]float64{0, 0, 1}, 3) + if !Equal(want, c) { + t.Errorf("cross: got %v", c.RawFloats()) + } +} + +func TestIntegrate(t *testing.T) { + y, _ := FromFloats([]float64{0, 1, 2, 3}, 4) + v, err := Integrate(y, 1) + if err != nil { + t.Fatal(err) + } + if math.Abs(v-4.5) > 1e-9 { + t.Errorf("trapezoid of ramp: %v, want 4.5", v) + } + ci, err := CumulativeIntegrate(y, 1) + if err != nil { + t.Fatal(err) + } + if ci.Len() != 4 || ci.FloatAt(3) != 4.5 || ci.FloatAt(0) != 0 { + t.Errorf("cumulative integrate: %v", ci.RawFloats()) + } +} + +func TestInterpolate(t *testing.T) { + xs, _ := FromFloats([]float64{0, 10}, 2) + ys, _ := FromFloats([]float64{0, 100}, 2) + q, _ := FromFloats([]float64{-5, 5, 15}, 3) + out, err := Interpolate(xs, ys, q) + if err != nil { + t.Fatal(err) + } + for i, want := range []float64{0, 50, 100} { + if v := out.FloatAt(i); math.Abs(v-want) > 1e-9 { + t.Errorf("interp[%d]: %v, want %v", i, v, want) + } + } +} + +func TestMoveAxis(t *testing.T) { + a, _ := FromFloats(make([]float64, 24), 2, 3, 4) + for i := range a.Len() { + a.SetFloatAt(i, float64(i)) + } + moved, err := MoveAxis(a, 2, 0) + if err != nil { + t.Fatal(err) + } + if moved.Shape()[0] != 4 { + t.Fatalf("MoveAxis shape: %v", moved.Shape()) + } + // Element at output [0, d, h] equals input [d, h, 0]. + vOut, _ := FloatAt(moved, 0, 1, 1) + vIn, _ := FloatAt(a, 1, 1, 0) + if vOut != vIn { + t.Errorf("MoveAxis mapping broken: %v vs %v", vOut, vIn) + } +} + +func TestSearchSorted(t *testing.T) { + hay, _ := FromFloats([]float64{10, 20, 30, 40}, 4) + needles, _ := FromFloats([]float64{15, 5, 50, 20}, 4) + out, err := SearchSorted(hay, needles) + if err != nil { + t.Fatal(err) + } + // The rightmost rule: an exact hit lands after its equals, so 20 + // inserts at 2, not before its twin at 1. + want := []int64{1, 0, 4, 2} + for i := range want { + if v := out.RawInts()[i]; v != want[i] { + t.Errorf("searchsorted[%d]: %v, want %v", i, v, want[i]) + } + } + + // Duplicates in the haystack: every equal element is skipped. + dupHay, _ := FromFloats([]float64{1, 2, 2, 3}, 4) + dupNeedles, _ := FromFloats([]float64{2}, 1) + dupOut, err := SearchSorted(dupHay, dupNeedles) + if err != nil { + t.Fatal(err) + } + if got := dupOut.RawInts()[0]; got != 3 { + t.Errorf("searchsorted duplicates: %d, want 3", got) + } + + if _, err := SearchSorted( + mustFromComplexes(t, []complex128{1}, 1), + mustFromFloats(t, []float64{1}, 1)); err == nil { + t.Error("searchsorted complex haystack: expected error") + } +} + +func TestAssignBins(t *testing.T) { + edges, err0 := FromFloats([]float64{0, 10, 20}, 3) + if err0 != nil { + t.Fatal(err0) + } + vals, _ := FromFloats([]float64{-3, 5, 12, 25}, 4) + bins, err := AssignBins(vals, edges) + if err != nil { + t.Fatal(err) + } + wantBins := []int64{0, 0, 1, 1} + for i := range wantBins { + if v := bins.RawInts()[i]; v != wantBins[i] { + t.Errorf("bin[%d]: %v, want %v", i, v, wantBins[i]) + } + } +} + +func TestCovarianceCorrelation(t *testing.T) { + x, _ := FromFloats([]float64{1, 2, 3, 4}, 4) + y, _ := FromFloats([]float64{2, 4, 6, 8}, 4) + cov, err := Covariance(x, y) + if err != nil { + t.Fatal(err) + } + if math.Abs(cov-10.0/3.0) > 1e-9 { + t.Errorf("covariance: %v, want 10/3", cov) + } + r, err := Correlation(x, y) + if err != nil { + t.Fatal(err) + } + if math.Abs(r-1) > 1e-9 { + t.Errorf("correlation of identical trend: %v, want 1", r) + } + yFlip, _ := FromFloats([]float64{8, 6, 4, 2}, 4) + rNeg, _ := Correlation(x, yFlip) + if math.Abs(rNeg+1) > 1e-9 { + t.Errorf("anti-correlation: %v, want -1", rNeg) + } +} + +// TestBytesDegenerateInputs pins the guard that used to panic: Bytes +// on non-int payloads. +func TestBytesDegenerateInputs(t *testing.T) { + f := mustFromFloats(t, []float64{1, 2}, 2) + if b := f.Bytes(); b != nil { + t.Errorf("Bytes on float array: got %v, want nil", b) + } + i := mustFromInts(t, []int64{65, 66}, 2) + if b := i.Bytes(); string(b) != "AB" { + t.Errorf("Bytes on int array: got %q, want %q", b, "AB") + } +} + +// TestSpMulValidatesAndPromotes pins SpMul's guards: shape mismatch, +// out-of-range indices and mixed dtypes used to panic through nil +// payloads instead of erroring or promoting. +func TestSpMulValidatesAndPromotes(t *testing.T) { + vals := mustFromFloats(t, []float64{2, 3}, 2) + idx := mustFromInts(t, []int64{0, 1, 1, 0}, 2, 2) + sp := &SparseCOO{Indices: idx, Values: vals, Shape: []int{2, 2}} + + wrongShape := mustFromFloats(t, []float64{1, 2, 3}, 3) + if _, err := SpMul(sp, wrongShape); err == nil { + t.Error("SpMul shape mismatch: expected error") + } + + intDense := mustFromInts(t, []int64{10, 20, 30, 40}, 2, 2) + got, err := SpMul(sp, intDense) + if err != nil { + t.Fatalf("SpMul mixed dtype: %v", err) + } + if got.Dtype() != Float { + t.Fatalf("SpMul promote dtype: %s", got.Dtype()) + } + // Entry 0: value 2 at (0,1) -> 2*20 = 40; entry 1: value 3 at + // (1,0) -> 3*30 = 90. + if v, _ := FloatAt(got, 0, 1); v != 40 { + t.Errorf("SpMul [0,1]: %v, want 40", v) + } + if v, _ := FloatAt(got, 1, 0); v != 90 { + t.Errorf("SpMul [1,0]: %v, want 90", v) + } +} diff --git a/internal/core/arrayutil_test.go b/internal/core/arrayutil_test.go new file mode 100644 index 0000000..f73e89b --- /dev/null +++ b/internal/core/arrayutil_test.go @@ -0,0 +1,397 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +func TestLinspace(t *testing.T) { + v, err := Linspace(0, 1, 5) + if err != nil { + t.Fatal(err) + } + want := []float64{0, 0.25, 0.5, 0.75, 1} + for i := range 5 { + if g, _ := FloatAt(v, i); math.Abs(g-want[i]) > 1e-12 { + t.Errorf("Linspace[%d]: got %v, want %v", i, g, want[i]) + } + } + one, _ := Linspace(5, 9, 1) + if v, _ := FloatAt(one, 0); v != 5 { + t.Errorf("Linspace n=1: got %v, want 5", v) + } + empty, _ := Linspace(0, 1, 0) + if empty.Len() != 0 { + t.Errorf("Linspace n=0: len %d", empty.Len()) + } + if _, err := Linspace(0, 1, -1); err == nil { + t.Error("Linspace: expected error for negative n") + } +} + +func TestRepeat(t *testing.T) { + a, _ := FromInts([]int64{1, 2, 3}, 3) + r, err := Repeat(a, 2, 0) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{1, 1, 2, 2, 3, 3}, 6), r) { + t.Errorf("Repeat: %v", r.RawInts()) + } + m, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) + r2, err := Repeat(m, 2, 0) + if err != nil { + t.Fatal(err) + } + // Repeating along dim 0 duplicates whole rows. + if !Equal(mustFromInts(t, []int64{1, 2, 1, 2, 3, 4, 3, 4}, 4, 2), r2) { + t.Errorf("Repeat dim0: %v", r2.RawInts()) + } + if _, err := Repeat(m, 2, 5); err == nil { + t.Error("Repeat: expected error for out-of-range dim") + } +} + +func TestTile(t *testing.T) { + a, _ := FromInts([]int64{1, 2, 3}, 3) + tw, err := Tile(a, 2) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{1, 2, 3, 1, 2, 3}, 6), tw) { + t.Errorf("Tile: %v", tw.RawInts()) + } + m, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) + t4, err := Tile(m, 2, 2) + if err != nil { + t.Fatal(err) + } + want := []int64{1, 2, 1, 2, 3, 4, 3, 4, 1, 2, 1, 2, 3, 4, 3, 4} + if !Equal(mustFromInts(t, want, 4, 4), t4) { + t.Errorf("Tile 2x2: %v", t4.RawInts()) + } + if _, err := Tile(a, -1); err == nil { + t.Error("Tile: expected error for negative reps") + } +} + +func TestFlip(t *testing.T) { + a, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) + f, err := Flip(a) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{4, 3, 2, 1}, 2, 2), f) { + t.Errorf("Flip all: %v", f.RawInts()) + } + f0, err := Flip(a, 0) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{3, 4, 1, 2}, 2, 2), f0) { + t.Errorf("Flip dim0: %v", f0.RawInts()) + } + if _, err := Flip(a, 9); err == nil { + t.Error("Flip: expected error for out-of-range dim") + } +} + +func TestRoll(t *testing.T) { + a, _ := FromInts([]int64{1, 2, 3, 4, 5}, 5) + r, err := Roll(a, 2, 0) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{4, 5, 1, 2, 3}, 5), r) { + t.Errorf("Roll +2: %v", r.RawInts()) + } + rNeg, err := Roll(a, -1, 0) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{2, 3, 4, 5, 1}, 5), rNeg) { + t.Errorf("Roll -1: %v", rNeg.RawInts()) + } + if _, err := Roll(a, 1, 3); err == nil { + t.Error("Roll: expected error for out-of-range dim") + } +} + +func TestUnique(t *testing.T) { + a, _ := FromInts([]int64{3, 1, 2, 1, 3, 5}, 6) + u, err := Unique(a) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{1, 2, 3, 5}, 4), u) { + t.Errorf("Unique int: %v", u.RawInts()) + } + f, _ := FromFloats([]float64{2.5, math.NaN(), 1.5, 2.5, math.NaN()}, 5) + uf, err := Unique(f) + if err != nil { + t.Fatal(err) + } + // Sorted: 1.5, 2.5, then one NaN. + if v, _ := FloatAt(uf, 0); v != 1.5 { + t.Errorf("Unique float[0]: %v", v) + } + if v, _ := FloatAt(uf, 1); v != 2.5 { + t.Errorf("Unique float[1]: %v", v) + } + if v, _ := FloatAt(uf, 2); !math.IsNaN(v) { + t.Errorf("Unique float[2]: %v, want NaN", v) + } +} + +func TestArgwhere(t *testing.T) { + a, _ := FromInts([]int64{0, 5, 0, 7}, 2, 2) + aw, err := Argwhere(a) + if err != nil { + t.Fatal(err) + } + if aw.Shape()[0] != 2 || aw.Shape()[1] != 2 { + t.Fatalf("Argwhere shape: %v", aw.Shape()) + } + // Nonzeros at (0,1) and (1,1). + if v, _ := IntAt(aw, 0, 0); v != 0 { + t.Errorf("Argwhere[0,0]: %v", v) + } + if v, _ := IntAt(aw, 0, 1); v != 1 { + t.Errorf("Argwhere[0,1]: %v", v) + } + if v, _ := IntAt(aw, 1, 1); v != 1 { + t.Errorf("Argwhere[1,1]: %v", v) + } +} + +func TestAstype(t *testing.T) { + i, _ := FromInts([]int64{1, 2, 3}, 3) + f, err := Astype(i, Float) + if err != nil { + t.Fatal(err) + } + if f.Dtype() != Float || f.RawFloats()[0] != 1 { + t.Errorf("Astype int to float: %s %v", f.Dtype(), f.RawFloats()) + } + back, err := Astype(f, Int) + if err != nil { + t.Fatal(err) + } + if !Equal(back, i) { + t.Errorf("Astype round-trip: %v", back.RawInts()) + } + c, _ := FromComplexes([]complex128{1 + 2i, 3 - 4i}, 2) + cf, err := Astype(c, Float) + if err != nil { + t.Fatal(err) + } + if cf.RawFloats()[0] != 1 || cf.RawFloats()[1] != 3 { + t.Errorf("Astype complex to float: %v", cf.RawFloats()) + } + if _, err := Astype(c, Int); err == nil || !strings.Contains(err.Error(), "narrow") { + t.Errorf("Astype complex to int: %v", err) + } +} + +// TestAstypeNarrowDtypes pins the conversion matrix the narrow element +// types add: exact widening out of them, range-checked narrowing into +// them with the loud first-failure error, the zero-test into bool, the +// exact 0/1 widening out of bool, and the legacy pairs keeping their +// historical cast semantics beside all of it. +func TestAstypeNarrowDtypes(t *testing.T) { + // Bool is the zero-test target: NaN reads true, and bool widens out + // as 0/1, exact into complex. + f, _ := FromFloats([]float64{0, 1.5, math.NaN()}, 3) + b, err := Astype(f, Bool) + if err != nil { + t.Fatal(err) + } + if got := b.RawBools(); got[0] || !got[1] || !got[2] { + t.Errorf("Astype float to bool = %v", got) + } + cx, _ := FromComplexes([]complex128{0, 3i}, 2) + cxb, err := Astype(cx, Bool) + if err != nil { + t.Fatal(err) + } + if got := cxb.RawBools(); got[0] || !got[1] { + t.Errorf("Astype complex to bool = %v", got) + } + cb, err := Astype(b, Complex) + if err != nil { + t.Fatal(err) + } + if got := cb.RawComplexes(); got[0] != 0 || got[1] != 1 || got[2] != 1 { + t.Errorf("Astype bool to complex = %v", got) + } + bi, err := Astype(b, Int) + if err != nil { + t.Fatal(err) + } + if got := bi.RawInts(); got[0] != 0 || got[1] != 1 || got[2] != 1 { + t.Errorf("Astype bool to int = %v", got) + } + + // Exact widening out of the narrow integers. + i8, _ := FromInt8s([]int8{-128, 0, 127}, 3) + wInt, err := Astype(i8, Int) + if err != nil { + t.Fatal(err) + } + if got := wInt.RawInts(); got[0] != -128 || got[2] != 127 { + t.Errorf("Astype int8 to int = %v", got) + } + wC, err := Astype(i8, Complex) + if err != nil { + t.Fatal(err) + } + if got := wC.RawComplexes(); got[0] != complex(-128, 0) || got[2] != complex(127, 0) { + t.Errorf("Astype int8 to complex = %v", got) + } + u32, _ := FromUint32s([]uint32{4294967295}, 1) + wF, err := Astype(u32, Float) + if err != nil { + t.Fatal(err) + } + if got := wF.RawFloats(); got[0] != 4294967295 { + t.Errorf("Astype uint32 to float = %v", got) + } + wI, err := Astype(u32, Int) + if err != nil { + t.Fatal(err) + } + if got := wI.RawInts(); got[0] != 4294967295 { + t.Errorf("Astype uint32 to int = %v", got) + } + + // Narrow to narrow: a contained value set widens exactly, a + // non-contained one fails at the first offending index. + i16, err := Astype(i8, Int16) + if err != nil { + t.Fatal(err) + } + if got := i16.RawInt16s(); got[0] != -128 || got[2] != 127 { + t.Errorf("Astype int8 to int16 = %v", got) + } + if _, err := Astype(i8, Uint8); err == nil || + !strings.Contains(err.Error(), "Astype: value -128 at index 0 does not fit uint8") { + t.Errorf("Astype int8 to uint8: %v", err) + } + if _, err := Astype(u32, Int32); err == nil || + !strings.Contains(err.Error(), "Astype: value 4294967295 at index 0 does not fit int32") { + t.Errorf("Astype uint32 to int32: %v", err) + } + + // An int source checks exactly in int64 space, and the reported + // failure is the lowest index even when several elements overflow. + big, _ := FromInts([]int64{5, 300, -2}, 3) + if _, err := Astype(big, Int8); err == nil || + !strings.Contains(err.Error(), "Astype: value 300 at index 1 does not fit int8") { + t.Errorf("Astype int to int8: %v", err) + } + two, _ := FromInts([]int64{400, 300}, 2) + if _, err := Astype(two, Int8); err == nil || + !strings.Contains(err.Error(), "Astype: value 400 at index 0 does not fit int8") { + t.Errorf("Astype int to int8 first failure: %v", err) + } + + // A float source must be finite, integral and in range; the error + // names the value in the source's own float64 space. + ff, _ := FromFloats([]float64{2, -3, 127}, 3) + fi8, err := Astype(ff, Int8) + if err != nil { + t.Fatal(err) + } + if got := fi8.RawInt8s(); got[0] != 2 || got[1] != -3 || got[2] != 127 { + t.Errorf("Astype float to int8 = %v", got) + } + fb, _ := FromFloats([]float64{1, 2.5}, 2) + if _, err := Astype(fb, Int8); err == nil || + !strings.Contains(err.Error(), "Astype: value 2.5 at index 1 does not fit int8") { + t.Errorf("Astype float to int8: %v", err) + } + fInf, _ := FromFloats([]float64{math.Inf(1)}, 1) + if _, err := Astype(fInf, Uint16); err == nil || + !strings.Contains(err.Error(), "Astype: value +Inf at index 0 does not fit uint16") { + t.Errorf("Astype +Inf to uint16: %v", err) + } + fNaN, _ := FromFloats([]float64{math.NaN()}, 1) + if _, err := Astype(fNaN, Int8); err == nil || + !strings.Contains(err.Error(), "Astype: value NaN at index 0 does not fit int8") { + t.Errorf("Astype NaN to int8: %v", err) + } + + // A complex source into a narrow numeric target keeps the + // historical loud refusal; only bool reaches it by zero-test. + if _, err := Astype(cx, Uint8); err == nil || + !strings.Contains(err.Error(), "cannot narrow complex to uint8") { + t.Errorf("Astype complex to uint8: %v", err) + } + + // Legacy pairs keep their cast semantics beside the new rules: + // float to int truncates, with no range error. + f3, _ := FromFloats([]float64{3.9, -1.2}, 2) + li, err := Astype(f3, Int) + if err != nil { + t.Fatal(err) + } + if got := li.RawInts(); got[0] != 3 || got[1] != -1 { + t.Errorf("Astype float to int legacy truncation = %v", got) + } + + // Same dtype copies through cloneArray for a narrow dtype; Equal + // still defaults to the complex payload for these dtypes, so the + // comparison reads the payload directly. + u8, _ := FromUint8s([]uint8{1, 2, 3}, 3) + cp, err := Astype(u8, Uint8) + if err != nil { + t.Fatal(err) + } + if cp.Dtype() != Uint8 { + t.Fatalf("Astype uint8 to uint8 dtype = %s", cp.Dtype()) + } + if got := cp.RawUint8s(); got[0] != 1 || got[1] != 2 || got[2] != 3 { + t.Errorf("Astype uint8 to uint8 = %v", got) + } +} + +func TestItem(t *testing.T) { + a, _ := FromFloats([]float64{3.25}, 1) + v, err := Item(a) + if err != nil { + t.Fatal(err) + } + if v != 3.25 { + t.Errorf("Item: %v", v) + } + b, _ := FromFloats([]float64{1, 2}, 2) + if _, err := Item(b); err == nil { + t.Error("Item: expected error for multi-element array") + } +} + +func TestDiag(t *testing.T) { + a, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) + d, err := Diag(a) + if err != nil { + t.Fatal(err) + } + if !Equal(mustFromInts(t, []int64{1, 4}, 2), d) { + t.Errorf("Diag 2-D: %v", d.RawInts()) + } + v, _ := FromInts([]int64{5, 6}, 2) + m, err := Diag(v) + if err != nil { + t.Fatal(err) + } + if m.Shape()[0] != 2 || m.Shape()[1] != 2 { + t.Fatalf("Diag 1-D shape: %v", m.Shape()) + } + if val, _ := IntAt(m, 1, 1); val != 6 { + t.Errorf("Diag 1-D [1,1]: %v", val) + } +} diff --git a/internal/core/axis.go b/internal/core/axis.go new file mode 100644 index 0000000..e6f2013 --- /dev/null +++ b/internal/core/axis.go @@ -0,0 +1,1563 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "cmp" + "math" + "slices" + "sync/atomic" +) + +// Axis-based reductions: the softmax, normalisation and loss +// primitive. Each reduces along one dimension and keeps the others in +// order; the result rank is NDim-1. Reducing the only dimension of a 1-D +// array is an error pointing at the global variant: tensor has no 0-d +// arrays. + +// axisOp selects the fold an axis reduction applies to each line. +type axisOp uint8 + +const ( + opSum axisOp = iota + opMean + opMin + opMax +) + +// SumAxis returns the sums along the given dimension. +func SumAxis(a *Array, dim int) (*Array, error) { + return a.reduceAxis(dim, "SumAxis", opSum) +} + +// MinAxis returns the smallest values along the given dimension; NaN +// elements never win, complex arrays have no ordering. +func MinAxis(a *Array, dim int) (*Array, error) { + return a.reduceAxis(dim, "MinAxis", opMin) +} + +// MaxAxis returns the largest values along the given dimension; NaN +// elements never win, complex arrays have no ordering. +func MaxAxis(a *Array, dim int) (*Array, error) { + return a.reduceAxis(dim, "MaxAxis", opMax) +} + +// MeanAxis returns the float means along the given dimension; complex +// arrays have no float mean. +func MeanAxis(a *Array, dim int) (*Array, error) { + if a.dt == Complex { + return nil, errf("MeanAxis: complex arrays have no float mean") + } + return a.reduceAxis(dim, "MeanAxis", opMean) +} + +// reduceAxis folds every element into the accumulator slot addressed by +// the source coordinate with dim dropped. The walk is line based, like +// the Norm and scanDim kernels: a line is the run of a.shape[dim] +// elements that share every surviving coordinate, so the destination +// index collapses to b*stride + s and the per-element odometer +// disappears. Whole lines are handed to each worker, which makes every +// accumulator slot single-writer, so there is no merge phase; the +// fan-out is capped by splitCapped so a small fold never pays a +// core-count spawn bill. +// The float and complex sums fold each line through the canonical +// partition the global sums use: fixed blocks of the line, the block +// partials combined through the balanced tree. The partition follows +// from the line length alone, so a line's answer is the same bits +// whatever the worker split, and a single-line fold answers the bits of +// Sum over the same elements; against the single chain it replaced the +// tree holds full accuracy on lines long enough for a chain's roundings +// to pile up, measured against big.Float in the accuracy test. The +// integer sums stay plain chains: wrapping addition is exact under any +// grouping. Extrema seed each line from its first non-NaN element, NaN +// candidates never win (they fail every comparison), and a line whose +// every element is NaN never seeds and reports NaN. mean divides by the +// reduced dimension afterwards and accumulates in float regardless of +// the input dtype; everything else keeps its dtype. float16 and float32 +// reductions fold in a float64 scratch and narrow once. +func (a *Array) reduceAxis(dim int, name string, op axisOp) (*Array, error) { + if dim < 0 || dim >= a.NDim() { + return nil, errf("%s: dimension %d is out of range for shape %s", name, dim, shapeText(a.shape)) + } + if a.NDim() == 1 { + return nil, errf("%s: reducing the only dimension of a 1-D array, use the global variant", name) + } + if a.dt == Complex && op != opSum { + return nil, errf("%s: complex arrays have no ordering", name) + } + + newShape := make([]int, 0, a.NDim()-1) + newShape = append(newShape, a.shape[:dim]...) + newShape = append(newShape, a.shape[dim+1:]...) + total := 1 + for _, d := range newShape { + total *= d + } + + acc := &Array{shape: newShape, dt: a.dt} + if op == opMean { + acc.dt = Float + } else if intClass(a.dt) && a.dt != Int { + // The scalar reductions answer Int scalars for the whole integer + // class; the axis folds mirror that with Int accumulators fed by + // exact widenings. + acc.dt = Int + } + // float16 and float32 reductions fold in a float64 scratch and + // narrow once; every other dtype folds straight into the + // accumulator. + scratch := acc + if (a.dt == Float16 || a.dt == Float32) && op != opMean { + scratch = &Array{shape: newShape, dt: Float} + } + scratch.alloc(total) + + stride := 1 + for k := dim + 1; k < a.NDim(); k++ { + stride *= a.shape[k] + } + line := a.shape[dim] + perLine := stride * line + lines := 0 + if perLine > 0 { + lines = a.Len() / perLine + } + seed := op == opMin || op == opMax + + // One line per accumulator slot, whole lines per worker: the slot at + // b*stride + s is written exactly once, by the worker that owns + // line b. The fan-out is capped by splitCapped so every worker + // carries at least reduceSplitFloor elements. + work := lines * perLine + switch a.dt { + case Int: + switch { + case seed: + foldExtremeInt(a.ints, scratch.ints, op == opMin, lines, work, perLine, stride, line) + case op == opMean: + foldMeanInt(a.ints, scratch.floats, lines, work, perLine, stride, line) + default: + foldSum(a.ints, scratch.ints, lines, work, perLine, stride, line) + } + case Bool: + switch { + case seed: + foldExtremeAxisBool(a.bools, scratch.ints, op == opMin, lines, work, perLine, stride, line) + case op == opMean: + foldMeanAxisBool(a.bools, scratch.floats, lines, work, perLine, stride, line) + default: + foldSumAxisBool(a.bools, scratch.ints, lines, work, perLine, stride, line) + } + case Int8: + axisFoldNarrow(a.i8s, scratch, op, seed, lines, work, perLine, stride, line) + case Uint8: + axisFoldNarrow(a.u8s, scratch, op, seed, lines, work, perLine, stride, line) + case Int16: + axisFoldNarrow(a.i16s, scratch, op, seed, lines, work, perLine, stride, line) + case Uint16: + axisFoldNarrow(a.u16s, scratch, op, seed, lines, work, perLine, stride, line) + case Int32: + axisFoldNarrow(a.i32s, scratch, op, seed, lines, work, perLine, stride, line) + case Uint32: + axisFoldNarrow(a.u32s, scratch, op, seed, lines, work, perLine, stride, line) + case Float16: + if seed { + foldExtremeHalf(a.halves, scratch.floats, op == opMin, lines, work, perLine, stride, line) + } else { + foldSumAxisF16(a.halves, scratch.floats, lines, work, perLine, stride, line) + } + case Float32: + if seed { + foldExtremeFloat(a.floats32, scratch.floats, op == opMin, lines, work, perLine, stride, line) + } else { + foldSumAxisF32(a.floats32, scratch.floats, lines, work, perLine, stride, line) + } + case Float: + if seed { + foldExtremeFloat(a.floats, scratch.floats, op == opMin, lines, work, perLine, stride, line) + } else { + foldSumAxis(a.floats, scratch.floats, lines, work, perLine, stride, line) + } + default: + foldSumAxis(a.complexes, scratch.complexes, lines, work, perLine, stride, line) + } + + // A zero-length reduction dimension leaves every line empty, so the + // walk never runs and nothing can seed: each float slot reports the + // missing value NaN, and int slots keep their zeros. This is the + // same outcome the all-NaN fill rule produced, and only float + // accumulators report NaN because int elements always seed. + if seed && line == 0 { + // Only the float64 accumulator can land here unseeded: int + // elements always seed the first line, and every scratch this + // path allocates is the float64 one (a float32 or half input + // widens into it), so no narrower scratch exists to fill. + for d := range total { + if scratch.dt == Float { + scratch.floats[d] = math.NaN() + } + } + } + + if op == opMean { + line := float64(a.shape[dim]) + for k := range acc.floats { + acc.floats[k] /= line + } + return acc, nil + } + if scratch != acc { + if acc.dt == Float16 { + acc.halves = make([]uint16, total) + for k, v := range scratch.floats { + acc.halves[k] = HalfFromFloat64(v) + } + } else { + acc.floats32 = make([]float32, total) + for k, v := range scratch.floats { + acc.floats32[k] = float32(v) + } + } + } + return acc, nil +} + +// reduceSplitFloor is the element count every spawned fold worker must +// carry before splitCapped grants it a goroutine: the folds stream about +// one element per instruction, so the spawn and its synchronisation only +// amortise above this many of them, and a smaller chunk costs more to +// schedule than to run (the axis sweep splits between 4,096 elements, +// where the lone walk wins outright, and 16,384, where the capped +// fan-out is ahead). +const reduceSplitFloor = 16_384 + +// foldSum adds every line's int64 elements into the slot the line +// shares, in ascending element order. Wrapping addition is associative, +// so no grouping can move a bit and the chain needs no partition. +func foldSum(src, dst []int64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + // Contiguous lines: one slice expression bounds-checks the + // whole line and the range walk streams it. + for b := ls; b < le; b++ { + var m int64 + for _, v := range src[b*perLine : b*perLine+line] { + m += v + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + var m int64 + for off := range line { + m += src[base+off*stride+s] + } + dst[b*stride+s] = m + } + } + }) +} + +// foldBlock is one canonical block's fold over the line elements +// src[base+off*stride] for off in [0, count): four interleaved chains +// combined as ((s0+s1)+(s2+s3)), the pairing foldRange keeps. A stride +// of one visits the block in exactly foldRange's order, so a contiguous +// line's block value is the value reduce.go's block fold gives. +func foldBlock[N float64 | complex128](src []N, base, stride, count int) N { + var s0, s1, s2, s3 N + i := 0 + for ; i+4 <= count; i += 4 { + o := base + i*stride + s0 += src[o] + s1 += src[o+stride] + s2 += src[o+2*stride] + s3 += src[o+3*stride] + } + for ; i < count; i++ { + s0 += src[base+i*stride] + } + return (s0 + s1) + (s2 + s3) +} + +// lineFoldParts folds one line through the canonical partition: fixed +// block boundaries at c·line/foldParts(line), each block's partial from +// the block closure, the partials combined through the balanced tree. +// The line length alone picks the shape, so the value is the same bits +// whatever the worker split, and a single-block line is one block fold. +func lineFoldParts[N float64 | complex128](line int, block func(lo, hi int) N) N { + parts := foldParts(line) + if parts == 1 { + return block(0, line) + } + partials := make([]N, parts) + for c := range parts { + lo, hi := c*line/parts, (c+1)*line/parts + partials[c] = block(lo, hi) + } + return treeSum(partials) +} + +// foldSumAxis folds every float64 or complex128 line through the +// canonical partition the global sums use: the line's blocks from +// foldBlock, the partials combined through treeSum. Against the single +// chain it replaced, the tree holds its accuracy on lines long enough +// for a chain's roundings to pile up and loses nothing on the short +// ones; on the large-plus-small counterpoint, where a chain's running +// total swallows the small elements outright, the tree's blocks keep +// them. Contiguous lines reuse reduce.go's own block folds, so a +// single-line fold answers Sum's exact bits. +func foldSumAxis[N float64 | complex128](src, dst []N, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + dst[b] = lineFoldParts(line, func(lo, hi int) N { + return foldBlock(row, lo, 1, hi-lo) + }) + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + dst[b*stride+s] = lineFoldParts(line, func(lo, hi int) N { + return foldBlock(src, base+lo*stride+s, stride, hi-lo) + }) + } + } + }) +} + +// foldBlockWidenF32 is foldBlock over a float32 payload: every element +// widens exactly, so the chains see the values an accessor read would. +func foldBlockWidenF32(src []float32, base, stride, count int) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= count; i += 4 { + o := base + i*stride + s0 += float64(src[o]) + s1 += float64(src[o+stride]) + s2 += float64(src[o+2*stride]) + s3 += float64(src[o+3*stride]) + } + for ; i < count; i++ { + s0 += float64(src[base+i*stride]) + } + return (s0 + s1) + (s2 + s3) +} + +// foldBlockWidenF16 is foldBlockWidenF32 for a uint16 half payload: the +// raw bits widen with HalfToFloat64, never cast as integers. +func foldBlockWidenF16(src []uint16, base, stride, count int) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= count; i += 4 { + o := base + i*stride + s0 += HalfToFloat64(src[o]) + s1 += HalfToFloat64(src[o+stride]) + s2 += HalfToFloat64(src[o+2*stride]) + s3 += HalfToFloat64(src[o+3*stride]) + } + for ; i < count; i++ { + s0 += HalfToFloat64(src[base+i*stride]) + } + return (s0 + s1) + (s2 + s3) +} + +// foldSumAxisF32 is foldSumAxis for a float32 payload folded into a +// float64 scratch through the same canonical partition; every widening +// is exact, so the values match the accessor walk. +func foldSumAxisF32(src []float32, dst []float64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + dst[b] = lineFoldParts(line, func(lo, hi int) float64 { + return foldBlockWidenF32(row, lo, 1, hi-lo) + }) + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + dst[b*stride+s] = lineFoldParts(line, func(lo, hi int) float64 { + return foldBlockWidenF32(src, base+lo*stride+s, stride, hi-lo) + }) + } + } + }) +} + +// foldSumAxisF16 is foldSumAxisF32 for a half payload's raw bit +// patterns, widened with HalfToFloat64 exactly as they are read. +func foldSumAxisF16(src []uint16, dst []float64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + dst[b] = lineFoldParts(line, func(lo, hi int) float64 { + return foldBlockWidenF16(row, lo, 1, hi-lo) + }) + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + dst[b*stride+s] = lineFoldParts(line, func(lo, hi int) float64 { + return foldBlockWidenF16(src, base+lo*stride+s, stride, hi-lo) + }) + } + } + }) +} + +// foldMeanInt is foldSumWiden for an int payload: the widened elements +// accumulate in float64 exactly as the per-element walk did. +func foldMeanInt(src []int64, dst []float64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + var m float64 + for _, v := range src[b*perLine : b*perLine+line] { + m += float64(v) + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + var m float64 + for off := range line { + m += float64(src[base+off*stride+s]) + } + dst[b*stride+s] = m + } + } + }) +} + +// foldExtremeInt writes the smallest (wantMin) or largest element of +// every line. Int lines always seed: the first element opens the line +// and the rest compare against it, in ascending order, exactly as the +// per-element walk compared them. +func foldExtremeInt(src, dst []int64, wantMin bool, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + m := row[0] + if wantMin { + for _, v := range row[1:] { + if v < m { + m = v + } + } + } else { + for _, v := range row[1:] { + if v > m { + m = v + } + } + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + m := src[base+s] + if wantMin { + for off := 1; off < line; off++ { + if v := src[base+off*stride+s]; v < m { + m = v + } + } + } else { + for off := 1; off < line; off++ { + if v := src[base+off*stride+s]; v > m { + m = v + } + } + } + dst[b*stride+s] = m + } + } + }) +} + +// foldExtremeFloat is foldExtremeInt for the float payloads, folded into +// a float64 scratch: the line seeds from its first non-NaN element, NaN +// candidates fail every comparison and can neither seed nor win, and a +// line whose every element is NaN never seeds and reports NaN. The seed +// scan and the compare walk visit the elements in the same order the +// single-loop walk did, so every selection is unchanged; splitting the +// two phases only lifts the per-element seeded test off the hot loop. +func foldExtremeFloat[F float32 | float64](src []F, dst []float64, wantMin bool, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + i, m := 0, 0.0 + for i < line { + if v := float64(row[i]); v == v { + m = v + break + } + i++ + } + if i == line { + // Every element of the line is NaN. + dst[b] = math.NaN() + continue + } + i++ // step past the seed + if wantMin { + for _, v := range row[i:] { + if w := float64(v); w < m { + m = w + } + } + } else { + for _, v := range row[i:] { + if w := float64(v); w > m { + m = w + } + } + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + i, m := 0, 0.0 + for i < line { + if v := float64(src[base+i*stride+s]); v == v { + m = v + break + } + i++ + } + if i == line { + // Every element of the line is NaN. + dst[b*stride+s] = math.NaN() + continue + } + i++ // step past the seed + if wantMin { + for off := i; off < line; off++ { + if w := float64(src[base+off*stride+s]); w < m { + m = w + } + } + } else { + for off := i; off < line; off++ { + if w := float64(src[base+off*stride+s]); w > m { + m = w + } + } + } + dst[b*stride+s] = m + } + } + }) +} + +// foldExtremeHalf is foldExtremeFloat for a uint16 half payload: the +// raw bits widen with HalfToFloat64, never cast as integers, so the +// seed scan and the comparisons see exactly the values floatAt would +// hand over. +func foldExtremeHalf(src []uint16, dst []float64, wantMin bool, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + i, m := 0, 0.0 + for i < line { + if v := HalfToFloat64(row[i]); v == v { + m = v + break + } + i++ + } + if i == line { + // Every element of the line is NaN. + dst[b] = math.NaN() + continue + } + i++ // step past the seed + if wantMin { + for _, v := range row[i:] { + if w := HalfToFloat64(v); w < m { + m = w + } + } + } else { + for _, v := range row[i:] { + if w := HalfToFloat64(v); w > m { + m = w + } + } + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + i, m := 0, 0.0 + for i < line { + if v := HalfToFloat64(src[base+i*stride+s]); v == v { + m = v + break + } + i++ + } + if i == line { + // Every element of the line is NaN. + dst[b*stride+s] = math.NaN() + continue + } + i++ // step past the seed + if wantMin { + for off := i; off < line; off++ { + if w := HalfToFloat64(src[base+off*stride+s]); w < m { + m = w + } + } + } else { + for off := i; off < line; off++ { + if w := HalfToFloat64(src[base+off*stride+s]); w > m { + m = w + } + } + } + dst[b*stride+s] = m + } + } + }) +} + +// axisFoldNarrow runs one narrow integer payload through reduceAxis's +// fold selection: sums and extrema widen exactly into the Int +// accumulator, means widen into the float64 scratch. +func axisFoldNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, scratch *Array, op axisOp, seed bool, lines, work, perLine, stride, line int) { + switch { + case seed: + foldExtremeAxisNarrow(src, scratch.ints, op == opMin, lines, work, perLine, stride, line) + case op == opMean: + foldMeanAxisNarrow(src, scratch.floats, lines, work, perLine, stride, line) + default: + foldSumAxisNarrow(src, scratch.ints, lines, work, perLine, stride, line) + } +} + +// foldSumAxisNarrow is foldSum for a narrow integer payload folded into +// an int64 accumulator: every element widens exactly. +func foldSumAxisNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, dst []int64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + var m int64 + for _, v := range src[b*perLine : b*perLine+line] { + m += int64(v) + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + var m int64 + for off := range line { + m += int64(src[base+off*stride+s]) + } + dst[b*stride+s] = m + } + } + }) +} + +// foldSumAxisBool is foldSumAxisNarrow for a bool payload: each line +// counts its true elements. +func foldSumAxisBool(src []bool, dst []int64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + var m int64 + for _, v := range src[b*perLine : b*perLine+line] { + if v { + m++ + } + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + var m int64 + for off := range line { + if src[base+off*stride+s] { + m++ + } + } + dst[b*stride+s] = m + } + } + }) +} + +// foldMeanAxisNarrow is foldMeanInt for a narrow integer payload: the +// widened elements accumulate in float64 exactly as the per-element walk +// did. +func foldMeanAxisNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, dst []float64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + var m float64 + for _, v := range src[b*perLine : b*perLine+line] { + m += float64(v) + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + var m float64 + for off := range line { + m += float64(src[base+off*stride+s]) + } + dst[b*stride+s] = m + } + } + }) +} + +// foldMeanAxisBool counts a bool line into the float64 mean scratch. +func foldMeanAxisBool(src []bool, dst []float64, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + var m float64 + for _, v := range src[b*perLine : b*perLine+line] { + if v { + m++ + } + } + dst[b] = m + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + var m float64 + for off := range line { + if src[base+off*stride+s] { + m++ + } + } + dst[b*stride+s] = m + } + } + }) +} + +// foldExtremeAxisNarrow is foldExtremeInt for a narrow integer payload: +// the line compares in the payload's own type, so no value ever meets a +// float64 rounding, and the winner widens exactly into the Int +// accumulator. +func foldExtremeAxisNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, dst []int64, wantMin bool, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + m := row[0] + if wantMin { + for _, v := range row[1:] { + if v < m { + m = v + } + } + } else { + for _, v := range row[1:] { + if v > m { + m = v + } + } + } + dst[b] = int64(m) + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + m := src[base+s] + if wantMin { + for off := 1; off < line; off++ { + if v := src[base+off*stride+s]; v < m { + m = v + } + } + } else { + for off := 1; off < line; off++ { + if v := src[base+off*stride+s]; v > m { + m = v + } + } + } + dst[b*stride+s] = int64(m) + } + } + }) +} + +// boolToInt64 widens a bool to the 0/1 the int64 accumulators of the +// extrema family carry. +func boolToInt64(b bool) int64 { + if b { + return 1 + } + return 0 +} + +// foldExtremeAxisBool is foldExtremeAxisNarrow for a bool payload: +// false below true, widened to the 0/1 the Int accumulator carries. Go +// orders no bool with < or >, so the strict improvement writes its own +// logic. +func foldExtremeAxisBool(src []bool, dst []int64, wantMin bool, lines, work, perLine, stride, line int) { + splitCapped(lines, work, reduceSplitFloor, func(ls, le int) { + if stride == 1 { + for b := ls; b < le; b++ { + row := src[b*perLine : b*perLine+line] + m := row[0] + if wantMin { + for _, v := range row[1:] { + if !v && m { + m = v + } + } + } else { + for _, v := range row[1:] { + if v && !m { + m = v + } + } + } + dst[b] = boolToInt64(m) + } + return + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + m := src[base+s] + if wantMin { + for off := 1; off < line; off++ { + if v := src[base+off*stride+s]; !v && m { + m = v + } + } + } else { + for off := 1; off < line; off++ { + if v := src[base+off*stride+s]; v && !m { + m = v + } + } + } + dst[b*stride+s] = boolToInt64(m) + } + } + }) +} + +// ArgMax returns the index of the largest element of a 1-D array; NaN +// elements are skipped as missing. An empty or all-NaN array, a complex +// array, or a non-1-D shape is an error. +func ArgMax(a *Array) (int, error) { + return a.argExtreme("ArgMax", false) +} + +// ArgMin returns the index of the smallest element of a 1-D array; NaN +// elements are skipped as missing. +func ArgMin(a *Array) (int, error) { + return a.argExtreme("ArgMin", true) +} + +func (a *Array) argExtreme(name string, wantMin bool) (int, error) { + if a.NDim() != 1 { + return 0, errf("%s: needs a 1-D array, got shape %s", name, shapeText(a.shape)) + } + if a.dt == Complex { + return 0, errf("%s: complex arrays have no ordering", name) + } + // Both dtype walks read the payload at the logical flat index, so a + // strided view is materialised first; a contiguous array is returned + // unchanged, keeping the hot path allocation-free. + a = a.materialise() + // An integer-class array compares in its own payload type: the reason + // the int64 walk exists is that floatAt rounds above 2^53, and the + // narrow widths and bool keep the same native-comparison contract. + if intClass(a.dt) { + n := a.Len() + switch a.dt { + case Int: + best := -1 + for i := range n { + v := a.ints[i] + if best < 0 { + best = i + continue + } + b := a.ints[best] + if (wantMin && v < b) || (!wantMin && v > b) { + best = i + } + } + if best < 0 { + return 0, errf("%s: the array is empty", name) + } + return best, nil + case Bool: + return argExtremeIndexBool(a.bools, n, name, wantMin) + case Int8: + return argExtremeIndex(a.i8s, n, name, wantMin) + case Uint8: + return argExtremeIndex(a.u8s, n, name, wantMin) + case Int16: + return argExtremeIndex(a.i16s, n, name, wantMin) + case Uint16: + return argExtremeIndex(a.u16s, n, name, wantMin) + case Int32: + return argExtremeIndex(a.i32s, n, name, wantMin) + default: + return argExtremeIndex(a.u32s, n, name, wantMin) + } + } + best := -1 + for i := range a.Len() { + v := a.floatAt(i) + if v != v { + continue + } + if best < 0 { + best = i + continue + } + b := a.floatAt(best) + if (wantMin && v < b) || (!wantMin && v > b) { + best = i + } + } + if best < 0 { + return 0, errf("%s: every element is NaN", name) + } + return best, nil +} + +// argExtremeIndex walks the first n elements of a narrow integer payload +// comparing in the payload's own type; such a payload never holds a NaN, +// so the first element always seeds the walk. +func argExtremeIndex[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int, name string, wantMin bool) (int, error) { + if n == 0 { + return 0, errf("%s: the array is empty", name) + } + best := 0 + for i := 1; i < n; i++ { + if (wantMin && src[i] < src[best]) || (!wantMin && src[i] > src[best]) { + best = i + } + } + return best, nil +} + +// argExtremeIndexBool is argExtremeIndex for a bool payload: Go orders +// no bool with < or >, so the strict improvement writes its own logic, +// and a tie keeps the earlier index exactly as the integer walk does. +func argExtremeIndexBool(src []bool, n int, name string, wantMin bool) (int, error) { + if n == 0 { + return 0, errf("%s: the array is empty", name) + } + best := 0 + for i := 1; i < n; i++ { + if (wantMin && !src[i] && src[best]) || (!wantMin && src[i] && !src[best]) { + best = i + } + } + return best, nil +} + +// ArgMaxAxis returns the indices of the maximum values along the given +// dimension. The result is an Int array with the same shape as the +// receiver except that the chosen dimension is dropped: the same +// reduction shape SumAxis / MaxAxis / MinAxis produce. Reducing the only +// dimension of a 1-D array is an error; use ArgMax instead. +func ArgMaxAxis(a *Array, dim int) (*Array, error) { + return a.argExtremeAxis(dim, false) +} + +// ArgMinAxis returns the indices of the minimum values along the given +// dimension. +func ArgMinAxis(a *Array, dim int) (*Array, error) { + return a.argExtremeAxis(dim, true) +} + +// argExtremeAxis is the shared implementation behind ArgMaxAxis and +// ArgMinAxis. The result drops the chosen dimension, like SumAxis / +// MaxAxis, and each element is the position of the extreme along it. +// NaN elements are skipped as missing, mirroring the 1-D ArgMax: each +// line seeds from its first non-NaN element, and a line with no finite +// element at all is an error. +func (a *Array) argExtremeAxis(dim int, wantMin bool) (*Array, error) { + name := "ArgMaxAxis" + global := "ArgMax" + if wantMin { + name = "ArgMinAxis" + global = "ArgMin" + } + if dim < 0 || dim >= a.NDim() { + return nil, errf("%s: dimension %d out of range for shape %s", name, dim, shapeText(a.shape)) + } + if a.dt == Complex { + return nil, errf("%s: complex arrays have no ordering", name) + } + if a.NDim() == 1 { + return nil, errf("%s: reducing the only dimension of a 1-D array, use %s", name, global) + } + if a.Len() == 0 { + return nil, errf("%s: an empty array has no result", name) + } + newShape := reduceShape(a.shape, dim) + out := &Array{shape: newShape, dt: Int} + total := 1 + for _, d := range newShape { + total *= d + } + out.alloc(total) + stride := 1 + for k := a.NDim() - 1; k > dim; k-- { + stride *= a.shape[k] + } + perLine := a.shape[dim] * stride + lines := a.Len() / perLine + lineLen := a.shape[dim] + var allNaN atomic.Bool + // The dtype dispatch sits outside the walk: elements come straight + // from the payload, int and float32 widening to float64 exactly. + switch a.dt { + case Int: + src := a.ints + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for line := ls; line < le; line++ { + base := line * perLine + for post := range stride { + // Compared as int64: widening to float64 would round + // above 2^53 and pick the wrong element of two + // neighbours. An int is never NaN, so the first + // element always seeds. + best, bestVal := 0, src[base+post] + for k := 1; k < lineLen; k++ { + if v := src[base+k*stride+post]; wantMin { + if v < bestVal { + best, bestVal = k, v + } + } else if v > bestVal { + best, bestVal = k, v + } + } + out.ints[line*stride+post] = int64(best) + } + } + }) + case Bool: + argExtremeAxisLineBool(a.bools, a.Len(), out, wantMin, lines, perLine, stride, lineLen) + case Int8: + argExtremeAxisLine(a.i8s, a.Len(), out, wantMin, lines, perLine, stride, lineLen) + case Uint8: + argExtremeAxisLine(a.u8s, a.Len(), out, wantMin, lines, perLine, stride, lineLen) + case Int16: + argExtremeAxisLine(a.i16s, a.Len(), out, wantMin, lines, perLine, stride, lineLen) + case Uint16: + argExtremeAxisLine(a.u16s, a.Len(), out, wantMin, lines, perLine, stride, lineLen) + case Int32: + argExtremeAxisLine(a.i32s, a.Len(), out, wantMin, lines, perLine, stride, lineLen) + case Uint32: + argExtremeAxisLine(a.u32s, a.Len(), out, wantMin, lines, perLine, stride, lineLen) + case Float16: + src := a.halves + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for line := ls; line < le; line++ { + base := line * perLine + for post := range stride { + // Seed from the first non-NaN element; NaN candidates + // never enter the comparison. The widening is exact, + // so the half values compare exactly in float64. + best := -1 + var bestVal float64 + for k := range lineLen { + v := HalfToFloat64(src[base+k*stride+post]) + if v != v { // NaN is skipped as missing + continue + } + if best < 0 { + best, bestVal = k, v + continue + } + if (wantMin && v < bestVal) || (!wantMin && v > bestVal) { + best, bestVal = k, v + } + } + if best < 0 { + // Every element of the line is NaN. + allNaN.Store(true) + continue + } + out.ints[line*stride+post] = int64(best) + } + } + }) + case Float32: + src := a.floats32 + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for line := ls; line < le; line++ { + base := line * perLine + for post := range stride { + // Seed from the first non-NaN element; NaN candidates + // never enter the comparison. + best := -1 + var bestVal float64 + for k := range lineLen { + v := float64(src[base+k*stride+post]) + if v != v { // NaN is skipped as missing + continue + } + if best < 0 { + best, bestVal = k, v + continue + } + if (wantMin && v < bestVal) || (!wantMin && v > bestVal) { + best, bestVal = k, v + } + } + if best < 0 { + // Every element of the line is NaN. + allNaN.Store(true) + continue + } + out.ints[line*stride+post] = int64(best) + } + } + }) + default: + src := a.floats + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for line := ls; line < le; line++ { + base := line * perLine + for post := range stride { + // Seed from the first non-NaN element; NaN candidates + // never enter the comparison. + best := -1 + var bestVal float64 + for k := range lineLen { + v := src[base+k*stride+post] + if v != v { // NaN is skipped as missing + continue + } + if best < 0 { + best, bestVal = k, v + continue + } + if (wantMin && v < bestVal) || (!wantMin && v > bestVal) { + best, bestVal = k, v + } + } + if best < 0 { + // Every element of the line is NaN. + allNaN.Store(true) + continue + } + out.ints[line*stride+post] = int64(best) + } + } + }) + } + if allNaN.Load() { + return nil, errf("%s: every element along dimension %d is NaN", name, dim) + } + return out, nil +} + +// argExtremeAxisLine is argExtremeAxis's walk for a narrow integer +// payload: each line seeds from its first element and the candidates +// compare in the payload's own type, so no value ever meets a float +// comparison; the winning position lands in the Int result unchanged. +func argExtremeAxisLine[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int, out *Array, wantMin bool, lines, perLine, stride, lineLen int) { + splitCapped(lines, n, reduceSplitFloor, func(ls, le int) { + for line := ls; line < le; line++ { + base := line * perLine + for post := range stride { + best, bestVal := 0, src[base+post] + for k := 1; k < lineLen; k++ { + if v := src[base+k*stride+post]; wantMin { + if v < bestVal { + best, bestVal = k, v + } + } else if v > bestVal { + best, bestVal = k, v + } + } + out.ints[line*stride+post] = int64(best) + } + } + }) +} + +// argExtremeAxisLineBool is argExtremeAxisLine for a bool payload: Go +// orders no bool with < or >, so the strict improvement writes its own +// logic, and a tie keeps the earlier position. +func argExtremeAxisLineBool(src []bool, n int, out *Array, wantMin bool, lines, perLine, stride, lineLen int) { + splitCapped(lines, n, reduceSplitFloor, func(ls, le int) { + for line := ls; line < le; line++ { + base := line * perLine + for post := range stride { + best, bestVal := 0, src[base+post] + for k := 1; k < lineLen; k++ { + v := src[base+k*stride+post] + if wantMin && !v && bestVal { + best, bestVal = k, v + } else if !wantMin && v && !bestVal { + best, bestVal = k, v + } + } + out.ints[line*stride+post] = int64(best) + } + } + }) +} + +// TopK returns the top-k values and their original indices along the +// given dimension, sorted descending by value. The result preserves +// the input shape except the chosen dimension is reduced to k. NaN +// elements never rank: they sort to the end of the output when the +// line holds fewer than k finite values. For a 1-D input it returns +// two 1-D arrays of length k. +func TopK(a *Array, k int, dim int) (values, indices *Array, err error) { + if a.dt == Complex { + return nil, nil, errf("TopK: complex arrays have no ordering") + } + if dim < 0 || dim >= a.NDim() { + return nil, nil, errf("TopK: dimension %d out of range for shape %s", dim, shapeText(a.shape)) + } + if a.shape[dim] == 0 { + return nil, nil, errf("TopK: dimension %d of shape %s is empty", dim, shapeText(a.shape)) + } + if a.Len() == 0 { + // A dimension of size zero elsewhere leaves the per-line stride at + // zero and the line count undefined; the same empty answer + // argExtremeAxis gives. + return nil, nil, errf("TopK: an empty array has no result") + } + if k < 0 { + return nil, nil, errf("TopK: k must be non-negative, got %d", k) + } + if k > a.shape[dim] { + return nil, nil, errf("TopK: k=%d exceeds dimension %d size %d", k, dim, a.shape[dim]) + } + outShape := a.Shape() + outShape[dim] = k + vals := &Array{shape: outShape, dt: a.dt} + idxs := &Array{shape: outShape, dt: Int} + valsTotal := 1 + for _, d := range outShape { + valsTotal *= d + } + vals.alloc(valsTotal) + idxs.alloc(valsTotal) + stride := 1 + for kk := a.NDim() - 1; kk > dim; kk-- { + stride *= a.shape[kk] + } + perLine := a.shape[dim] * stride + totalLines := a.Len() / perLine + n := a.shape[dim] + // The candidate positions are the same identity list for every line + // and the loop only reads them, so one shared snapshot serves all + // workers. + idxSnap := make([]int, n) + for i := range n { + idxSnap[i] = i + } + // The split counts element visits, not elements: each line is + // snapshotted once and then scanned k election rounds, so the fold + // touches every element about k+1 times and the spawn floor applies + // to that count. + splitCapped(totalLines, a.Len()*(k+1), reduceSplitFloor, func(ls, le int) { + // Per-worker scratch reused across lines: the value snapshots are + // rewritten in full each pass and the elected flags cleared, so + // no state leaks between lines. + valsSnap := make([]float64, n) + intsSnap := make([]int64, n) + used := make([]bool, n) + // Repeated argmax costs k full scans per line, which is quadratic + // when k approaches n. Wide requests switch to one sort of the + // line under the same total order the elections implement: + // value descending, ties by first occurrence, NaN last. + sortPath := k*8 > n + var pairs []topkPair + if sortPath { + pairs = make([]topkPair, n) + } + for line := ls; line < le; line++ { + baseFlat := line * perLine + for post := range stride { + // Snapshot the line from the raw payload: the dtype + // dispatch sits outside the walk, float32 widens to + // float64 exactly, and an integer-class line is + // snapshotted as int64 because the float64 detour would + // round neighbours above 2^53 together (see + // argExtreme). Every narrow widening into int64 is + // exact, so the int64 ranking carries each payload's own + // order whatever the width. + switch a.dt { + case Int: + src := a.ints + for i := range n { + intsSnap[i] = src[baseFlat+i*stride+post] + } + case Bool: + src := a.bools + for i := range n { + intsSnap[i] = boolToInt64(src[baseFlat+i*stride+post]) + } + case Int8: + topkSnapNarrow(a.i8s, intsSnap, baseFlat, stride, post) + case Uint8: + topkSnapNarrow(a.u8s, intsSnap, baseFlat, stride, post) + case Int16: + topkSnapNarrow(a.i16s, intsSnap, baseFlat, stride, post) + case Uint16: + topkSnapNarrow(a.u16s, intsSnap, baseFlat, stride, post) + case Int32: + topkSnapNarrow(a.i32s, intsSnap, baseFlat, stride, post) + case Uint32: + topkSnapNarrow(a.u32s, intsSnap, baseFlat, stride, post) + case Float16: + src := a.halves + for i := range n { + valsSnap[i] = HalfToFloat64(src[baseFlat+i*stride+post]) + } + case Float32: + src := a.floats32 + for i := range n { + valsSnap[i] = float64(src[baseFlat+i*stride+post]) + } + default: + src := a.floats + for i := range n { + valsSnap[i] = src[baseFlat+i*stride+post] + } + } + if sortPath { + topkSortLine(a.dt, valsSnap, intsSnap, idxSnap, pairs, k, + vals, idxs, line, stride, post) + continue + } + clear(used) + // Repeated argmax over the finite values: NaN candidates + // are skipped, so they can never be elected, and a line + // with fewer than k finite values fills its remaining + // slots with NaN. + for outK := range k { + bestIdx := -1 + var bestVal float64 + var bestInt int64 + for i := range n { + if used[i] { + continue + } + if intClass(a.dt) { + // Integer-class candidates compare in exact + // int64, never float64, and an integer is + // never NaN: the first one always seeds. + if v := intsSnap[i]; bestIdx < 0 || v > bestInt { + bestIdx, bestInt = i, v + } + continue + } + if v := valsSnap[i]; v == v && (bestIdx < 0 || v > bestVal) { + bestIdx = i + bestVal = v + } + } + // The output line keeps the (k, suffix) row-major + // order: the reduced dimension carries stride, not + // the suffix. + outIdx := line*k*stride + outK*stride + post + if bestIdx < 0 { + // No finite value left; the float payloads fill + // with NaN. An integer-class line always seeds, so + // its zero fill never surfaces. + if intClass(a.dt) { + topkStore(vals, outIdx, 0) + } else { + switch a.dt { + case Float16: + vals.halves[outIdx] = halfNaN + case Float32: + vals.floats32[outIdx] = float32(math.NaN()) + default: + vals.floats[outIdx] = math.NaN() + } + } + idxs.ints[outIdx] = 0 + continue + } + if intClass(a.dt) { + // The original payload element widened exactly, + // never a rounded float64 image. + topkStore(vals, outIdx, intsSnap[bestIdx]) + } else { + switch a.dt { + case Float16: + vals.halves[outIdx] = HalfFromFloat64(valsSnap[bestIdx]) + case Float32: + vals.floats32[outIdx] = float32(valsSnap[bestIdx]) + default: + vals.floats[outIdx] = valsSnap[bestIdx] + } + } + idxs.ints[outIdx] = int64(idxSnap[bestIdx]) + used[bestIdx] = true + } + } + } + }) + return vals, idxs, nil +} + +// topkPair is one line element of the sort-based TopK path. +type topkPair struct { + val float64 + ival int64 + idx int +} + +// topkSnapNarrow snapshots one narrow integer line into the int64 +// snapshot the elections rank: every widening is exact, so the int64 +// comparison carries the payload's own order whatever the width. +func topkSnapNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, intsSnap []int64, baseFlat, stride, post int) { + for i := range intsSnap { + intsSnap[i] = int64(src[baseFlat+i*stride+post]) + } +} + +// topkStore writes an integer-class TopK value back into the values +// array's own payload: the implicit-store cast the ladder carries. +func topkStore(vals *Array, outIdx int, v int64) { + switch vals.dt { + case Int: + vals.ints[outIdx] = v + case Bool: + vals.bools[outIdx] = v != 0 + case Int8: + vals.i8s[outIdx] = int8(v) + case Uint8: + vals.u8s[outIdx] = uint8(v) + case Int16: + vals.i16s[outIdx] = int16(v) + case Uint16: + vals.u16s[outIdx] = uint16(v) + case Int32: + vals.i32s[outIdx] = int32(v) + default: + vals.u32s[outIdx] = uint32(v) + } +} + +// topkSortLine elects the top-k of one line by sorting it under the +// elections' own total order: value descending, ties by first +// occurrence, NaN last. The output slots, the NaN fill and the index +// reporting match the repeated-argmax path exactly, so the two paths +// are interchangeable for any k. +func topkSortLine(dt Dtype, valsSnap []float64, intsSnap []int64, idxSnap []int, pairs []topkPair, k int, + vals, idxs *Array, line, stride, post int) { + n := len(pairs) + for i := range n { + pairs[i] = topkPair{val: valsSnap[i], ival: intsSnap[i], idx: idxSnap[i]} + } + slices.SortFunc(pairs, func(a, b topkPair) int { + if intClass(dt) { + // The exact int64 snapshot: every narrow widening preserves + // the payload's own order, so this is the native comparison. + if c := cmp.Compare(b.ival, a.ival); c != 0 { + return c + } + return cmp.Compare(a.idx, b.idx) + } + aNaN, bNaN := math.IsNaN(a.val), math.IsNaN(b.val) + switch { + case aNaN && bNaN: + return cmp.Compare(a.idx, b.idx) + case aNaN: + return 1 + case bNaN: + return -1 + } + if c := cmp.Compare(b.val, a.val); c != 0 { + return c + } + return cmp.Compare(a.idx, b.idx) + }) + for outK := range k { + outIdx := line*k*stride + outK*stride + post + if outK < n && (intClass(dt) || !math.IsNaN(pairs[outK].val)) { + if intClass(dt) { + topkStore(vals, outIdx, pairs[outK].ival) + } else { + switch dt { + case Float16: + vals.halves[outIdx] = HalfFromFloat64(pairs[outK].val) + case Float32: + vals.floats32[outIdx] = float32(pairs[outK].val) + default: + vals.floats[outIdx] = pairs[outK].val + } + } + idxs.ints[outIdx] = int64(pairs[outK].idx) + continue + } + // Fewer than k finite values: the float payloads fill with NaN; + // an integer-class line always seeds, so its zero fill never + // surfaces. + if intClass(dt) { + topkStore(vals, outIdx, 0) + } else { + switch dt { + case Float16: + vals.halves[outIdx] = halfNaN + case Float32: + vals.floats32[outIdx] = float32(math.NaN()) + default: + vals.floats[outIdx] = math.NaN() + } + } + idxs.ints[outIdx] = 0 + } +} diff --git a/internal/core/axis_accuracy_test.go b/internal/core/axis_accuracy_test.go new file mode 100644 index 0000000..76d1e0f --- /dev/null +++ b/internal/core/axis_accuracy_test.go @@ -0,0 +1,345 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The axis sums fold each line through the canonical partition the +// global sums use: fixed blocks of the line, the block partials combined +// through the balanced tree. The claims pinned here against big.Float at +// 200 bits: on lines long enough for a chain's roundings to pile up the +// tree holds its accuracy where the chain drifts; on the large-plus-small +// counterpoint the chain's running total swallows the small elements +// outright and the tree's blocks keep them; the partition follows from +// the line length alone, so no worker count moves a bit; and a +// single-line fold answers Sum's own bits. + +// axisChain is the shape the axis fold replaced: one accumulator per +// line, addends in ascending element order. +func axisChain(v []float64) float64 { + var s float64 + for _, x := range v { + s += x + } + return s +} + +func TestAxisSumAccuracyAgainstBigFloat(t *testing.T) { + for _, n := range []int{1 << 13, 1 << 16, 1 << 20} { + v := foldFixture(n) + a, err := FromFloats(v, 1, n) + if err != nil { + t.Fatal(err) + } + got, err := SumAxis(a, 1) + if err != nil { + t.Fatal(err) + } + want := bigSum(v) + axisErr := relErr(got.FloatAt(0), want) + chainErr := relErr(axisChain(v), want) + t.Logf("n=%d: axis fold %.3e, one chain %.3e", n, axisErr, chainErr) + // A single contiguous line is one canonical fold: the global + // Sum over the same elements must give the same bits. + one, err := FromFloats(v, n) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(got.FloatAt(0)) != math.Float64bits(Sum(one).Float()) { + t.Fatalf("n=%d: the single-line axis fold %v disagrees with Sum %v", + n, got.FloatAt(0), Sum(one).Float()) + } + // The tree never rounds an order of magnitude past the chain + // where the chain happens to win by luck. + if axisErr > 4*chainErr && axisErr > 1e-16 { + t.Errorf("n=%d: the axis fold rounds an order worse than the chain: %.3e against %.3e", + n, axisErr, chainErr) + } + } +} + +// TestAxisSumCounterpoint pins the large-plus-small line: one element at +// 1e16 and the rest ones. The chain's running total sits on the 1e16 +// grid, where an ulp is 2, and swallows every one it meets; the tree's +// blocks sum the ones among themselves before they meet the giant. The +// exact referent is the big.Float sum. +func TestAxisSumCounterpoint(t *testing.T) { + const n = 1 << 13 + v := make([]float64, n) + for i := range v { + v[i] = 1 + } + v[0] = 1e16 + want := bigSum(v) + a, err := FromFloats(v, 1, n) + if err != nil { + t.Fatal(err) + } + got, err := SumAxis(a, 1) + if err != nil { + t.Fatal(err) + } + w, _ := want.Float64() + chainLoss := math.Abs(axisChain(v) - w) + axisLoss := math.Abs(got.FloatAt(0) - w) + t.Logf("counterpoint: axis fold loses %.0f, one chain loses %.0f", axisLoss, chainLoss) + // The chain drops all 8191 ones; the tree's three clean blocks keep + // three quarters of them and lose only the ones sharing the giant's + // own block. + if axisLoss >= chainLoss { + t.Fatalf("the axis fold loses %.0f against the chain's %.0f on the counterpoint", axisLoss, chainLoss) + } +} + +// TestAxisSumStridedAccuracyAgainstBigFloat holds the strided walk +// (reducing the leading dimension) against the same reference: the +// stride must not reintroduce a chain. +func TestAxisSumStridedAccuracyAgainstBigFloat(t *testing.T) { + const rows, cols = 8, 1 << 13 + v := foldFixture(rows * cols) + a, err := FromFloats(v, rows, cols) + if err != nil { + t.Fatal(err) + } + got, err := SumAxis(a, 0) + if err != nil { + t.Fatal(err) + } + for c := range cols { + // The strided fold and the packed fold share one canonical + // partition, so they must agree bit for bit. + column := make([]float64, rows) + scale := 0.0 + for r := range rows { + column[r] = v[r*cols+c] + scale = math.Max(scale, math.Abs(column[r])) + } + packed, err := FromFloats(column, 1, rows) + if err != nil { + t.Fatal(err) + } + want, err := SumAxis(packed, 1) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(got.FloatAt(c)) != math.Float64bits(want.FloatAt(0)) { + t.Fatalf("column %d: the strided fold %v disagrees with the packed fold %v", + c, got.FloatAt(c), want.FloatAt(0)) + } + // The reference check is absolute against the addend scale: a + // cancelling column makes the relative error meaningless. + acc := new(big.Float).SetPrec(200) + for r := range rows { + acc.Add(acc, new(big.Float).SetPrec(200).SetFloat64(column[r])) + } + w, _ := acc.Float64() + if d := math.Abs(got.FloatAt(c) - w); d > 1e-13*scale { + t.Fatalf("column %d: the strided fold is %.3e from the reference (scale %.3g)", c, d, scale) + } + } + // The same bits under a different worker count: the strided lines + // split across workers the same way the contiguous ones do. + prev := engine.SetNumWorkers(1) + defer engine.SetNumWorkers(prev) + one, err := SumAxis(a, 0) + if err != nil { + t.Fatal(err) + } + for _, w := range []int{2, 8, 32} { + engine.SetNumWorkers(w) + many, err := SumAxis(a, 0) + if err != nil { + t.Fatal(err) + } + for c := range cols { + if math.Float64bits(many.FloatAt(c)) != math.Float64bits(one.FloatAt(c)) { + t.Fatalf("workers=%d moved strided column %d", w, c) + } + } + } +} + +// TestAxisSumMeanDeterministic pins the fixed partition across worker +// counts for the fold and the mean, contiguous and strided, float64 and +// float32. +func TestAxisSumMeanDeterministic(t *testing.T) { + prev := engine.SetNumWorkers(1) + defer engine.SetNumWorkers(prev) + build := func() (*Array, *Array) { + v := foldFixture(16 * (1 << 13)) + contig, err := FromFloats(v, 16, 1<<13) + if err != nil { + t.Fatal(err) + } + strid, err := FromFloats(v, 1<<13, 16) + if err != nil { + t.Fatal(err) + } + return contig, strid + } + contig, strid := build() + snapshot := func(a *Array, dim int, mean bool) []float64 { + var out *Array + var err error + if mean { + out, err = MeanAxis(a, dim) + } else { + out, err = SumAxis(a, dim) + } + if err != nil { + t.Fatal(err) + } + vals := make([]float64, out.Len()) + copy(vals, out.floats) + return vals + } + for _, mean := range []bool{false, true} { + cOne := snapshot(contig, 1, mean) + sOne := snapshot(strid, 0, mean) + for _, w := range []int{2, 7, 32} { + engine.SetNumWorkers(w) + for i, got := range snapshot(contig, 1, mean) { + if math.Float64bits(got) != math.Float64bits(cOne[i]) { + t.Fatalf("workers=%d mean=%v moved contiguous slot %d", w, mean, i) + } + } + for i, got := range snapshot(strid, 0, mean) { + if math.Float64bits(got) != math.Float64bits(sOne[i]) { + t.Fatalf("workers=%d mean=%v moved strided slot %d", w, mean, i) + } + } + } + } +} + +// TestAxisMeanMatchesMean pins the mean's consistency: a single-line +// MeanAxis divides the canonical fold by the line length, which is +// Mean's own computation. +func TestAxisMeanMatchesMean(t *testing.T) { + const n = 1 << 13 + v := foldFixture(n) + a, err := FromFloats(v, 1, n) + if err != nil { + t.Fatal(err) + } + got, err := MeanAxis(a, 1) + if err != nil { + t.Fatal(err) + } + one, err := FromFloats(v, n) + if err != nil { + t.Fatal(err) + } + m, err := Mean(one) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(got.FloatAt(0)) != math.Float64bits(m) { + t.Fatalf("MeanAxis %.17g against Mean %.17g", got.FloatAt(0), m) + } + // And the mean keeps the reference's accuracy: the fold error + // divided by n. + want := new(big.Float).SetPrec(200).Quo(bigSum(v), new(big.Float).SetPrec(200).SetFloat64(n)) + if d := relErr(got.FloatAt(0), want); d > 1e-15 { + t.Fatalf("MeanAxis rounds %.3e from the big.Float mean", d) + } +} + +// TestAxisSumComplexAgainstBigFloat pins the complex fold: the real and +// imaginary parts accumulate through the same canonical partition, so +// each part sits at the reference's floor. +func TestAxisSumComplexAgainstBigFloat(t *testing.T) { + const n = 1 << 13 + re := foldFixture(n) + im := foldFixture(n / 2) + im = append(im, im...) + c := make([]complex128, n) + for i := range c { + c[i] = complex(re[i], im[i]) + } + a, err := FromComplexes(c, 1, n) + if err != nil { + t.Fatal(err) + } + got, err := SumAxis(a, 1) + if err != nil { + t.Fatal(err) + } + v := got.complexAt(0) + for _, part := range []struct { + name string + got float64 + want *big.Float + chain float64 + }{ + {"real", real(v), bigSum(re), axisChain(re)}, + {"imaginary", imag(v), bigSum(im), axisChain(im)}, + } { + axisErr, chainErr := relErr(part.got, part.want), relErr(part.chain, part.want) + t.Logf("%s line: axis fold %.3e, one chain %.3e", part.name, axisErr, chainErr) + if axisErr > 4*chainErr && axisErr > 1e-16 { + t.Fatalf("the %s part rounds an order worse than the chain: %.3e against %.3e", + part.name, axisErr, chainErr) + } + } +} + +// TestAxisSumWidenedAccuracy pins the float32 line. Every widening is +// exact, so the scratch fold sums exactly the values an accessor walk +// would; the accumulator then narrows to float32, which caps the +// published answer at the dtype's own resolution. Pinned here: the +// published value is the narrowed scratch fold, and the scratch fold +// itself never rounds an order past the chain the fold replaced. +func TestAxisSumWidenedAccuracy(t *testing.T) { + const n = 1 << 16 + fv := foldFixture(n) + f32 := make([]float32, n) + for i, x := range fv { + f32[i] = float32(x) + } + a32, err := FromFloat32s(f32, 1, n) + if err != nil { + t.Fatal(err) + } + got, err := SumAxis(a32, 1) + if err != nil { + t.Fatal(err) + } + want := new(big.Float).SetPrec(200) + for _, x := range f32 { + want.Add(want, new(big.Float).SetPrec(200).SetFloat64(float64(x))) + } + // The scratch fold is the canonical one; the published value is its + // float32 narrowing. + scratch := floatFoldSumF32(f32) + if math.Float64bits(float64(float32(scratch))) != math.Float64bits(got.FloatAt(0)) { + t.Fatalf("the published float32 sum %.9g disagrees with the narrowed scratch fold %.9g", + got.FloatAt(0), float32(scratch)) + } + // One float32 ulp of the reference bounds the published answer. + w, _ := want.Float64() + ulp := math.Abs(w) * 1.19e-7 + if d := math.Abs(float64(got.FloatAt(0)) - w); d > 2*ulp { + t.Fatalf("the published float32 sum is %.3e from the reference, above two float32 ulp (%.3e)", d, ulp) + } + // And the scratch fold holds the chain's level on the long line. + var chain float64 + for _, x := range f32 { + chain += float64(x) + } + scratchErr, chainErr := relErr(scratch, want), relErr(chain, want) + // Both forms sit ten orders below the float32 resolution the + // published answer carries, so which of them is luckier on one + // length is noise: pin only the level, not the race. + t.Logf("float32 scratch fold %.3e, one chain %.3e", scratchErr, chainErr) + if scratchErr > 1e-13 { + t.Fatalf("the float32 scratch fold rounds %.3e from the reference, far past the dtype's floor", scratchErr) + } +} diff --git a/internal/core/axis_bench_test.go b/internal/core/axis_bench_test.go new file mode 100644 index 0000000..231be8d --- /dev/null +++ b/internal/core/axis_bench_test.go @@ -0,0 +1,132 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Along-dimension axis reductions: the softmax, normalisation and loss +// primitives. Run with -bench before and after any change to axis.go. +// The 2-D shapes line up with BenchmarkNorm2D and BenchmarkProd2D; the +// 3-D 32×17×9 shape exercises every dim position against a +// non-power-of-two stride. + +// benchFloat2D builds an n×n float64 array with a deterministic mix of +// small values. +func benchFloat2D(b *testing.B, n int) *Array { + b.Helper() + a, _ := FromFloats(make([]float64, n*n), n, n) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%9) - 4 + } + return a +} + +// benchFloat3D builds a 32×17×9 float64 array with a deterministic mix +// of small values. +func benchFloat3D(b *testing.B) *Array { + b.Helper() + a, _ := FromFloats(make([]float64, 32*17*9), 32, 17, 9) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%11) - 5 + } + return a +} + +func benchAxis1(b *testing.B, f func(*Array, int) (*Array, error)) { + b.Helper() + a := benchFloat2D(b, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := f(a, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSumAxis2D(b *testing.B) { benchAxis1(b, SumAxis) } +func BenchmarkMeanAxis2D(b *testing.B) { benchAxis1(b, MeanAxis) } +func BenchmarkMinAxis2D(b *testing.B) { benchAxis1(b, MinAxis) } +func BenchmarkMaxAxis2D(b *testing.B) { benchAxis1(b, MaxAxis) } +func BenchmarkArgMaxAxis2D(b *testing.B) { benchArgAxis1(b, false) } +func BenchmarkArgMinAxis2D(b *testing.B) { benchArgAxis1(b, true) } + +func benchArgAxis1(b *testing.B, wantMin bool) { + b.Helper() + a := benchFloat2D(b, 512) + b.ReportAllocs() + for b.Loop() { + var err error + if wantMin { + _, err = ArgMinAxis(a, 1) + } else { + _, err = ArgMaxAxis(a, 1) + } + if err != nil { + b.Fatal(err) + } + } +} + +// The 3-D cases walk every dim position: dim 2 is the contiguous +// trailing run, dim 0 the coarsest split, dim 1 the interior stride. +func BenchmarkSumAxis3DDim0(b *testing.B) { benchAxis3D(b, SumAxis, 0) } +func BenchmarkSumAxis3DDim1(b *testing.B) { benchAxis3D(b, SumAxis, 1) } +func BenchmarkSumAxis3DDim2(b *testing.B) { benchAxis3D(b, SumAxis, 2) } +func BenchmarkMinAxis3DDim0(b *testing.B) { benchAxis3D(b, MinAxis, 0) } +func BenchmarkMinAxis3DDim1(b *testing.B) { benchAxis3D(b, MinAxis, 1) } +func BenchmarkMinAxis3DDim2(b *testing.B) { benchAxis3D(b, MinAxis, 2) } + +func benchAxis3D(b *testing.B, f func(*Array, int) (*Array, error), dim int) { + b.Helper() + a := benchFloat3D(b) + b.ReportAllocs() + for b.Loop() { + if _, err := f(a, dim); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSumAxis2DInt covers the int payload of SumAxis; the int path +// folds straight into the accumulator with no float scratch. +func BenchmarkSumAxis2DInt(b *testing.B) { + a, _ := FromInts(make([]int64, 512*512), 512, 512) + for i := range a.RawInts() { + a.RawInts()[i] = int64(i%9) - 4 + } + b.ReportAllocs() + for b.Loop() { + if _, err := SumAxis(a, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSumAxis2DFloat32 covers the float32 scratch path: folds go +// through a float64 scratch and round once at the end. +func BenchmarkSumAxis2DFloat32(b *testing.B) { + af, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range af.RawFloats() { + af.RawFloats()[i] = float64(i%9) - 4 + } + a, _ := Astype(af, Float32) + b.ReportAllocs() + for b.Loop() { + if _, err := SumAxis(a, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkTopK2D covers the repeated-argmax selection along the +// trailing dimension, the loss and beam-search hot path. +func BenchmarkTopK2D(b *testing.B) { + a := benchFloat2D(b, 512) + b.ReportAllocs() + for b.Loop() { + if _, _, err := TopK(a, 16, 1); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/axis_test.go b/internal/core/axis_test.go new file mode 100644 index 0000000..3482815 --- /dev/null +++ b/internal/core/axis_test.go @@ -0,0 +1,266 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "slices" + "strings" + "testing" +) + +func TestSumAxis(t *testing.T) { + m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + + // Column sums: shape (2,) gives [6, 15]. + cols, err := SumAxis(m, 1) + if err != nil { + t.Fatalf("SumAxis(1): %v", err) + } + want := mustFromInts(t, []int64{6, 15}, 2) + if !Equal(want, cols) { + t.Fatalf("SumAxis(1): %s", cols) + } + + // Row sums: shape (3,) gives [5, 7, 9]. + rows, err := SumAxis(m, 0) + if err != nil { + t.Fatalf("SumAxis(0): %v", err) + } + wantRows := mustFromInts(t, []int64{5, 7, 9}, 3) + if !Equal(wantRows, rows) { + t.Fatalf("SumAxis(0): %s", rows) + } + + // Float and complex keep their dtypes. + f := mustFromFloats(t, []float64{0.5, 1.5, 2.5, 3.5}, 2, 2) + fs, _ := SumAxis(f, 0) + if fs.Dtype() != Float { + t.Fatalf("SumAxis float dtype: %s", fs.Dtype()) + } + c := mustFromComplexes(t, []complex128{1, complex(0, 1), 0, 0}, 2, 2) + cs, err := SumAxis(c, 1) + if err != nil { + t.Fatalf("SumAxis complex: %v", err) + } + if v, _ := ComplexAt(cs, 0); v != complex(1, 1) { + t.Fatalf("SumAxis complex value: %v", v) + } + + if _, err := SumAxis(m, 2); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("SumAxis dim: %v", err) + } + v := mustFromInts(t, []int64{1, 2}, 2) + if _, err := SumAxis(v, 0); err == nil || !strings.Contains(err.Error(), "global variant") { + t.Fatalf("SumAxis 1-D: %v", err) + } +} + +func TestMinMaxAxis(t *testing.T) { + m := mustFromInts(t, []int64{1, 5, 2, 6, 4, 3}, 2, 3) + + mn, err := MinAxis(m, 1) + if err != nil { + t.Fatalf("MinAxis: %v", err) + } + if !Equal(mustFromInts(t, []int64{1, 3}, 2), mn) { + t.Fatalf("MinAxis: %s", mn) + } + mx, err := MaxAxis(m, 1) + if err != nil { + t.Fatalf("MaxAxis: %v", err) + } + if !Equal(mustFromInts(t, []int64{5, 6}, 2), mx) { + t.Fatalf("MaxAxis: %s", mx) + } + + // Along rows. + rowsMin, _ := MinAxis(m, 0) + if !Equal(mustFromInts(t, []int64{1, 4, 2}, 3), rowsMin) { + t.Fatalf("MinAxis(0): %s", rowsMin) + } + + // A zero start must never win: all-negative values. + neg := mustFromInts(t, []int64{-5, -1, -9, -2}, 2, 2) + nmin, _ := MinAxis(neg, 1) + if !Equal(mustFromInts(t, []int64{-5, -9}, 2), nmin) { + t.Fatalf("MinAxis negatives: %s", nmin) + } + nmax, _ := MaxAxis(neg, 0) + if !Equal(mustFromInts(t, []int64{-5, -1}, 2), nmax) { + t.Fatalf("MaxAxis negatives: %s", nmax) + } + + // A NaN element never wins: a line starting with NaN takes its + // first finite element, and a line that is all NaN has no extreme + // and reports NaN. + f := mustFromFloats(t, []float64{math.NaN(), 1.0, 3.0}, 1, 3) + fmin, err := MinAxis(f, 1) + if err != nil { + t.Fatalf("MinAxis NaN: %v", err) + } + if v, _ := FloatAt(fmin, 0); v != 1.0 { + t.Fatalf("MinAxis NaN never wins: got %v, want 1", v) + } + fmax, err := MaxAxis(f, 1) + if err != nil { + t.Fatalf("MaxAxis NaN: %v", err) + } + if v, _ := FloatAt(fmax, 0); v != 3.0 { + t.Fatalf("MaxAxis NaN never wins: got %v, want 3", v) + } + allNaN := mustFromFloats(t, []float64{math.NaN(), math.NaN(), math.NaN(), 2.0}, 2, 2) + anmin, _ := MinAxis(allNaN, 1) + if v, _ := FloatAt(anmin, 0); !math.IsNaN(v) { + t.Fatalf("MinAxis all-NaN line: got %v, want NaN", v) + } + if v, _ := FloatAt(anmin, 1); v != 2.0 { + t.Fatalf("MinAxis all-NaN neighbour: got %v, want 2", v) + } + + c := mustFromComplexes(t, []complex128{1, 2, 3, 4}, 2, 2) + if _, err := MinAxis(c, 0); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("MinAxis complex: %v", err) + } +} + +func TestMeanAxis(t *testing.T) { + m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + + mean, err := MeanAxis(m, 1) + if err != nil { + t.Fatalf("MeanAxis: %v", err) + } + if mean.Dtype() != Float { + t.Fatalf("MeanAxis dtype: %s", mean.Dtype()) + } + want := []float64{2.0, 5.0} // (1+2+3)/3, (4+5+6)/3 + for i := range 2 { + if v, _ := FloatAt(mean, i); v != want[i] { + t.Fatalf("MeanAxis[%d]: %v", i, v) + } + } + + byRows, _ := MeanAxis(m, 0) + wantRows := []float64{2.5, 3.5, 4.5} // column means of [[1,2,3],[4,5,6]] + for i := range 3 { + if v, _ := FloatAt(byRows, i); v != wantRows[i] { + t.Fatalf("MeanAxis(0)[%d]: %v", i, v) + } + } + + c := mustFromComplexes(t, []complex128{1, 2, 3, 4}, 2, 2) + if _, err := MeanAxis(c, 0); err == nil || !strings.Contains(err.Error(), "no float mean") { + t.Fatalf("MeanAxis complex: %v", err) + } +} + +func TestArgExtreme(t *testing.T) { + a := mustFromInts(t, []int64{3, 1, 2}, 3) + imax, err := ArgMax(a) + if err != nil || imax != 0 { + t.Fatalf("ArgMax: %d %v", imax, err) + } + imin, err := ArgMin(a) + if err != nil || imin != 1 { + t.Fatalf("ArgMin: %d %v", imin, err) + } + + // NaN elements are skipped; a finite one still wins. + f := mustFromFloats(t, []float64{math.NaN(), 2.0, math.NaN(), 1.0}, 4) + fmin, err := ArgMin(f) + if err != nil || fmin != 3 { + t.Fatalf("ArgMin NaN skip: %d %v", fmin, err) + } + + allNaN := mustFromFloats(t, []float64{math.NaN()}, 1) + if _, err := ArgMax(allNaN); err == nil || !strings.Contains(err.Error(), "every element is NaN") { + t.Fatalf("ArgMax all NaN: %v", err) + } + + m := mustFromInts(t, []int64{1, 2}, 1, 2) + if _, err := ArgMax(m); err == nil || !strings.Contains(err.Error(), "needs a 1-D array") { + t.Fatalf("ArgMax 2-D: %v", err) + } + c := mustFromComplexes(t, []complex128{1}, 1) + if _, err := ArgMin(c); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("ArgMin complex: %v", err) + } +} + +// TestAxisReductionsNarrowDtypes pins the axis rules the scalar +// reductions carry: integer-class axis sums answer Int-axis results +// under exact widening, means answer float results, extrema compare +// natively in each payload type, and the arg extremes keep their Int +// index contract. +func TestAxisReductionsNarrowDtypes(t *testing.T) { + i8, err := FromInt8s([]int8{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + sum, err := SumAxis(i8, 1) + if err != nil { + t.Fatalf("SumAxis int8: %v", err) + } + if sum.Dtype() != Int { + t.Fatalf("SumAxis int8 answered %s, want the Int axis result", sum.Dtype()) + } + if want := []int64{6, 15}; !slices.Equal(sum.RawInts(), want) { + t.Fatalf("SumAxis int8 = %v, want %v", sum.RawInts(), want) + } + mean, err := MeanAxis(i8, 1) + if err != nil || mean.Dtype() != Float { + t.Fatalf("MeanAxis int8: %s %v", mean.Dtype(), err) + } + if want := []float64{2, 5}; !slices.Equal(mean.RawFloats(), want) { + t.Fatalf("MeanAxis int8 = %v, want %v", mean.RawFloats(), want) + } + mn, err := MinAxis(i8, 1) + if err != nil || mn.Dtype() != Int || !slices.Equal(mn.RawInts(), []int64{1, 4}) { + t.Fatalf("MinAxis int8 = %s %v %v", mn.Dtype(), mn.RawInts(), err) + } + mx, err := MaxAxis(i8, 1) + if err != nil || mx.Dtype() != Int || !slices.Equal(mx.RawInts(), []int64{3, 6}) { + t.Fatalf("MaxAxis int8 = %s %v %v", mx.Dtype(), mx.RawInts(), err) + } + + // Bool axis sums count trues per line into the Int axis result. + bl, err := FromBools([]bool{true, false, true, true}, 2, 2) + if err != nil { + t.Fatal(err) + } + bs, err := SumAxis(bl, 1) + if err != nil || bs.Dtype() != Int || !slices.Equal(bs.RawInts(), []int64{1, 2}) { + t.Fatalf("SumAxis bool = %s %v %v", bs.Dtype(), bs.RawInts(), err) + } + bmn, err := MinAxis(bl, 1) + if err != nil || !slices.Equal(bmn.RawInts(), []int64{0, 1}) { + t.Fatalf("MinAxis bool = %v %v", bmn.RawInts(), err) + } + + // The arg extremes compare natively per payload dtype and keep the + // Int index contract. + i16, err := FromInt16s([]int16{1, 9, 3, 8, 2, 7}, 2, 3) + if err != nil { + t.Fatal(err) + } + amax, err := ArgMaxAxis(i16, 1) + if err != nil || amax.Dtype() != Int || !slices.Equal(amax.RawInts(), []int64{1, 0}) { + t.Fatalf("ArgMaxAxis int16 = %s %v %v", amax.Dtype(), amax.RawInts(), err) + } + u32, err := FromUint32s([]uint32{1, 5, 3}, 3) + if err != nil { + t.Fatal(err) + } + if idx, err := ArgMax(u32); err != nil || idx != 1 { + t.Fatalf("ArgMax uint32 = %d %v, want 1", idx, err) + } + bl1, err := FromBools([]bool{true, false, true}, 3) + if err != nil { + t.Fatal(err) + } + if idx, err := ArgMin(bl1); err != nil || idx != 1 { + t.Fatalf("ArgMin bool = %d %v, want the first false at 1", idx, err) + } +} diff --git a/internal/core/basebridge.go b/internal/core/basebridge.go new file mode 100644 index 0000000..3717d09 --- /dev/null +++ b/internal/core/basebridge.go @@ -0,0 +1,36 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// The bridge delegates to base wholesale: the error prefix, the shape +// rendering and the range helper must not drift from what the domain +// packages see, so no local copy of them exists here. + +// epsF is the float64 machine epsilon, bridged from base.EpsF. +const epsF = base.EpsF + +// scalar is the element type the generic kernels serve. +type scalar = base.Scalar + +func errf(format string, args ...any) error { return base.Errf(format, args...) } + +func shapeText(shape []int) string { return base.ShapeText(shape) } + +func rangeN(n int) []int { return base.RangeN(n) } + +func factor[T scalar](m [][]T) ([]int, int) { return base.Factor(m) } + +func checkSingular[T scalar](name string, m [][]T) error { return base.CheckSingular(name, m) } + +func solveColumn[T scalar](m [][]T, col []T) { base.SolveColumn(m, col) } + +func permuteColumn[T scalar](col []T, perm []int) { base.PermuteColumn(col, perm) } + +func solveSystem[T scalar](name string, m [][]T, rhs [][]T) ([][]T, error) { + return base.SolveSystem(name, m, rhs) +} diff --git a/internal/core/bench_einsum2_test.go b/internal/core/bench_einsum2_test.go new file mode 100644 index 0000000..6902959 --- /dev/null +++ b/internal/core/bench_einsum2_test.go @@ -0,0 +1,639 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "slices" + "strconv" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The general engine's slot walk and the axis folds the reduction-only +// specs dispatch to: their benchmarks, and the oracles that pin the +// walk's bits. + +// einsumSlotSumOdometer steps every summed axis through one shared +// odometer. It is the oracle for the peeled kernel: a slot must hold +// exactly what this walk puts there, bit for bit. +func einsumSlotSumOdometer[T int64 | float64 | complex128](rv [][]T, out []T, slot int, base, off, coord, delta, sizes []int, total int) { + nSum := len(sizes) + copy(off, base) + clear(coord[:nSum]) + acc := T(1) + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] += acc + for range total - 1 { + t := nSum - 1 + for t >= 0 { + coord[t]++ + for p := range rv { + off[p] += delta[p*nSum+t] + } + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + for p := range rv { + off[p] -= delta[p*nSum+t] * sizes[t] + } + t-- + } + acc = T(1) + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] += acc + } +} + +// einsumSlotSumF32Odometer is einsumSlotSumOdometer for a float32 +// result, narrowing every addend into the slot as it arrives. +func einsumSlotSumF32Odometer(rv [][]float64, out []float32, slot int, base, off, coord, delta, sizes []int, total int) { + nSum := len(sizes) + copy(off, base) + clear(coord[:nSum]) + acc := 1.0 + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] = float32(float64(out[slot]) + acc) + for range total - 1 { + t := nSum - 1 + for t >= 0 { + coord[t]++ + for p := range rv { + off[p] += delta[p*nSum+t] + } + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + for p := range rv { + off[p] -= delta[p*nSum+t] * sizes[t] + } + t-- + } + acc = 1.0 + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] = float32(float64(out[slot]) + acc) + } +} + +// einsumProbeFloats fills n elements with distinct values of mixed +// magnitude: any visit order other than the reference's moves addends +// across a rounding boundary and changes the low bits of the sum. +func einsumProbeFloats(n int) []float64 { + v := make([]float64, n) + for i := range v { + v[i] = float64(i%11)*1.5 - 4 + math.Ldexp(1, i%17-8) + float64(i)/512 + } + return v +} + +// einsumSlotConfig is one slot-walk configuration: the axis extents, the +// per-operand cursor and the per-operand, per-axis strides. +type einsumSlotConfig struct { + sizes []int + base []int + delta []int + reach []int +} + +// configs walks the operand counts and the summed-axis counts, giving +// every kernel shape a configuration with in-bounds cursors; the +// zero-summed-axis shapes come last. +func einsumSlotConfigs() []einsumSlotConfig { + var out []einsumSlotConfig + for nOps := 1; nOps <= 5; nOps++ { + for nSum := 1; nSum <= 3; nSum++ { + cfg := einsumSlotConfig{ + sizes: make([]int, nSum), + base: make([]int, nOps), + delta: make([]int, nOps*nSum), + reach: make([]int, nOps), + } + for t := range nSum { + cfg.sizes[t] = 2 + (t*5+nOps)%3 + } + for p := range nOps { + cfg.base[p] = p + cfg.reach[p] = p + for t := range nSum { + cfg.delta[p*nSum+t] = 1 + (p*3+t)%4 + cfg.reach[p] += (cfg.sizes[t] - 1) * cfg.delta[p*nSum+t] + } + } + out = append(out, cfg) + } + } + for nOps := 1; nOps <= 4; nOps++ { + cfg := einsumSlotConfig{ + base: make([]int, nOps), + delta: make([]int, nOps), + reach: make([]int, nOps), + } + for p := range nOps { + cfg.base[p] = p + cfg.reach[p] = p + } + out = append(out, cfg) + } + return out +} + +// slotTotal is the number of visits one slot's walk takes. +func (c einsumSlotConfig) slotTotal() int { + total := 1 + for _, s := range c.sizes { + total *= s + } + return total +} + +// name spells the configuration for a failure message. +func (c einsumSlotConfig) name() string { + return "ops=" + strconv.Itoa(len(c.base)) + "/sum=" + strconv.Itoa(len(c.sizes)) +} + +// einsumCheckSlotKernel runs one configuration through the peeled +// kernel and the odometer reference for a payload type and compares the +// two slots by the given identity, so every operand count and cursor +// reaches both loops. +func einsumCheckSlotKernel[T int64 | float64 | complex128]( + t *testing.T, label string, cfg einsumSlotConfig, + mk func(reach []int) [][]T, same func(T, T) bool, +) { + t.Helper() + rv := mk(cfg.reach) + got := make([]T, 1) + want := make([]T, 1) + nSum := len(cfg.sizes) + einsumSlotSum(rv, got, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal()) + einsumSlotSumOdometer(rv, want, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal()) + if !same(got[0], want[0]) { + t.Fatalf("%s %s: slot %v, want %v", label, cfg.name(), got[0], want[0]) + } +} + +// einssumCheckSlotKernelF32 pins the float32 slot kernel, which narrows +// every addend into the slot, against its own odometer form. +func einssumCheckSlotKernelF32(t *testing.T, cfg einsumSlotConfig, mk func(reach []int) [][]float64) { + t.Helper() + rv := mk(cfg.reach) + got := make([]float32, 1) + want := make([]float32, 1) + nSum := len(cfg.sizes) + einsumSlotSumF32(rv, got, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal()) + einsumSlotSumF32Odometer(rv, want, 0, cfg.base, make([]int, len(cfg.base)), make([]int, nSum), cfg.delta, cfg.sizes, cfg.slotTotal()) + if math.Float32bits(got[0]) != math.Float32bits(want[0]) { + t.Fatalf("float32 %s: slot %v, want %v", cfg.name(), got[0], want[0]) + } +} + +// floatsPerReach builds one probe reader per operand, each long enough +// for that operand's whole walk. +func floatsPerReach(reach []int) [][]float64 { + rv := make([][]float64, len(reach)) + for p, r := range reach { + rv[p] = einsumProbeFloats(r + 1) + } + return rv +} + +// TestEinsumSlotSumMatchesOdometer pins the peeled slot kernels against +// the odometer form across operand counts, summed-axis counts, axis +// extents and cursors, on every payload type the kernels serve. The +// distinct probe values make a moved addend visible in the low bits. +func TestEinsumSlotSumMatchesOdometer(t *testing.T) { + configs := einsumSlotConfigs() + for _, cfg := range configs { + // The generic kernel, instantiated for each payload type. + einsumCheckSlotKernel(t, "float", cfg, floatsPerReach, + func(a, b float64) bool { return math.Float64bits(a) == math.Float64bits(b) }) + einsumCheckSlotKernel(t, "int", cfg, + func(reach []int) [][]int64 { + rv := make([][]int64, len(reach)) + for p, r := range reach { + rv[p] = make([]int64, r+1) + for i := range rv[p] { + rv[p][i] = int64(i%13)*7 - 9 + } + } + return rv + }, func(a, b int64) bool { return a == b }) + einsumCheckSlotKernel(t, "complex", cfg, + func(reach []int) [][]complex128 { + rv := make([][]complex128, len(reach)) + for p, r := range reach { + rv[p] = make([]complex128, r+1) + for i := range rv[p] { + rv[p][i] = complex(float64(i%7)-3, float64(i%5)-2) + } + } + return rv + }, func(a, b complex128) bool { + return math.Float64bits(real(a)) == math.Float64bits(real(b)) && + math.Float64bits(imag(a)) == math.Float64bits(imag(b)) + }) + einssumCheckSlotKernelF32(t, cfg, floatsPerReach) + } +} + +// einsumOperand builds a deterministic contiguous operand of the given +// dtype. +func einsumOperand(tb testing.TB, dt Dtype, shape ...int) *Array { + tb.Helper() + n := 1 + for _, d := range shape { + n *= d + } + var a *Array + var err error + switch dt { + case Int: + v := make([]int64, n) + for i := range v { + v[i] = int64(i%9)*3 - 4 + int64(i)*7 + } + a, err = FromInts(v, shape...) + case Float32: + v := make([]float32, n) + for i := range v { + v[i] = float32(i%7)*0.5 - 2 + float32(i)/64 + } + a, err = FromFloat32s(v, shape...) + case Complex: + v := make([]complex128, n) + for i := range v { + v[i] = complex(float64(i%5)-2, float64(i%3)-1) + } + a, err = FromComplexes(v, shape...) + default: + a, err = FromFloats(einsumProbeFloats(n), shape...) + } + if err != nil { + tb.Fatalf("operand %s %v: %v", dt, shape, err) + } + return a +} + +// einsumStridedOperand builds a read-only view of the given shape whose +// axes step by the given strides over a longer payload. +func einsumStridedOperand(t testing.TB, dt Dtype, shape, strides []int) *Array { + t.Helper() + n := 1 + for d, s := range shape { + n += (s - 1) * strides[d] + } + flat := einsumOperand(t, dt, n) + view := &Array{shape: slices.Clone(shape), dt: dt, strides: slices.Clone(strides)} + switch dt { + case Int: + view.ints = flat.RawInts() + case Float32: + view.floats32 = flat.RawFloat32s() + case Complex: + view.complexes = flat.RawComplexes() + default: + view.floats = flat.RawFloats() + } + return view +} + +// einsumBitsEqual compares two results by dtype, shape and raw payload +// bits, so a NaN matches only its own bit pattern. +func einsumBitsEqual(a, b *Array) bool { + if a.Dtype() != b.Dtype() || !slices.Equal(a.Shape(), b.Shape()) { + return false + } + switch a.dt { + case Int: + return slices.Equal(a.RawInts(), b.RawInts()) + case Float32: + return slices.EqualFunc(a.RawFloat32s(), b.RawFloat32s(), func(x, y float32) bool { + return math.Float32bits(x) == math.Float32bits(y) + }) + case Float: + return slices.EqualFunc(a.RawFloats(), b.RawFloats(), func(x, y float64) bool { + return math.Float64bits(x) == math.Float64bits(y) + }) + case Complex: + return slices.EqualFunc(a.RawComplexes(), b.RawComplexes(), func(x, y complex128) bool { + return math.Float64bits(real(x)) == math.Float64bits(real(y)) && + math.Float64bits(imag(x)) == math.Float64bits(imag(y)) + }) + } + return false +} + +// TestEinsumSlotSplitMatchesSerial runs the whole dispatch with the +// worker pool as it stands and with it pinned to one worker, which +// walks every slot in one chunk. The slot walk must write the same bits +// either way, for every operand count, dtype and stride pattern the +// engine accepts. +func TestEinsumSlotSplitMatchesSerial(t *testing.T) { + cases := []struct { + spec string + shapes [][]int + }{ + {"...ij,jk->...ik", [][]int{{4, 6, 5}, {5, 3}}}, + {"...ij,...jk->...ik", [][]int{{3, 4, 5}, {3, 5, 2}}}, + {"ik,kj,jl->il", [][]int{{6, 5}, {5, 4}, {4, 7}}}, + {"ijk->kji", [][]int{{3, 4, 5}}}, + {"ij,jk,kl->il", [][]int{{5, 4}, {4, 6}, {6, 3}}}, + {"ij,jk,lk->il", [][]int{{5, 4}, {4, 6}, {3, 6}}}, + {"i,...i->...", [][]int{{5}, {3, 5}}}, + {"...i,...i->...", [][]int{{2, 4}, {4}}}, + {"kji->k", [][]int{{3, 4, 5}}}, + {"iij->ij", [][]int{{4, 4, 3}}}, + {"ii->i", [][]int{{5, 5}}}, + {"...->...", [][]int{{3, 4}}}, + {"ij->i", [][]int{{7, 5}}}, + {"ab,bc->ac", [][]int{{6, 5}, {5, 4}}}, + {"ij,ij->", [][]int{{6, 5}, {6, 5}}}, + } + for _, tc := range cases { + for _, dt := range []Dtype{Int, Float32, Float, Complex} { + ops := make([]*Array, len(tc.shapes)) + for i, shape := range tc.shapes { + ops[i] = einsumOperand(t, dt, shape...) + } + want, werr := Einsum(tc.spec, ops...) + restore := engine.SetNumWorkers(1) + got, gerr := Einsum(tc.spec, ops...) + engine.SetNumWorkers(restore) + if werr != nil || gerr != nil { + t.Fatalf("%s/%s: split error %v, serial error %v", tc.spec, dt, werr, gerr) + } + if !einsumBitsEqual(got, want) { + t.Fatalf("%s/%s: split result %s differs from the serial walk %s", + tc.spec, dt, got, want) + } + } + } + // A strided operand is gathered through the accessor on every dtype + // the engine widens; the walk must read the view's own elements. + for _, dt := range []Dtype{Int, Float32, Float, Complex} { + ops := []*Array{ + einsumStridedOperand(t, dt, []int{3, 4}, []int{8, 2}), + einsumOperand(t, dt, 3, 4), + } + want, werr := Einsum("ij,ij->i", ops...) + restore := engine.SetNumWorkers(1) + got, gerr := Einsum("ij,ij->i", ops...) + engine.SetNumWorkers(restore) + if werr != nil || gerr != nil { + t.Fatalf("strided/%s: split error %v, serial error %v", dt, werr, gerr) + } + if !einsumBitsEqual(got, want) { + t.Fatalf("strided/%s: split result %s differs from the serial walk %s", dt, got, want) + } + } +} + +// TestEinsumSumAxesMatchesEngine pins the reduction-only dispatch to the +// general engine: "ij->i" and its relatives must fold to the same bits +// the engine's walk produced, for every shape, axis order and dtype. The +// float32 and complex operands stay with the engine by design and are +// compared here as well, so a later change to the dispatch cannot move +// them silently. +func TestEinsumSumAxesMatchesEngine(t *testing.T) { + shapes := map[int][]int{2: {6, 5}, 3: {4, 6, 5}, 4: {3, 4, 6, 2}} + specs := []string{ + "ij->i", "ij->j", + "ijk->ij", "ijk->ik", "ijk->ji", "ijk->i", "ijk->j", "ijk->k", + "kji->k", "kji->ki", "kji->jk", + "ijkl->il", "ijkl->lj", "ijkl->ji", "ijkl->i", "ijkl->l", + } + for _, spec := range specs { + lhsStr, rhsStr, ok := strings.Cut(spec, "->") + shape, ok2 := shapes[len(lhsStr)] + if !ok || !ok2 { + continue + } + for _, dt := range []Dtype{Int, Float32, Float, Complex} { + a := einsumOperand(t, dt, shape...) + want, werr := einsumGeneral([]string{lhsStr}, rhsStr, true, []*Array{a}) + got, gerr := Einsum(spec, a) + if werr != nil || gerr != nil { + t.Fatalf("%s/%s %v: dispatch error %v, engine error %v", spec, dt, shape, gerr, werr) + } + if !einsumBitsEqual(got, want) { + t.Fatalf("%s/%s %v: dispatched result %s differs from the engine %s", + spec, dt, shape, got, want) + } + } + } +} + +// einsumDenseOf copies a view's logical elements into a contiguous +// array of the same dtype. +func einsumDenseOf(t *testing.T, a *Array) *Array { + t.Helper() + n := a.Len() + var out *Array + var err error + switch a.dt { + case Int: + v := make([]int64, n) + for i := range v { + v[i] = a.ints[a.physIndex(i)] + } + out, err = FromInts(v, a.Shape()...) + case Float32: + v := make([]float32, n) + for i := range v { + v[i] = a.floats32[a.physIndex(i)] + } + out, err = FromFloat32s(v, a.Shape()...) + case Complex: + v := make([]complex128, n) + for i := range v { + v[i] = a.complexAt(i) + } + out, err = FromComplexes(v, a.Shape()...) + default: + v := make([]float64, n) + for i := range v { + v[i] = a.floatAt(i) + } + out, err = FromFloats(v, a.Shape()...) + } + if err != nil { + t.Fatalf("dense copy of %s: %v", a, err) + } + return out +} + +// TestEinsumStridedViewMatchesDense pins the engine's readers against a +// view's own elements: a strided operand must contract exactly as a +// dense array holding the same logical elements, for every dtype the +// engine widens. The patterns are the ones the engine owns; the table +// paths read payload windows and take contiguous operands only, as the +// Array contract states. +func TestEinsumStridedViewMatchesDense(t *testing.T) { + for _, dt := range []Dtype{Int, Float32, Float, Complex} { + view := einsumStridedOperand(t, dt, []int{3, 4}, []int{8, 2}) + dense := einsumDenseOf(t, view) + other := einsumOperand(t, dt, 3, 4) + vec := einsumOperand(t, dt, 4) + for _, tc := range []struct { + spec string + which []bool // true takes the strided view, false the dense operand + }{ + {"ij->i", []bool{true}}, + {"ij->j", []bool{true}}, + {"ij,ij->i", []bool{true, false}}, + {"ij,ij->ji", []bool{true, false}}, + {"ij,kj->ik", []bool{true, false}}, + {"...i,i->...", []bool{true, false}}, + } { + withView := make([]*Array, len(tc.which)) + withDense := make([]*Array, len(tc.which)) + for i, isView := range tc.which { + if isView { + withView[i], withDense[i] = view, dense + continue + } + if tc.spec == "...i,i->..." { + withView[i], withDense[i] = vec, vec + continue + } + withView[i], withDense[i] = other, other + } + got, gerr := Einsum(tc.spec, withView...) + want, werr := Einsum(tc.spec, withDense...) + if werr != nil || gerr != nil { + t.Fatalf("%s/%s: dense error %v, view error %v", tc.spec, dt, werr, gerr) + } + if !einsumBitsEqual(got, want) { + t.Fatalf("%s/%s: view result %s differs from the dense array %s", tc.spec, dt, got, want) + } + } + } +} + +// benchEinsum measures one spec with the worker pool as it stands and +// with it pinned to one worker: the in-process A/B of the slot walk's +// split, where the pinned pool walks every slot in one chunk. +func benchEinsum(b *testing.B, spec string, ops ...*Array) { + b.Run("pool", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum(spec, ops...); err != nil { + b.Fatal(err) + } + } + }) + b.Run("one worker", func(b *testing.B) { + restore := engine.SetNumWorkers(1) + defer engine.SetNumWorkers(restore) + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum(spec, ops...); err != nil { + b.Fatal(err) + } + } + }) +} + +// BenchmarkEinsumEllipsisBatch is a contraction carrying a batch axis: +// the pattern the dispatch table cannot take, walked slot by slot. +func BenchmarkEinsumEllipsisBatch(b *testing.B) { + a := benchMat(b, 11, 8, 48, 48) + c := benchMat(b, 12, 48, 48) + benchEinsum(b, "...ij,jk->...ik", a, c) +} + +// BenchmarkEinsumEllipsisFloat32 is the same contraction on float32 +// payloads: the operands widen into the walk and every addend narrows +// into the float32 slot. +func BenchmarkEinsumEllipsisFloat32(b *testing.B) { + a := einsumOperand(b, Float32, 8, 48, 48) + c := einsumOperand(b, Float32, 48, 48) + benchEinsum(b, "...ij,jk->...ik", a, c) +} + +// BenchmarkEinsumThreeOperandChain sums one shared label across three +// operands, the shape attention and tensordot patterns reduce to. +func BenchmarkEinsumThreeOperandChain(b *testing.B) { + a := benchMat(b, 13, 16, 32) + c := benchMat(b, 14, 32, 32) + d := benchMat(b, 15, 32, 16) + benchEinsum(b, "ik,kj,jl->il", a, c, d) +} + +// BenchmarkEinsumSmallBatchedProduct is the small batched product the +// dispatch table's own batched kernel takes, kept as the fixed-cost +// control beside the general engine's shapes. +func BenchmarkEinsumSmallBatchedProduct(b *testing.B) { + a := benchMat(b, 16, 4, 32, 32) + c := benchMat(b, 17, 4, 32, 32) + benchEinsum(b, "bij,bjk->bik", a, c) +} + +// BenchmarkEinsumReduceOnly measures a reduction-only spec through the +// axis fold the dispatch now takes, against the same spec on the +// general engine, which is where it went before. +func BenchmarkEinsumReduceOnly(b *testing.B) { + a := benchMat(b, 18, 512, 512) + b.Run("axis fold", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum("ij->i", a); err != nil { + b.Fatal(err) + } + } + }) + b.Run("general engine", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := einsumGeneral([]string{"ij"}, "i", true, []*Array{a}); err != nil { + b.Fatal(err) + } + } + }) + b.Run("general engine, one worker", func(b *testing.B) { + restore := engine.SetNumWorkers(1) + defer engine.SetNumWorkers(restore) + b.ReportAllocs() + for b.Loop() { + if _, err := einsumGeneral([]string{"ij"}, "i", true, []*Array{a}); err != nil { + b.Fatal(err) + } + } + }) +} + +// BenchmarkEinsumSlotKernelWalk is one slot's walk shaped like a +// contraction's, two operands over a summed axis of 48 with the second +// operand's cursor stepping a row: the peeled kernel against the +// odometer oracle, in one process. +func BenchmarkEinsumSlotKernelWalk(b *testing.B) { + sizes := []int{48} + base := []int{0, 0} + delta := []int{1, 48} + rv := [][]float64{einsumProbeFloats(48), einsumProbeFloats(48 * 48)} + out := make([]float64, 1) + off := make([]int, 2) + coord := make([]int, 1) + b.Run("peeled", func(b *testing.B) { + for b.Loop() { + einsumSlotSum(rv, out, 0, base, off, coord, delta, sizes, 48) + } + }) + b.Run("odometer", func(b *testing.B) { + for b.Loop() { + einsumSlotSumOdometer(rv, out, 0, base, off, coord, delta, sizes, 48) + } + }) +} diff --git a/internal/core/bench_kernels_ab_test.go b/internal/core/bench_kernels_ab_test.go new file mode 100644 index 0000000..50364e4 --- /dev/null +++ b/internal/core/bench_kernels_ab_test.go @@ -0,0 +1,274 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "fmt" + "testing" +) + +// The A/B territory benchmarks: every kernel under review timed through +// its exported entry point. The file only touches the exported API, so +// the same file runs against the previous revision in a second worktree +// and the two sides interleave round by round. + +// abLCG is a small deterministic fill so neither side depends on a +// random package's stream. +type abLCG struct{ s uint64 } + +func (g *abLCG) next() float64 { + g.s = g.s*6364136223846793005 + 1442695040888963407 + return float64(g.s>>11) / (1 << 53) +} + +func abFloats(n int, seed uint64) *Array { + g := &abLCG{s: seed} + vals := make([]float64, n) + for i := range vals { + vals[i] = g.next() + } + a, _ := FromFloats(vals, n) + return a +} + +func BenchmarkABCovariance(b *testing.B) { + x := abFloats(1<<20, 1) + y := abFloats(1<<20, 2) + for b.Loop() { + if _, err := Covariance(x, y); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(x.Len()) * 8) +} + +func BenchmarkABCorrelation(b *testing.B) { + x := abFloats(1<<20, 1) + y := abFloats(1<<20, 2) + for b.Loop() { + if _, err := Correlation(x, y); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(x.Len()) * 8) +} + +func BenchmarkABIntegrate(b *testing.B) { + y := abFloats(1<<20, 3) + for b.Loop() { + if _, err := Integrate(y, 0.1); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(y.Len()) * 8) +} + +func BenchmarkABCumSum1D(b *testing.B) { + y := abFloats(1<<20, 4) + for b.Loop() { + if _, err := CumSum(y, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(y.Len()) * 8) +} + +func BenchmarkABCumSumShortLines(b *testing.B) { + // 4096x2: the short-line scan the existing suite also guards, where + // the per-line compensation overhead should be visible if anywhere. + vals := make([]float64, 8192) + for i := range vals { + vals[i] = float64(i % 97) + } + y, _ := FromFloats(vals, 4096, 2) + for b.Loop() { + if _, err := CumSum(y, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkABSign(b *testing.B) { + y := abFloats(1<<20, 5) + for b.Loop() { + if _, err := Sign(y); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(y.Len()) * 8) +} + +func BenchmarkABSignSmall(b *testing.B) { + y := abFloats(256, 6) + for b.Loop() { + if _, err := Sign(y); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkABEvaluatePolynomial(b *testing.B) { + coeffs := abFloats(16, 7) + x := abFloats(20000, 8) + for b.Loop() { + if _, err := EvaluatePolynomial(coeffs, x); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(x.Len()) * 8) +} + +func BenchmarkABInterpolateGrid(b *testing.B) { + for _, tc := range []struct { + side, m int + }{{128, 200000}, {8, 50000}} { + dims := 2 + if tc.side == 8 { + dims = 4 + } + shape := make([]int, dims) + for i := range shape { + shape[i] = tc.side + } + size := 1 + for range dims { + size *= tc.side + } + grid := abFloats(size, uint64(tc.side)) + grid = mustABReshape(b, grid, shape...) + g := &abLCG{s: uint64(tc.m)} + q := make([]float64, tc.m*dims) + for i := range q { + q[i] = g.next() * float64(tc.side-1) + } + queries, _ := FromFloats(q, tc.m, dims) + origins := make([]float64, dims) + steps := make([]float64, dims) + for i := range steps { + steps[i] = 1 + } + name := fmt.Sprintf("dims=%d", dims) + b.Run(name, func(b *testing.B) { + for b.Loop() { + if _, err := InterpolateGrid(grid, origins, steps, queries); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(tc.m) * 8) + }) + } +} + +func BenchmarkABSobol(b *testing.B) { + for _, dim := range []int{2, 8} { + b.Run(fmt.Sprintf("dim=%d", dim), func(b *testing.B) { + for b.Loop() { + if _, err := SobolPoints(1<<20, dim, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(1<<20) * int64(dim) * 8) + }) + } +} + +func BenchmarkABHalton(b *testing.B) { + for _, dim := range []int{2, 6} { + b.Run(fmt.Sprintf("dim=%d", dim), func(b *testing.B) { + for b.Loop() { + if _, err := HaltonPoints(1<<20, dim, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(1<<20) * int64(dim) * 8) + }) + } +} + +func BenchmarkABConcatPromoted(b *testing.B) { + a, _ := FromFloat32s(make([]float32, 1<<19), 1<<19) + bv := abFloats(1<<19, 9) + for b.Loop() { + if _, err := Concat(a, bv, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(a.Len()+bv.Len()) * 8) +} + +func BenchmarkABConcatSame(b *testing.B) { + a := abFloats(1<<19, 10) + bv := abFloats(1<<19, 11) + for b.Loop() { + if _, err := Concat(a, bv, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(a.Len()+bv.Len()) * 8) +} + +func BenchmarkABDiff(b *testing.B) { + y := abFloats(1<<20, 12) + for b.Loop() { + if _, err := Diff(y, 1, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(y.Len()) * 8) +} + +func BenchmarkABDiff2D(b *testing.B) { + // 4096x256 along axis 0: 4096 independent head runs, the shape the + // run split can fan out over. + n := 4096 * 256 + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(i % 511) + } + y, _ := FromFloats(vals, 4096, 256) + for b.Loop() { + if _, err := Diff(y, 1, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(n) * 8) +} + +func BenchmarkABOneHot(b *testing.B) { + g := &abLCG{s: 13} + codes := make([]int64, 1<<20) + for i := range codes { + codes[i] = int64(g.next() * 64) + } + a, _ := FromInts(codes, 1<<20) + for b.Loop() { + if _, err := OneHot(a, 64); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(len(codes)) * 8) +} + +func BenchmarkABMoveAxis(b *testing.B) { + vals := make([]float64, 128*128*128) + for i := range vals { + vals[i] = float64(i) + } + a, _ := FromFloats(vals, 128, 128, 128) + for b.Loop() { + if _, err := MoveAxis(a, 2, 0); err != nil { + b.Fatal(err) + } + } + b.SetBytes(int64(len(vals)) * 8) +} + +func mustABReshape(b *testing.B, a *Array, shape ...int) *Array { + b.Helper() + out, err := Reshape(a, shape...) + if err != nil { + b.Fatal(err) + } + return out +} diff --git a/internal/core/bench_mat2_test.go b/internal/core/bench_mat2_test.go new file mode 100644 index 0000000..52dc0dc --- /dev/null +++ b/internal/core/bench_mat2_test.go @@ -0,0 +1,188 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Dtype and split companions to the matmul benchmarks: bench_test.go, +// mat_tile_bench_test.go and mat_kernel_bench_test.go pin the float64 +// shapes, so these pin the complex 2-D product, the float32, int and +// complex vector walks, and the vector split's per-worker floor. The +// operands are fixed literals, so a run is comparable to the next. + +// benchMatMulComplexRect runs the n×k·k×m complex product under the +// default worker policy. +func benchMatMulComplexRect(b *testing.B, n, k, m int) { + b.Helper() + av := make([]complex128, n*k) + bv := make([]complex128, k*m) + for i := range av { + av[i] = complex(float64(i%7)-3, float64(i%5)-2) + } + for i := range bv { + bv[i] = complex(float64(i%5)-2, float64(i%11)-5) + } + a, err := FromComplexes(av, n, k) + if err != nil { + b.Fatal(err) + } + c, err := FromComplexes(bv, k, m) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMatMulComplex256 is the smallest complex square the split +// engages, one size above matMulComplexParallelMin. +func BenchmarkMatMulComplex256(b *testing.B) { benchMatMulComplexRect(b, 256, 256, 256) } + +// BenchmarkMatMulComplex512 is the complex square product: its b +// payload is sixteen bytes per element, so sharing one b row across a +// panel carries the most weight here. +func BenchmarkMatMulComplex512(b *testing.B) { benchMatMulComplexRect(b, 512, 512, 512) } + +// The complex rectangular extremes, mirroring the float64 shape set. +func BenchmarkMatMulComplexFat64x512x2048(b *testing.B) { + benchMatMulComplexRect(b, 64, 512, 2048) +} + +func BenchmarkMatMulComplexTall512x64x2048(b *testing.B) { + benchMatMulComplexRect(b, 512, 64, 2048) +} + +func BenchmarkMatMulComplexWide2048x64x512(b *testing.B) { + benchMatMulComplexRect(b, 2048, 64, 512) +} + +func BenchmarkMatMulComplexSkinny2048x512x64(b *testing.B) { + benchMatMulComplexRect(b, 2048, 512, 64) +} + +// benchMatVecTyped runs the 2-D×1-D (swap true) or 1-D×2-D (swap +// false) product of a typed pair built from fixed literals, under the +// default worker policy. The vector carries the inner dimension: the +// matrix's column count for a row-dot product, its row count for a +// column-dot one. +func benchMatVecTyped(b *testing.B, dt Dtype, rows, cols int, swap bool) { + b.Helper() + n := rows * cols + flat := make([]float64, n) + for i := range flat { + flat[i] = float64(i%7) - 3 + } + vlen := cols + if !swap { + vlen = rows + } + vecf := make([]float64, vlen) + for i := range vecf { + vecf[i] = float64(i%5) - 2 + } + var mat, vec *Array + var err error + switch dt { + case Float32: + f32 := make([]float32, n) + for i, v := range flat { + f32[i] = float32(v) + } + v32 := make([]float32, vlen) + for i, v := range vecf { + v32[i] = float32(v) + } + mat, err = FromFloat32s(f32, rows, cols) + if err == nil { + vec, err = FromFloat32s(v32, vlen) + } + case Int: + iv := make([]int64, n) + for i, v := range flat { + iv[i] = int64(v) + } + vv := make([]int64, vlen) + for i, v := range vecf { + vv[i] = int64(v) + } + mat, err = FromInts(iv, rows, cols) + if err == nil { + vec, err = FromInts(vv, vlen) + } + default: + cv := make([]complex128, n) + for i, v := range flat { + cv[i] = complex(v, float64(i%5)-2) + } + cvv := make([]complex128, vlen) + for i, v := range vecf { + cvv[i] = complex(v, float64(i%3)-1) + } + mat, err = FromComplexes(cv, rows, cols) + if err == nil { + vec, err = FromComplexes(cvv, vlen) + } + } + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + var err error + if swap { + _, err = MatMul2D(mat, vec) + } else { + _, err = MatMul2D(vec, mat) + } + if err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMatVecF32RowDots512x512 pins the float32 row dots, whose +// four-row walk is the float64 twin's structure. +func BenchmarkMatVecF32RowDots512x512(b *testing.B) { + benchMatVecTyped(b, Float32, 512, 512, true) +} + +// BenchmarkMatVecIntRowDots512x512 pins the int row dots. +func BenchmarkMatVecIntRowDots512x512(b *testing.B) { benchMatVecTyped(b, Int, 512, 512, true) } + +// BenchmarkMatVecComplexRowDots256x256 pins the complex row dots, whose +// two-row walk is what a sixteen-byte sum chain can hold in registers. +func BenchmarkMatVecComplexRowDots256x256(b *testing.B) { + benchMatVecTyped(b, Complex, 256, 256, true) +} + +// BenchmarkVecMatF32ColDots128x2048 pins the float32 column dots. +func BenchmarkVecMatF32ColDots128x2048(b *testing.B) { + benchMatVecTyped(b, Float32, 128, 2048, false) +} + +// BenchmarkVecMatComplexColDots128x2048 pins the complex column dots. +func BenchmarkVecMatComplexColDots128x2048(b *testing.B) { + benchMatVecTyped(b, Complex, 128, 2048, false) +} + +// BenchmarkMatVecRowDots128x128 sits just above the vector split's +// per-worker floor: 16,384 terms fill at most one worker, so the shape +// measures the lone walk the floor keeps it on. +func BenchmarkMatVecRowDots128x128(b *testing.B) { benchMatVecTyped(b, Float, 128, 128, true) } + +// BenchmarkVecMatColDots16x16384 pins the wide column dots: sixteen +// rows by 16,384 output columns, so each worker's band is far wider +// than a cache line. +func BenchmarkVecMatColDots16x16384(b *testing.B) { benchMatVecTyped(b, Float, 16, 16384, false) } + +// BenchmarkVecMatColNarrowDots16384x16 is the thin-band shape: sixteen +// output columns of 16,384 terms, where the split hands a worker a band +// no wider than a cache line. +func BenchmarkVecMatColNarrowDots16384x16(b *testing.B) { + benchMatVecTyped(b, Float, 16384, 16, false) +} diff --git a/internal/core/bench_narrow_test.go b/internal/core/bench_narrow_test.go new file mode 100644 index 0000000..1946c2d --- /dev/null +++ b/internal/core/bench_narrow_test.go @@ -0,0 +1,342 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// The narrow element types' kernels: element-wise arithmetic, the +// comparisons, the scalar maps, the full-array reductions and Where, +// each against the int64 or float64 kernel of the same element count, +// so a narrow walk is judged against the widest loop it could have +// widened into. Every payload holds one million elements, the shape the +// BenchmarkAdd1M family reports at. + +// narrowBenchLen is the element count every payload in this file fills. +const narrowBenchLen = 1 << 20 + +// benchNarrowInts fills an n-element signed integer slice with a +// deterministic mix of positive and negative values. +func benchNarrowInts[T ~int8 | ~int16 | ~int32](n int) []T { + src := make([]T, n) + for i := range src { + src[i] = T(i%251) - 125 + } + return src +} + +// benchNarrowUints fills an n-element unsigned integer slice with a +// deterministic mix of small values. +func benchNarrowUints[T ~uint8 | ~uint16 | ~uint32](n int) []T { + src := make([]T, n) + for i := range src { + src[i] = T(i % 251) + } + return src +} + +// benchNarrowPair builds two one-million-element arrays of dtype dt +// with the same deterministic fill, the pair the binary kernels walk. +func benchNarrowPair(dt Dtype) (*Array, *Array) { + must := func(a *Array, err error) *Array { + if err != nil { + panic(err) + } + return a + } + fillInt := func() []int64 { + v := make([]int64, narrowBenchLen) + for i := range v { + v[i] = int64(i%251) - 125 + } + return v + } + fillFloat := func() []float64 { + v := make([]float64, narrowBenchLen) + for i := range v { + v[i] = float64(i%251) - 125 + } + return v + } + switch dt { + case Int8: + return must(Int8sFromArray(benchNarrowInts[int8](narrowBenchLen), narrowBenchLen)), + must(Int8sFromArray(benchNarrowInts[int8](narrowBenchLen), narrowBenchLen)) + case Uint8: + return must(Uint8sFromArray(benchNarrowUints[uint8](narrowBenchLen), narrowBenchLen)), + must(Uint8sFromArray(benchNarrowUints[uint8](narrowBenchLen), narrowBenchLen)) + case Int16: + return must(Int16sFromArray(benchNarrowInts[int16](narrowBenchLen), narrowBenchLen)), + must(Int16sFromArray(benchNarrowInts[int16](narrowBenchLen), narrowBenchLen)) + case Uint16: + return must(Uint16sFromArray(benchNarrowUints[uint16](narrowBenchLen), narrowBenchLen)), + must(Uint16sFromArray(benchNarrowUints[uint16](narrowBenchLen), narrowBenchLen)) + case Int32: + return must(Int32sFromArray(benchNarrowInts[int32](narrowBenchLen), narrowBenchLen)), + must(Int32sFromArray(benchNarrowInts[int32](narrowBenchLen), narrowBenchLen)) + case Uint32: + return must(Uint32sFromArray(benchNarrowUints[uint32](narrowBenchLen), narrowBenchLen)), + must(Uint32sFromArray(benchNarrowUints[uint32](narrowBenchLen), narrowBenchLen)) + case Int: + return must(FromInts(fillInt(), narrowBenchLen)), must(FromInts(fillInt(), narrowBenchLen)) + default: // Float + return must(FromFloats(fillFloat(), narrowBenchLen)), must(FromFloats(fillFloat(), narrowBenchLen)) + } +} + +func benchNarrowAdd(b *testing.B, dt Dtype) { + b.Helper() + a, c := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowAddInt8(b *testing.B) { benchNarrowAdd(b, Int8) } +func BenchmarkNarrowAddUint8(b *testing.B) { benchNarrowAdd(b, Uint8) } +func BenchmarkNarrowAddInt16(b *testing.B) { benchNarrowAdd(b, Int16) } +func BenchmarkNarrowAddUint16(b *testing.B) { benchNarrowAdd(b, Uint16) } +func BenchmarkNarrowAddInt32(b *testing.B) { benchNarrowAdd(b, Int32) } +func BenchmarkNarrowAddUint32(b *testing.B) { benchNarrowAdd(b, Uint32) } +func BenchmarkNarrowAddInt64(b *testing.B) { benchNarrowAdd(b, Int) } +func BenchmarkNarrowAddFloat64(b *testing.B) { benchNarrowAdd(b, Float) } + +func benchNarrowMul(b *testing.B, dt Dtype) { + b.Helper() + a, c := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + if _, err := Mul(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowMulInt8(b *testing.B) { benchNarrowMul(b, Int8) } +func BenchmarkNarrowMulUint16(b *testing.B) { benchNarrowMul(b, Uint16) } +func BenchmarkNarrowMulInt32(b *testing.B) { benchNarrowMul(b, Int32) } +func BenchmarkNarrowMulInt64(b *testing.B) { benchNarrowMul(b, Int) } + +func benchNarrowLt(b *testing.B, dt Dtype) { + b.Helper() + a, c := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + if _, err := Lt(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowLtInt8(b *testing.B) { benchNarrowLt(b, Int8) } +func BenchmarkNarrowLtUint8(b *testing.B) { benchNarrowLt(b, Uint8) } +func BenchmarkNarrowLtInt16(b *testing.B) { benchNarrowLt(b, Int16) } +func BenchmarkNarrowLtUint32(b *testing.B) { benchNarrowLt(b, Uint32) } +func BenchmarkNarrowLtInt64(b *testing.B) { benchNarrowLt(b, Int) } +func BenchmarkNarrowLtFloat64(b *testing.B) { benchNarrowLt(b, Float) } + +func benchNarrowEqI(b *testing.B, dt Dtype) { + b.Helper() + a, _ := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + if _, err := EqI(a, 17); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowEqIInt8(b *testing.B) { benchNarrowEqI(b, Int8) } +func BenchmarkNarrowEqIUint16(b *testing.B) { benchNarrowEqI(b, Uint16) } +func BenchmarkNarrowEqIInt64(b *testing.B) { benchNarrowEqI(b, Int) } + +func benchNarrowLtF(b *testing.B, dt Dtype) { + b.Helper() + a, _ := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + if _, err := LtF(a, 100.5); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowLtFInt32(b *testing.B) { benchNarrowLtF(b, Int32) } +func BenchmarkNarrowLtFInt64(b *testing.B) { benchNarrowLtF(b, Int) } + +func benchNarrowSum(b *testing.B, dt Dtype) { + b.Helper() + a, _ := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + Sum(a) + } +} + +func BenchmarkNarrowSumInt8(b *testing.B) { benchNarrowSum(b, Int8) } +func BenchmarkNarrowSumUint8(b *testing.B) { benchNarrowSum(b, Uint8) } +func BenchmarkNarrowSumInt16(b *testing.B) { benchNarrowSum(b, Int16) } +func BenchmarkNarrowSumUint16(b *testing.B) { benchNarrowSum(b, Uint16) } +func BenchmarkNarrowSumInt32(b *testing.B) { benchNarrowSum(b, Int32) } +func BenchmarkNarrowSumUint32(b *testing.B) { benchNarrowSum(b, Uint32) } +func BenchmarkNarrowSumInt64(b *testing.B) { benchNarrowSum(b, Int) } +func BenchmarkNarrowSumFloat64(b *testing.B) { benchNarrowSum(b, Float) } +func BenchmarkNarrowSumBool(b *testing.B) { + a, _ := FromBools(make([]bool, narrowBenchLen), narrowBenchLen) + bs := a.RawBools() + for i := range bs { + bs[i] = i%3 != 0 + } + b.ReportAllocs() + for b.Loop() { + Sum(a) + } +} + +func benchNarrowMin(b *testing.B, dt Dtype) { + b.Helper() + a, _ := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + if _, err := Min(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowMinInt8(b *testing.B) { benchNarrowMin(b, Int8) } +func BenchmarkNarrowMinUint32(b *testing.B) { benchNarrowMin(b, Uint32) } +func BenchmarkNarrowMinInt64(b *testing.B) { benchNarrowMin(b, Int) } + +func benchNarrowDot(b *testing.B, dt Dtype) { + b.Helper() + a, c := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + if _, err := Dot(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowDotInt8(b *testing.B) { benchNarrowDot(b, Int8) } +func BenchmarkNarrowDotUint16(b *testing.B) { benchNarrowDot(b, Uint16) } +func BenchmarkNarrowDotInt32(b *testing.B) { benchNarrowDot(b, Int32) } +func BenchmarkNarrowDotInt64(b *testing.B) { benchNarrowDot(b, Int) } + +func benchNarrowAddF(b *testing.B, dt Dtype) { + b.Helper() + a, _ := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + AddF(a, 0.5) + } +} + +func BenchmarkNarrowAddFInt8(b *testing.B) { benchNarrowAddF(b, Int8) } +func BenchmarkNarrowAddFInt32(b *testing.B) { benchNarrowAddF(b, Int32) } +func BenchmarkNarrowAddFInt64(b *testing.B) { benchNarrowAddF(b, Int) } + +func benchNarrowMulI(b *testing.B, dt Dtype) { + b.Helper() + a, _ := benchNarrowPair(dt) + b.ReportAllocs() + for b.Loop() { + MulI(a, 3) + } +} + +func BenchmarkNarrowMulIInt8(b *testing.B) { benchNarrowMulI(b, Int8) } +func BenchmarkNarrowMulIUint32(b *testing.B) { benchNarrowMulI(b, Uint32) } +func BenchmarkNarrowMulIInt64(b *testing.B) { benchNarrowMulI(b, Int) } + +func benchNarrowWhere(b *testing.B, dt Dtype) { + b.Helper() + a, c := benchNarrowPair(dt) + cond, _ := FromInts(make([]int64, narrowBenchLen), narrowBenchLen) + ci := cond.RawInts() + for i := range ci { + if i%2 == 0 { + ci[i] = 1 + } + } + b.ReportAllocs() + for b.Loop() { + if _, err := Where(cond, a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowWhereInt8(b *testing.B) { benchNarrowWhere(b, Int8) } +func BenchmarkNarrowWhereUint16(b *testing.B) { benchNarrowWhere(b, Uint16) } +func BenchmarkNarrowWhereInt64(b *testing.B) { benchNarrowWhere(b, Int) } + +func BenchmarkNarrowAndBool(b *testing.B) { + a, _ := FromBools(make([]bool, narrowBenchLen), narrowBenchLen) + c, _ := FromBools(make([]bool, narrowBenchLen), narrowBenchLen) + ab, cb := a.RawBools(), c.RawBools() + for i := range ab { + ab[i] = i%3 != 0 + cb[i] = i%5 != 0 + } + b.ReportAllocs() + for b.Loop() { + if _, err := And(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowNotBool(b *testing.B) { + a, _ := FromBools(make([]bool, narrowBenchLen), narrowBenchLen) + ab := a.RawBools() + for i := range ab { + ab[i] = i%3 != 0 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Not(a); err != nil { + b.Fatal(err) + } + } +} + +// The mixed-width pairs: a narrow operand against a different narrow +// width and against Int, both promoted walks a single dense payload +// read from each side. +func BenchmarkNarrowMixedAddI8I16(b *testing.B) { + a, _ := benchNarrowPair(Int8) + c, _ := benchNarrowPair(Int16) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowMixedAddU8Int(b *testing.B) { + a, _ := benchNarrowPair(Uint8) + c, _ := benchNarrowPair(Int) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNarrowMixedLtI8U16(b *testing.B) { + a, _ := benchNarrowPair(Int8) + c, _ := benchNarrowPair(Uint16) + b.ReportAllocs() + for b.Loop() { + if _, err := Lt(a, c); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/bench_ops2_test.go b/internal/core/bench_ops2_test.go new file mode 100644 index 0000000..35aeecc --- /dev/null +++ b/internal/core/bench_ops2_test.go @@ -0,0 +1,499 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// Benchmarks for the walks whose per-element body changed: the binary +// arithmetic maps across every dtype, the mask comparisons and gathers, +// the scan and product kernels with short lines, the Lp norm, the +// broadcast fill and the magnitude map. Inputs are fixed literals, so +// every run feeds the same bytes. + +// opFloats builds an n-element float array whose i-th element is +// deterministic and nonzero. +func opFloats(b *testing.B, n int) *Array { + b.Helper() + a, _ := FromFloats(make([]float64, n), n) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%97) - 48 + } + return a +} + +// opFloat32s is opFloats in single precision, built through the +// float32 payload so the values are the same ones the dtype round-trips. +func opFloat32s(b *testing.B, n int) *Array { + b.Helper() + f64 := opFloats(b, n) + a, _ := Astype(f64, Float32) + return a +} + +// opInts is opFloats over an int payload; no element is zero, so the +// division benchmarks never take the error path. +func opInts(b *testing.B, n int) *Array { + b.Helper() + a, _ := FromInts(make([]int64, n), n) + for i := range a.RawInts() { + a.RawInts()[i] = int64(i%97) + 1 + } + return a +} + +// --- binary arithmetic maps --- + +// BenchmarkSub1M and BenchmarkMul1M measure the two relations the +// add-size sweep in math_bench_extra_test.go does not cover. +func BenchmarkSub1M(b *testing.B) { + a, c := opFloats(b, 1<<20), opFloats(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Sub(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMul1M(b *testing.B) { + a, c := opFloats(b, 1<<20), opFloats(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Mul(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkPow1M walks the float branch of Pow, whose body holds the +// math.Pow call instead of a func value. +func BenchmarkPow1M(b *testing.B) { + a, c := opFloats(b, 1<<20), opFloats(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Pow(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkAddInt1M pins the int payload walk: the int64 sum wraps, so +// the loop body is the add itself. +func BenchmarkAddInt1M(b *testing.B) { + a, c := opInts(b, 1<<20), opInts(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkAddFloat32_1M pins the narrow float walk, whose widening is +// exact and whose result narrows once. +func BenchmarkAddFloat32_1M(b *testing.B) { + a, c := opFloat32s(b, 1<<20), opFloat32s(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkAddFloat16_1M pins the half walk: both payloads widen to +// float64 and the sum narrows once per element. +func BenchmarkAddFloat16_1M(b *testing.B) { + f64a, f64b := opFloats(b, 1<<20), opFloats(b, 1<<20) + a, err := Astype(f64a, Float16) + if err != nil { + b.Fatal(err) + } + c, err := Astype(f64b, Float16) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMulComplex100k measures the widest payload: a complex product +// is four multiplies and two adds per element over 16 bytes each side. +func BenchmarkMulComplex100k(b *testing.B) { + ac, _ := FromComplexes(make([]complex128, 100_000), 100_000) + bc, _ := FromComplexes(make([]complex128, 100_000), 100_000) + for i := range ac.RawComplexes() { + v := complex(float64(i%97)-48, float64(i%31)-15) + ac.RawComplexes()[i] = v + bc.RawComplexes()[i] = complex(float64(i%53)-26, float64(i%17)-8) + } + b.ReportAllocs() + for b.Loop() { + if _, err := Mul(ac, bc); err != nil { + b.Fatal(err) + } + } +} + +// --- mask comparisons, selection and gather --- + +// BenchmarkLtFloat64_1M measures the array-with-array comparison at the +// size where the walk has to fan out. +func BenchmarkLtFloat64_1M(b *testing.B) { + a, c := opFloats(b, 1<<20), opFloats(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Lt(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkGeIFloat64_1M measures the scalar comparison, which compares +// exactly against the int and hoists the mirrored relation. +func BenchmarkGeIFloat64_1M(b *testing.B) { + a := opFloats(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := GeI(a, 3); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkEqComplex1M pins the complex equality walk. +func BenchmarkEqComplex1M(b *testing.B) { + ac, _ := FromComplexes(make([]complex128, 1<<20), 1<<20) + bc, _ := FromComplexes(make([]complex128, 1<<20), 1<<20) + for i := range ac.RawComplexes() { + ac.RawComplexes()[i] = complex(float64(i%97)-48, float64(i%31)-15) + bc.RawComplexes()[i] = complex(float64(i%53)-26, float64(i%17)-8) + } + b.ReportAllocs() + for b.Loop() { + if _, err := Eq(ac, bc); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkWhereFloat64_1M measures the dense int-conditioned select +// over float payloads: one condition read and one gathered value per +// element, the payload-slots fast path. +func BenchmarkWhereFloat64_1M(b *testing.B) { + n := 1 << 20 + cond, _ := FromInts(make([]int64, n), n) + for i := range cond.RawInts() { + cond.RawInts()[i] = int64(i % 2) + } + x, y := opFloats(b, n), opFloats(b, n) + b.ReportAllocs() + for b.Loop() { + if _, err := Where(cond, x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkWhereMixed100k measures the accessor fallback: the bool mask +// a comparison answers, an int operand against a float one promoting to +// float, so the walk reads through the accessors. +func BenchmarkWhereMixed100k(b *testing.B) { + n := 100_000 + cond, err := GeI(opFloats(b, n), 0) + if err != nil { + b.Fatal(err) + } + x := opInts(b, n) + y := opFloats(b, n) + b.ReportAllocs() + for b.Loop() { + if _, err := Where(cond, x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSelectHalf1M gathers every other element: the count and the +// compaction carry half a million hits each. +func BenchmarkSelectHalf1M(b *testing.B) { + n := 1 << 20 + a := opFloats(b, n) + m, _ := FromInts(make([]int64, n), n) + for i := range m.RawInts() { + m.RawInts()[i] = int64(i % 2) + } + b.ReportAllocs() + for b.Loop() { + if _, err := Select(a, m); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSelectSparse1M gathers one element in sixty-four: the count +// pass dominates, since the compaction writes almost nothing. +func BenchmarkSelectSparse1M(b *testing.B) { + n := 1 << 20 + a := opFloats(b, n) + m, _ := FromInts(make([]int64, n), n) + for i := range m.RawInts() { + if i%64 == 0 { + m.RawInts()[i] = 1 + } + } + b.ReportAllocs() + for b.Loop() { + if _, err := Select(a, m); err != nil { + b.Fatal(err) + } + } +} + +// --- scans and products with short lines --- + +// BenchmarkCumSumShortLines scans a 4096x2 array along the two-element +// trailing dimension: there are many lines but each carries almost no +// work, so the fan-out must not be granted per line. +func BenchmarkCumSumShortLines(b *testing.B) { + a, _ := FromFloats(make([]float64, 4096*2), 4096, 2) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%97) - 48 + } + b.ReportAllocs() + for b.Loop() { + if _, err := CumSum(a, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkCumProdShortLines is the multiply twin at the same shape. +func BenchmarkCumProdShortLines(b *testing.B) { + a, _ := FromFloats(make([]float64, 4096*2), 4096, 2) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%7) + 1 + } + b.ReportAllocs() + for b.Loop() { + if _, err := CumProd(a, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkProdShortLines reduces a 4096x8 array along the trailing +// dimension, the product twin of the same short-line shape. +func BenchmarkProdShortLines(b *testing.B) { + a, _ := FromFloats(make([]float64, 4096*8), 4096, 8) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%7) + 1 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Prod(a, 1, false); err != nil { + b.Fatal(err) + } + } +} + +// --- the Lp norm --- + +// BenchmarkNormL1_512 walks the branch whose sum multiplies instead of +// calling Pow. +func BenchmarkNormL1_512(b *testing.B) { + a, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%97) - 48 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Norm(a, 1, 1, false); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkNormP3_512 walks the general exponent, which calls math.Pow +// and math.Abs per element and therefore carries the low fan-out floor. +func BenchmarkNormP3_512(b *testing.B) { + a, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%97) - 48 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Norm(a, 3, 1, false); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkNormInfShortLines is the infinity norm over 64 lines of 4096 +// elements: few lines, each carrying a long max-abs fold. +func BenchmarkNormInfShortLines(b *testing.B) { + a, _ := FromFloats(make([]float64, 64*4096), 64, 4096) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%97) - 48 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Norm(a, math.Inf(1), 1, false); err != nil { + b.Fatal(err) + } + } +} + +// --- broadcast --- + +// BenchmarkBroadcastToTiny expands a 16x16 target from a 16x1 source: +// 256 elements, far below any spawn floor. +func BenchmarkBroadcastToTiny(b *testing.B) { + src, _ := FromFloats(make([]float64, 16), 16, 1) + b.ReportAllocs() + for b.Loop() { + if _, err := BroadcastTo(src, 16, 16); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkBroadcastToRun1M expands 256x1x1 to 256x64x64: one constant +// run of 4096 elements per outer step. +func BenchmarkBroadcastToRun1M(b *testing.B) { + src, _ := FromFloats(make([]float64, 256), 256, 1, 1) + for i := range src.RawFloats() { + src.RawFloats()[i] = float64(i%97) - 48 + } + b.ReportAllocs() + for b.Loop() { + if _, err := BroadcastTo(src, 256, 64, 64); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkBroadcastToRunCached expands 4x1 to 4x1024: the target is a +// 32k working set, so the run fill dominates the walk instead of the +// allocation and the first touch of a fresh target. +func BenchmarkBroadcastToRunCached(b *testing.B) { + src, _ := FromFloats(make([]float64, 4), 4, 1) + for i := range src.RawFloats() { + src.RawFloats()[i] = float64(i) + 0.5 + } + b.ReportAllocs() + for b.Loop() { + if _, err := BroadcastTo(src, 4, 1024); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkBroadcastToSerial and BenchmarkBroadcastToParallel fix the +// crossover the spawn floor sits on: the same 1x256 to 256x256 +// expansion, walked by one worker and by every worker. +func BenchmarkBroadcastToSerial(b *testing.B) { + SetNumCPU(1) + defer SetNumCPU(0) + opBroadcastNC(b) +} + +func BenchmarkBroadcastToParallel(b *testing.B) { opBroadcastNC(b) } + +func opBroadcastNC(b *testing.B) { + src, _ := FromFloats(make([]float64, 256), 1, 256) + for i := range src.RawFloats() { + src.RawFloats()[i] = float64(i%97) - 48 + } + b.ReportAllocs() + for b.Loop() { + if _, err := BroadcastTo(src, 256, 256); err != nil { + b.Fatal(err) + } + } +} + +// --- magnitude, mean and integer quotient --- + +// BenchmarkAbs1k pins the small end of the magnitude map: the walk is +// cheaper than the spawn it used to pay for. +func BenchmarkAbs1k(b *testing.B) { + a := opFloats(b, 1_000) + b.ReportAllocs() + for b.Loop() { + Abs(a) + } +} + +// BenchmarkAbs1M pins the large end, where the fan-out pays. +func BenchmarkAbs1M(b *testing.B) { + a := opFloats(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + Abs(a) + } +} + +// BenchmarkAbsComplex100k walks the magnitude of a complex payload: a +// math.Hypot per element, so the low fan-out floor. +func BenchmarkAbsComplex100k(b *testing.B) { + ac, _ := FromComplexes(make([]complex128, 100_000), 100_000) + for i := range ac.RawComplexes() { + ac.RawComplexes()[i] = complex(float64(i%97)-48, float64(i%31)-15) + } + b.ReportAllocs() + for b.Loop() { + Abs(ac) + } +} + +// BenchmarkMean1M measures the overflow guard and the fold behind it. +func BenchmarkMean1M(b *testing.B) { + a := opFloats(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Mean(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkQuo1M measures the integer quotient, whose zero check now +// rides the division walk instead of scanning the divisors first. +func BenchmarkQuo1M(b *testing.B) { + a, c := opInts(b, 1<<20), opInts(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Quo(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSum1M is the global sum over a payload large enough that the +// fold is worth partitioning. +func BenchmarkSum1M(b *testing.B) { + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + _ = Sum(a) + } +} + +// BenchmarkSum1MInt is the integer fold, which is exact under any +// partition and therefore needs no interleaved partials. +func BenchmarkSum1MInt(b *testing.B) { + a := benchIntsShape(1<<20, 1<<20) + b.ReportAllocs() + for b.Loop() { + _ = Sum(a) + } +} diff --git a/internal/core/bench_probes_test.go b/internal/core/bench_probes_test.go new file mode 100644 index 0000000..4994f5e --- /dev/null +++ b/internal/core/bench_probes_test.go @@ -0,0 +1,744 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "cmp" + "math" + "slices" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Paired probes for two dispatch decisions: the Argwhere packing +// strategy and the sort radix crew cap and digit width. Each probe +// interleaves every style on the same deterministic +// fixture in one process, so a drift in the machine's speed lands on +// all styles equally; the reported ns/op per style is the paired +// comparison, and the ns/op of a group as a whole is not comparable +// across groups. + +// probeSink accumulates checksums so no measured walk is elided. +var probeSink int64 + +// --- Argwhere variants ------------------------------------------------- + +// probeArgwhereBaseline is the append-growth Argwhere the merge +// replaced: per-chunk slices grown by append, concatenated in chunk +// order, then copied once more through FromInts. +func probeArgwhereBaseline(a *Array) *Array { + if a.dt == Complex { + return nil + } + n := a.Len() + parts := make([][]int64, 1) + chunk := n + if n >= copyMinPerWorker { + w := workersFor(n) + chunk = (n + w - 1) / w + parts = make([][]int64, (n+chunk-1)/chunk) + } + parallelMin(len(parts), 1, func(s, e int) { + for c := s; c < e; c++ { + start := c * chunk + parts[c] = probeArgwhereRun(a, start, min(start+chunk, n)) + } + }) + total := 0 + for _, p := range parts { + total += len(p) + } + rows := make([]int64, 0, total) + for _, p := range parts { + rows = append(rows, p...) + } + nnz := len(rows) / a.NDim() + out, _ := FromInts(rows, nnz, a.NDim()) + return out +} + +// probeArgwhereRun is the baseline per-chunk collector: the odometer +// seeded from the flat start, coordinates appended as they are met. +func probeArgwhereRun(a *Array, start, end int) []int64 { + ndim := a.NDim() + coord := make([]int, ndim) + if start > 0 { + rest := start + for d := ndim - 1; d >= 0; d-- { + coord[d] = rest % a.shape[d] + rest /= a.shape[d] + } + } + var rows []int64 + appendCoord := func() { + for d := range ndim { + rows = append(rows, int64(coord[d])) + } + } + if !a.isContiguous() { + for i := start; i < end; i++ { + if !isZero(a, i) { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + return rows + } + switch a.dt { + case Int: + p := a.ints + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Float16: + p := a.halves + for i := start; i < end; i++ { + if p[i]&0x7FFF != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Float32: + p := a.floats32 + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + default: + p := a.floats + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + } + return rows +} + +// probeArgwhereCountFill is the count-then-fill variant: an +// odometer-free count pass, then one fill pass writing exact disjoint +// slots of a single allocation. +func probeArgwhereCountFill(a *Array) *Array { + if a.dt == Complex { + return nil + } + n := a.Len() + ndim := a.NDim() + chunk := n + chunks := 1 + if n >= copyMinPerWorker { + w := workersFor(n) + chunk = (n + w - 1) / w + chunks = (n + chunk - 1) / chunk + } + counts := make([]int, chunks) + parallelMin(chunks, 1, func(s, e int) { + for c := s; c < e; c++ { + start := c * chunk + counts[c] = probeArgwhereCount(a, start, min(start+chunk, n)) + } + }) + total := 0 + for _, cnt := range counts { + total += cnt + } + out := &Array{shape: []int{total, ndim}, dt: Int} + out.alloc(total * ndim) + offsets := make([]int, chunks) + off := 0 + for c, cnt := range counts { + offsets[c] = off + off += cnt * ndim + } + parallelMin(chunks, 1, func(s, e int) { + for c := s; c < e; c++ { + start := c * chunk + probeArgwhereFill(a, out.ints[offsets[c]:], start, min(start+chunk, n)) + } + }) + return out +} + +func probeArgwhereCount(a *Array, start, end int) int { + count := 0 + if !a.isContiguous() { + for i := start; i < end; i++ { + if !isZero(a, i) { + count++ + } + } + return count + } + switch a.dt { + case Int: + p := a.ints + for i := start; i < end; i++ { + if p[i] != 0 { + count++ + } + } + case Float16: + p := a.halves + for i := start; i < end; i++ { + if p[i]&0x7FFF != 0 { + count++ + } + } + case Float32: + p := a.floats32 + for i := start; i < end; i++ { + if p[i] != 0 { + count++ + } + } + default: + p := a.floats + for i := start; i < end; i++ { + if p[i] != 0 { + count++ + } + } + } + return count +} + +func probeArgwhereFill(a *Array, dst []int64, start, end int) { + ndim := a.NDim() + coord := make([]int, ndim) + if start > 0 { + rest := start + for d := ndim - 1; d >= 0; d-- { + coord[d] = rest % a.shape[d] + rest /= a.shape[d] + } + } + pos := 0 + appendCoord := func() { + for d := range ndim { + dst[pos] = int64(coord[d]) + pos++ + } + } + if !a.isContiguous() { + for i := start; i < end; i++ { + if !isZero(a, i) { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + return + } + switch a.dt { + case Int: + p := a.ints + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Float16: + p := a.halves + for i := start; i < end; i++ { + if p[i]&0x7FFF != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + case Float32: + p := a.floats32 + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + default: + p := a.floats + for i := start; i < end; i++ { + if p[i] != 0 { + appendCoord() + } + advanceOdometer(coord, a.shape) + } + } +} + +// probeArgwhereMerge is the single-pass merge with a configurable first +// buffer capacity per chunk; bufMin 0 is pure append growth. +func probeArgwhereMerge(a *Array, bufMin int) *Array { + if a.dt == Complex { + return nil + } + n := a.Len() + ndim := a.NDim() + chunk := n + chunks := 1 + if n >= copyMinPerWorker { + w := workersFor(n) + chunk = (n + w - 1) / w + chunks = (n + chunk - 1) / chunk + } + parts := make([][]int64, chunks) + parallelMin(chunks, 1, func(s, e int) { + for c := s; c < e; c++ { + start := c * chunk + end := min(start+chunk, n) + buf := make([]int64, 0, min(bufMin, (end-start)*ndim)) + parts[c] = argwhereAppend(a, buf, start, end) + } + }) + total := 0 + for _, p := range parts { + total += len(p) + } + out := &Array{shape: []int{total / ndim, ndim}, dt: Int} + out.alloc(total) + off := 0 + for _, p := range parts { + copy(out.ints[off:off+len(p)], p) + off += len(p) + } + return out +} + +// argwhereStyles lists the Argwhere variants both probes walk. +func argwhereStyles() []struct { + name string + run func(*Array) *Array +} { + return []struct { + name string + run func(*Array) *Array + }{ + {"production", func(a *Array) *Array { out, _ := Argwhere(a); return out }}, + {"baseline-growth", probeArgwhereBaseline}, + {"count-then-fill", probeArgwhereCountFill}, + {"merge-pure", func(a *Array) *Array { return probeArgwhereMerge(a, 0) }}, + } +} + +// BenchmarkArgwhereVariantsAB interleaves the Argwhere packing variants +// on the 4096x4096 sparse fixture of BenchmarkArgwhere4M. +func BenchmarkArgwhereVariantsAB(b *testing.B) { + a := benchSparseInts(4096*4096, 64) + styles := argwhereStyles() + elapsed := make([]time.Duration, len(styles)) + for b.Loop() { + for i, s := range styles { + start := time.Now() + out := s.run(a) + elapsed[i] += time.Since(start) + probeSink += int64(out.Len()) + if out.Len() > 0 { + probeSink += out.RawInts()[0] + out.RawInts()[out.Len()-1] + } + } + } + for i, s := range styles { + b.ReportMetric(float64(elapsed[i].Nanoseconds())/float64(b.N), "ns/op-"+s.name) + } +} + +// BenchmarkArgwhereVariantAllocs reports each variant's deterministic +// allocation count on the same fixture, one sub-benchmark per style. +func BenchmarkArgwhereVariantAllocs(b *testing.B) { + a := benchSparseInts(4096*4096, 64) + for _, s := range argwhereStyles() { + b.Run(s.name, func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + out := s.run(a) + probeSink += int64(out.Len()) + } + }) + } +} + +// --- Sort radix configurations ---------------------------------------- + +// probeDigitSpec returns the digit shifts, bucket count and exact digit +// mask of a probe radix over the key range diff at the given width. +func probeDigitSpec(diff uint64, bits int) (shifts []uint, buckets int, digits uint32) { + if bits == 8 { + shifts = make([]uint, 8) + for j := range shifts { + shifts[j] = uint(8 * j) + if diff>>shifts[j]&0xFF != 0 { + digits |= 1 << j + } + } + return shifts, 256, digits + } + shifts = append([]uint(nil), radix12Shifts[:]...) + return shifts, radixBuckets, digitMask12Of(diff) +} + +func probeWorkerCap(n, wcap int) int { + if wcap == 0 { + return engine.WorkersFor(n) + } + return min(engine.WorkersFor(n), wcap) +} + +// probeArgSortRadixUint64 mirrors argSortRadixUint64 with the parallel +// crew cap and digit width exposed; the serial paths are production's +// own functions, so only the parallel core varies. +func probeArgSortRadixUint64(keys []uint64, idx []int, diff uint64, wcap, bits int) { + n := len(idx) + if n < radixSeqMin { + slices.SortStableFunc(idx, func(x, y int) int { + return cmp.Compare(keys[x], keys[y]) + }) + return + } + if diff == 0 { + return + } + if n < radixParMin { + argSortRadixSerial(keys, idx, digitMaskOf(diff)) + return + } + buf := make([]int, n) + src, dst := idx, buf + shifts, buckets, digits := probeDigitSpec(diff, bits) + w := probeWorkerCap(n, wcap) + chunk := (n + w - 1) / w + if chunk < radixParMin || chunk >= n { + w = 1 + chunk = n + } + hist := make([]int, w*buckets) + passes := 0 + for j, shift := range shifts { + if digits&(1<>shift)&(buckets-1))]++ + } + }) + sum := 0 + for bb := range buckets { + for c := range w { + cnt := hist[c*buckets+bb] + hist[c*buckets+bb] = sum + sum += cnt + } + } + forEachRadixChunk(n, chunk, func(row, start, end int) { + base := row * buckets + for _, i := range src[start:end] { + d := base + (int(keys[i]>>shift) & (buckets - 1)) + dst[hist[d]] = i + hist[d]++ + } + }) + src, dst = dst, src + passes++ + } + if passes%2 == 1 { + copy(idx, buf) + } +} + +// probeRadixSortUint64 mirrors radixSortUint64 with the same exposures. +func probeRadixSortUint64(vals []uint64, diff uint64, wcap, bits int) { + n := len(vals) + if n < radixSeqMin { + slices.Sort(vals) + return + } + if diff == 0 { + return + } + if n < radixParMin { + probeRadixSortSerial(vals, diff) + return + } + buf := make([]uint64, n) + src, dst := vals, buf + shifts, buckets, digits := probeDigitSpec(diff, bits) + w := probeWorkerCap(n, wcap) + chunk := (n + w - 1) / w + if chunk < radixParMin || chunk >= n { + w = 1 + chunk = n + } + hist := make([]int, w*buckets) + passes := 0 + for j, shift := range shifts { + if digits&(1<>shift)&(buckets-1))]++ + } + }) + sum := 0 + for bb := range buckets { + for c := range w { + cnt := hist[c*buckets+bb] + hist[c*buckets+bb] = sum + sum += cnt + } + } + forEachRadixChunk(n, chunk, func(row, start, end int) { + base := row * buckets + for _, v := range src[start:end] { + d := base + (int(v>>shift) & (buckets - 1)) + dst[hist[d]] = v + hist[d]++ + } + }) + src, dst = dst, src + passes++ + } + if passes%2 == 1 { + copy(vals, buf) + } +} + +// probeRadixSortSerial is the byte-digit serial counting sort for the +// probe's small-n path. +func probeRadixSortSerial(vals []uint64, diff uint64) { + n := len(vals) + buf := make([]uint64, n) + src, dst := vals, buf + mask := digitMaskOf(diff) + passes := 0 + var count [256]int + for b := range 8 { + if mask&(1<>shift&0xFF]++ + } + sum := 0 + for bb, c := range count { + count[bb], sum = sum, sum+c + } + for _, v := range src { + d := v >> shift & 0xFF + dst[count[d]] = v + count[d]++ + } + src, dst = dst, src + passes++ + } + if passes%2 == 1 { + copy(vals, buf) + } +} + +// probeSortInt mirrors Sort's int path with the probe radix. +func probeSortInt(a *Array, wcap, bits int) *Array { + ints, _, _, _, _ := a.cloneData() + keys := make([]uint64, len(ints)) + hi, lo := uint64(0), ^uint64(0) + for i, v := range ints { + k := uint64(v) ^ (1 << 63) + keys[i] = k + hi |= k + lo &= k + } + probeRadixSortUint64(keys, hi^lo, wcap, bits) + for i, k := range keys { + ints[i] = int64(k ^ (1 << 63)) + } + return &Array{shape: a.Shape(), dt: Int, ints: ints} +} + +// probeSortFloatsRadix mirrors sortFloatsRadix with the probe radix. +func probeSortFloatsRadix(floats []float64, wcap, bits int) []float64 { + keys := make([]uint64, len(floats)) + finite := floats[:0] + m := 0 + hi, lo := uint64(0), ^uint64(0) + for _, v := range floats { + if v != v { + continue + } + k := floatSortKey(v) + finite = append(finite, v) + keys[m] = k + m++ + hi |= k + lo &= k + } + keys = keys[:m] + probeRadixSortUint64(keys, hi^lo, wcap, bits) + for i, k := range keys { + finite[i] = floatKeyToFloat(k) + } + for range len(floats) - m { + finite = append(finite, math.NaN()) + } + return finite +} + +// probeSortFloat mirrors Sort's float path with the probe radix. +func probeSortFloat(a *Array, wcap, bits int) *Array { + _, _, _, floats, _ := a.cloneData() + return &Array{shape: a.Shape(), dt: Float, floats: probeSortFloatsRadix(floats, wcap, bits)} +} + +// probeArgSort mirrors ArgSort's int and float paths with the probe +// radix; the fixture dtypes of the named benchmarks are int and float, +// which are the paths this mirror carries. +func probeArgSort(a *Array, wcap, bits int) *Array { + n := a.Len() + idx := make([]int, n) + m := a.materialise() + if a.dt == Int { + ints := m.ints + keys := make([]uint64, n) + hi, lo := uint64(0), ^uint64(0) + for i := range n { + k := uint64(ints[i]) ^ (1 << 63) + keys[i] = k + idx[i] = i + hi |= k + lo &= k + } + probeArgSortRadixUint64(keys, idx, hi^lo, wcap, bits) + } else { + vals := m.floats + keys := make([]uint64, n) + kept := make([]int, 0, n) + var nans []int + hi, lo := uint64(0), ^uint64(0) + for i := range n { + v := vals[i] + if v != v { + nans = append(nans, i) + continue + } + k := floatSortKey(v) + keys[i] = k + kept = append(kept, i) + hi |= k + lo &= k + } + probeArgSortRadixUint64(keys, kept, hi^lo, wcap, bits) + copy(idx, kept) + copy(idx[len(kept):], nans) + } + out := make([]int64, len(idx)) + parallelMin(len(out), copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out[i] = int64(idx[i]) + } + }) + return &Array{shape: []int{len(out)}, dt: Int, ints: out} +} + +// probeFloatSeed7 rebuilds the seed-7 float fixture of BenchmarkSort1M +// and BenchmarkArgSort1M. +func probeFloatSeed7() *Array { + a, _ := FromFloats(make([]float64, 1<<20), 1<<20) + g := NewGenerator(7) + src, _ := Floats(g, 1<<20) + copy(a.RawFloats(), src.RawFloats()) + return a +} + +// probeSortConfigs are the crew-cap and digit-width combinations the +// interleaved probe walks; wcap 0 means the machine's full worker count. +func probeSortConfigs() []struct { + name string + wcap int + bits int +} { + return []struct { + name string + wcap int + bits int + }{ + {"cap=8/12b", 8, 12}, + {"cap=16/12b", 16, 12}, + {"cap=uncapped/12b", 0, 12}, + {"cap=8/8b", 8, 8}, + {"cap=16/8b", 16, 8}, + {"cap=uncapped/8b", 0, 8}, + } +} + +// BenchmarkSortRadixConfigAB interleaves the six radix configurations +// on the fixtures of the five named sort benchmarks. +func BenchmarkSortRadixConfigAB(b *testing.B) { + fixtures := []struct { + name string + mk func() *Array + run func(*Array, int, int) *Array + }{ + {"SortInt1M", func() *Array { return benchIntsShape(1<<10, 1<<20) }, probeSortInt}, + {"Sort1M", probeFloatSeed7, probeSortFloat}, + {"SortFloat1M", func() *Array { return benchFloatsShape(1 << 20) }, probeSortFloat}, + {"ArgSortInt1M", func() *Array { return benchIntsShape(1<<10, 1<<20) }, probeArgSort}, + {"ArgSort1M", probeFloatSeed7, probeArgSort}, + } + cfgs := probeSortConfigs() + for _, f := range fixtures { + b.Run(f.name, func(b *testing.B) { + a := f.mk() + elapsed := make([]time.Duration, len(cfgs)) + for b.Loop() { + for i, c := range cfgs { + start := time.Now() + out := f.run(a, c.wcap, c.bits) + elapsed[i] += time.Since(start) + probeSink += int64(out.Len()) + } + } + for i, c := range cfgs { + b.ReportMetric(float64(elapsed[i].Nanoseconds())/float64(b.N), "ns/op-"+c.name) + } + }) + } +} + +// BenchmarkSortRadixConfigAllocs reports each configuration's +// deterministic allocation count on the SortFloat1M and SortInt1M +// fixtures, one sub-benchmark per configuration. +func BenchmarkSortRadixConfigAllocs(b *testing.B) { + for _, fx := range []struct { + name string + mk func() *Array + run func(*Array, int, int) *Array + }{ + {"SortInt1M", func() *Array { return benchIntsShape(1<<10, 1<<20) }, probeSortInt}, + {"SortFloat1M", func() *Array { return benchFloatsShape(1 << 20) }, probeSortFloat}, + } { + b.Run(fx.name, func(b *testing.B) { + a := fx.mk() + for _, c := range probeSortConfigs() { + b.Run(c.name, func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + out := fx.run(a, c.wcap, c.bits) + probeSink += int64(out.Len()) + } + }) + } + }) + } +} diff --git a/internal/core/bench_sort2_test.go b/internal/core/bench_sort2_test.go new file mode 100644 index 0000000..f456a9f --- /dev/null +++ b/internal/core/bench_sort2_test.go @@ -0,0 +1,372 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "fmt" + "testing" +) + +// Benchmarks for the ordering, search, selection and layout paths that +// misc_bench_test.go and axis_bench_test.go do not cover: the int radix, +// the mirror and take copies, the rightmost search, the bin and the +// interpolation searches, the payload walks of the mask family, the +// gather and the repeated, flipped and rolled layout kernels. + +// benchFloatsShape builds a deterministic float array of the given shape +// from the generator's uniform range. +func benchFloatsShape(shape ...int) *Array { + n := 1 + for _, d := range shape { + n *= d + } + g := NewGenerator(11) + src, err := Floats(g, n) + if err != nil { + panic(err) + } + a, err := FromFloats(src.RawFloats(), shape...) + if err != nil { + panic(err) + } + return a +} + +// benchIntsShape builds a deterministic int array of the given shape +// with values in [0, span). +func benchIntsShape(span int, shape ...int) *Array { + n := 1 + for _, d := range shape { + n *= d + } + g := NewGenerator(11) + src, err := Ints(g, n, 0, int64(span)) + if err != nil { + panic(err) + } + a, err := FromInts(src.RawInts(), shape...) + if err != nil { + panic(err) + } + return a +} + +func BenchmarkSortInt1M(b *testing.B) { + a := benchIntsShape(1<<10, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Sort(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSortIntWide1M(b *testing.B) { + a := benchIntsShape(1<<30, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Sort(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSortFloat1M(b *testing.B) { + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + if _, err := Sort(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkArgSortInt1M(b *testing.B) { + a := benchIntsShape(1<<10, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := ArgSort(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkReverseFloat1M(b *testing.B) { + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + _ = Reverse(a) + } +} + +func BenchmarkSearchSorted16k(b *testing.B) { + hay, err := Linspace(0, 1, 1<<15) + if err != nil { + b.Fatal(err) + } + needles := benchFloatsShape(1 << 13) + b.ReportAllocs() + for b.Loop() { + if _, err := SearchSorted(hay, needles); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkAssignBins1M256(b *testing.B) { + edges, err := Linspace(0, 1, 256) + if err != nil { + b.Fatal(err) + } + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + if _, err := AssignBins(a, edges); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkInterpolate128k(b *testing.B) { + xs, err := Linspace(0, 1, 1024) + if err != nil { + b.Fatal(err) + } + ys := benchFloatsShape(1024) + query := benchFloatsShape(1 << 17) + b.ReportAllocs() + for b.Loop() { + if _, err := Interpolate(xs, ys, query); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkAstypeIntToFloat1M(b *testing.B) { + a := benchIntsShape(1<<20, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Astype(a, Float); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkIsNaN1M(b *testing.B) { + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + if _, err := IsNaN(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCountNonzero1M(b *testing.B) { + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + if _, err := CountNonzero(a); err != nil { + b.Fatal(err) + } + } +} + +// benchSparseInts builds a deterministic int array whose every period-th +// element is non-zero, the sparse shape Argwhere walks. +func benchSparseInts(n, period int) *Array { + a := benchIntsShape(1<<20, n) + p := a.RawInts() + for i := range p { + if i%period != 0 { + p[i] = 0 + } + } + return a +} + +func BenchmarkArgwhere4M(b *testing.B) { + a := benchSparseInts(4096*4096, 64) + b.ReportAllocs() + for b.Loop() { + if _, err := Argwhere(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkGrid1k(b *testing.B) { + x := benchFloatsShape(1024) + y := benchFloatsShape(1024) + b.ReportAllocs() + for b.Loop() { + if _, _, err := Grid(x, y); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkIntegrate1M(b *testing.B) { + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + if _, err := Integrate(a, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCumulativeIntegrate1M(b *testing.B) { + a := benchFloatsShape(1 << 20) + b.ReportAllocs() + for b.Loop() { + if _, err := CumulativeIntegrate(a, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSliceCols1k(b *testing.B) { + a := benchFloatsShape(1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := Slice(a, 1, 100, 900); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkGather2D(b *testing.B) { + src := benchFloatsShape(1024, 1024) + idx := benchIntsShape(1024, 1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := Gather(src, 0, idx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkTake1M(b *testing.B) { + src := benchFloatsShape(1 << 20) + idx := benchIntsShape(1<<20, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Take(src, idx); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkRepeatRows(b *testing.B) { + a := benchFloatsShape(1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := Repeat(a, 2, 0); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkTileRows(b *testing.B) { + a := benchFloatsShape(1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := Tile(a, 2, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkFlipRows(b *testing.B) { + a := benchFloatsShape(2048, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := Flip(a, 0); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkFlipCols(b *testing.B) { + a := benchFloatsShape(2048, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := Flip(a, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkRollRows(b *testing.B) { + a := benchFloatsShape(2048, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := Roll(a, 7, 0); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLowerTriangle1k(b *testing.B) { + a := benchFloatsShape(1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := LowerTriangle(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkDiff2D(b *testing.B) { + a := benchFloatsShape(1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := Diff(a, 1, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkConcatRows(b *testing.B) { + a := benchFloatsShape(512, 1024) + c := benchFloatsShape(512, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := Concat(a, c, 0); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkStackRows(b *testing.B) { + a := benchFloatsShape(1024, 1024) + c := benchFloatsShape(1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := Stack(a, c, 0); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSearchSortedScaling reports the bisection against the needle +// and haystack sizes, both quoted in the report's scaling table. +func BenchmarkSearchSortedScaling(b *testing.B) { + for _, n := range []int{1024, 4096, 16384} { + hay, err := Linspace(0, 1, n) + if err != nil { + b.Fatal(err) + } + needles := benchFloatsShape(n) + b.Run(fmt.Sprintf("needles=%d/hay=%d", n, n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := SearchSorted(hay, needles); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/internal/core/bench_test.go b/internal/core/bench_test.go new file mode 100644 index 0000000..f723422 --- /dev/null +++ b/internal/core/bench_test.go @@ -0,0 +1,318 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// The matmul hot loop of inference: benchmarks keep the loop +// order honest. Compare with -bench=BenchmarkMatMul before and after any +// change to the kernels. + +func BenchmarkMatMul(b *testing.B) { + a, _ := FromFloats(make([]float64, 128*128), 128, 128) + for i := range a.Len() { + a.RawFloats()[i] = float64(i%7) - 3 + } + c, _ := FromFloats(make([]float64, 128*128), 128, 128) + for i := range c.Len() { + c.RawFloats()[i] = float64(i%5) - 2 + } + b.ReportAllocs() + for b.Loop() { + _, err := MatMul2D(a, c) + if err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkAdd(b *testing.B) { + x, _ := FromFloats(make([]float64, 100_000), 100_000) + y, _ := FromFloats(make([]float64, 100_000), 100_000) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(x, y); err != nil { + b.Fatal(err) + } + } +} + +// Parallel-vs-serial comparisons: run with -bench=. to see the +// speedup the parallel kernels deliver on the host's core count. The +// serial variants pin the worker count to one. + +func BenchmarkMatMulParallel(b *testing.B) { + benchMatMulN(b, 512) +} + +func BenchmarkMatMulSerial(b *testing.B) { + SetNumCPU(1) + defer SetNumCPU(0) + benchMatMulN(b, 512) +} + +func benchMatMulN(b *testing.B, n int) { + a, _ := FromFloats(make([]float64, n*n), n, n) + c, _ := FromFloats(make([]float64, n*n), n, n) + b.ResetTimer() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkAddParallel(b *testing.B) { benchAddN(b, 1<<22) } + +func BenchmarkAddSerial(b *testing.B) { + SetNumCPU(1) + defer SetNumCPU(0) + benchAddN(b, 1<<22) +} + +func benchAddN(b *testing.B, n int) { + a, _ := FromFloats(make([]float64, n), n) + c, _ := FromFloats(make([]float64, n), n) + b.ResetTimer() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMatMulFloat32 exercises the production inference dtype at +// sizes the embedding/table workloads actually see. +func BenchmarkMatMulFloat32_128(b *testing.B) { benchMatMulF32N(b, 128) } + +func BenchmarkMatMulFloat32_512(b *testing.B) { benchMatMulF32N(b, 512) } + +func benchMatMulF32N(b *testing.B, n int) { + af, _ := FromFloats(make([]float64, n*n), n, n) + bf, _ := FromFloats(make([]float64, n*n), n, n) + a, _ := Astype(af, Float32) + c, _ := Astype(bf, Float32) + b.ResetTimer() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BroadcastTo sits inside every bias add and normalisation, so its walk +// is on the training hot path. The bias-shaped case (1, C, 1, 1) to +// (N, C, H, W) is the one the run-fill rewrite targets. +func BenchmarkBroadcastToBiasNCHW(b *testing.B) { + src, _ := FromFloats(make([]float64, 64), 1, 64, 1, 1) + b.ReportAllocs() + for b.Loop() { + if _, err := BroadcastTo(src, 8, 64, 28, 28); err != nil { + b.Fatal(err) + } + } +} + +// The trailing-replication case (1, C) to (N, C): the run is the whole +// channel block, the other extreme from a fully strided walk. +func BenchmarkBroadcastToBiasNC(b *testing.B) { + src, _ := FromFloats(make([]float64, 256), 1, 256) + b.ReportAllocs() + for b.Loop() { + if _, err := BroadcastTo(src, 128, 256); err != nil { + b.Fatal(err) + } + } +} + +// The non-broadcast expansion (prepending leading dimensions only): +// every element is distinct, so this is the worst case for the rewrite. +func BenchmarkBroadcastToPrepend(b *testing.B) { + src, _ := FromFloats(make([]float64, 128*256), 128, 256) + b.ReportAllocs() + for b.Loop() { + if _, err := BroadcastTo(src, 4, 128, 256); err != nil { + b.Fatal(err) + } + } +} + +// Reduction and vector-shaped hot paths: Dot and Min walk a single +// payload, MatVec is the 2-D×1-D product. Run with -bench before and +// after any change to reduce.go or the vector kernels in mat.go. + +func benchVecF64(b *testing.B, n int) (*Array, *Array) { + a, _ := FromFloats(make([]float64, n), n) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%11) - 5 + } + c, _ := FromFloats(make([]float64, n), n) + for i := range c.RawFloats() { + c.RawFloats()[i] = float64(i%7) - 3 + } + return a, c +} + +func BenchmarkDot(b *testing.B) { + a, c := benchVecF64(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Dot(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMinMax(b *testing.B) { + a, _ := benchVecF64(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Max(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMatVec(b *testing.B) { + m, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range m.RawFloats() { + m.RawFloats()[i] = float64(i%13) - 6 + } + v, _ := FromFloats(make([]float64, 512), 512) + for i := range v.RawFloats() { + v.RawFloats()[i] = float64(i%7) - 3 + } + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(m, v); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkTranspose(b *testing.B) { + a, _ := FromFloats(make([]float64, 256*256), 256, 256) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%9) - 4 + } + b.ReportAllocs() + for b.Loop() { + Transpose(a) + } +} + +func BenchmarkRowCol(b *testing.B) { + a, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%9) - 4 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Row(a, 7); err != nil { + b.Fatal(err) + } + if _, err := Col(a, 7); err != nil { + b.Fatal(err) + } + } +} + +// ColTall isolates the column gather from output allocation: the tall +// shape makes the per-row strided read the dominant cost. +func BenchmarkColTall(b *testing.B) { + a, _ := FromFloats(make([]float64, 65536*8), 65536, 8) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%9) - 4 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Col(a, 3); err != nil { + b.Fatal(err) + } + } +} + +// Along-dim reductions: Norm folds with math.Pow per element, Prod is +// the multiply twin of the axis sum. Run with -bench before and after +// any change to reduction2.go. +func BenchmarkNorm2D(b *testing.B) { + a, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%9) - 4 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Norm(a, 2, 1, false); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkProd2D(b *testing.B) { + a, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%9) + 1 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Prod(a, 1, false); err != nil { + b.Fatal(err) + } + } +} + +// CumSum walks the prefix scan along the trailing dimension. +func BenchmarkCumSum2D(b *testing.B) { + a, _ := FromFloats(make([]float64, 512*512), 512, 512) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%9) - 4 + } + b.ReportAllocs() + for b.Loop() { + if _, err := CumSum(a, 1); err != nil { + b.Fatal(err) + } + } +} + +// Element-wise math walks the payload through a real function; Exp and +// Sqrt are the training hot ones. Run with -bench before and after any +// change to mathfunc.go. + +func benchmarkRealFunc(b *testing.B, f func(*Array) (*Array, error)) { + a, _ := FromFloats(make([]float64, 1<<20), 1<<20) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%97) + 1 + } + b.ReportAllocs() + for b.Loop() { + if _, err := f(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkExp(b *testing.B) { benchmarkRealFunc(b, Exp) } + +func BenchmarkSqrt(b *testing.B) { benchmarkRealFunc(b, Sqrt) } + +func BenchmarkProd1D(b *testing.B) { + a, _ := FromFloats(prodFixture(1<<20), 1<<20) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if _, err := Prod(a, 0, false); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNorm1D(b *testing.B) { + a, _ := FromFloats(prodFixture(1<<20), 1<<20) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if _, err := Norm(a, 2, 0, false); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/bessel_climb_pin_test.go b/internal/core/bessel_climb_pin_test.go new file mode 100644 index 0000000..151c06d --- /dev/null +++ b/internal/core/bessel_climb_pin_test.go @@ -0,0 +1,81 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// Where BesselIn's upward recurrence takes over from the Miller walk, +// and the accuracy the handover must keep. + +// TestBesselInClimbBoundary pins where BesselIn stops climbing. The +// upward recurrence is the only affordable route to an order far below a +// large argument, but its seeds are besselI0 and besselI1, which below +// the asymptotic branch are polynomial fits worth about 2.5e-8: a short +// climb over a small argument multiplies that weak seed by the condition +// number and lands coarser than the downward walk. Below +// besselIUpwardMinX the value must therefore be the Miller result +// unchanged, and from there up it must agree with Miller, which is the +// independent route and the accuracy reference. +func TestBesselInClimbBoundary(t *testing.T) { + same := func(t *testing.T, n int, x float64) { + t.Helper() + a := mustFloats(t, []float64{x}, 1) + got, err := BesselIn(n, a) + if err != nil { + t.Fatalf("BesselIn(%d, %g): %v", n, x, err) + } + want := besselInMiller(n, x) + if math.Float64bits(got.FloatAt(0)) != math.Float64bits(want) { + t.Fatalf("BesselIn(%d, %g) = %v, the downward walk gives %v", n, x, got.FloatAt(0), want) + } + } + + // Under the bound the climb must not be taken at all: bit-for-bit + // the downward walk, which is what the pre-climb library returned. + for _, c := range []struct { + n int + x float64 + }{ + {2, 1}, {3, 2.25}, {3, 3}, {3, 11.9}, {5, 8}, {17, 11.999}, + } { + same(t, c.n, c.x) + } + + // From the bound up the climb is taken, and it must agree with the + // downward walk to a few ulps: the asymptotic seeds are near full + // precision there, so the condition number no longer amplifies a + // weak seed. Miller is affordable at these arguments; the ones where + // it is not are pinned by their cost in einsum_norm_special_pins_test.go. + for _, c := range []struct { + n int + x float64 + }{ + {3, 12}, {3, 12.1}, {3, 30}, {5, 30}, {8, 64}, {20, 400}, + } { + a := mustFloats(t, []float64{c.x}, 1) + got, err := BesselIn(c.n, a) + if err != nil { + t.Fatalf("BesselIn(%d, %g): %v", c.n, c.x, err) + } + want := besselInMiller(c.n, c.x) + if math.IsInf(want, 1) { + // Past the float64 range both routes report the overflow; + // there is no accuracy left to compare. + if !math.IsInf(got.FloatAt(0), 1) { + t.Errorf("BesselIn(%d, %g) = %v, the downward walk overflows to +Inf", c.n, c.x, got.FloatAt(0)) + } + continue + } + if want == 0 { + t.Fatalf("BesselIn(%d, %g): the reference is zero, the case is not usable", c.n, c.x) + } + if rel := math.Abs(got.FloatAt(0)-want) / math.Abs(want); rel > 1e-13 { + t.Errorf("BesselIn(%d, %g) = %v against the downward walk %v: relative error %.3g", + c.n, c.x, got.FloatAt(0), want, rel) + } + } +} diff --git a/internal/core/bessel_realorder_accuracy_test.go b/internal/core/bessel_realorder_accuracy_test.go new file mode 100644 index 0000000..9f9c0cb --- /dev/null +++ b/internal/core/bessel_realorder_accuracy_test.go @@ -0,0 +1,330 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + "testing" +) + +// BesselJRealOrder against an independent high-precision reference: the +// Frobenius series summed in math/big at a working size far past the +// float64 grid, with the gamma function from a Stirling series whose +// Bernoulli coefficients are generated exactly (big.Rat) and whose +// argument is lifted by the recurrence before the series applies. The +// reference shares no code with the implementation under test. + +const besselRefPrec = 800 + +func bf(x float64) *big.Float { + return new(big.Float).SetPrec(besselRefPrec).SetFloat64(x) +} + +func bfFromInt(n int64) *big.Float { + return new(big.Float).SetPrec(besselRefPrec).SetInt64(n) +} + +func bfMul(a, b *big.Float) *big.Float { + return new(big.Float).SetPrec(besselRefPrec).Mul(a, b) +} + +func bfQuo(a, b *big.Float) *big.Float { + return new(big.Float).SetPrec(besselRefPrec).Quo(a, b) +} + +func bfAdd(a, b *big.Float) *big.Float { + return new(big.Float).SetPrec(besselRefPrec).Add(a, b) +} + +func bfSub(a, b *big.Float) *big.Float { + return new(big.Float).SetPrec(besselRefPrec).Sub(a, b) +} + +// bfExp evaluates e^x by the halving ladder: the argument halves until a +// Taylor run converges past the working resolution, then the squarings +// unwind. Each squaring costs about one ulp, and the largest arguments +// here spend a dozen of them out of 800 bits. +func bfExp(x *big.Float) *big.Float { + if x.Sign() == 0 { + return bf(1) + } + halved := new(big.Float).SetPrec(besselRefPrec).Set(x) + squarings := 0 + for halved.MantExp(nil) > 0 { + halved.Quo(halved, bfFromInt(2)) + squarings++ + } + term := bf(1) + sum := bf(1) + for k := int64(1); ; k++ { + term = bfMul(term, bfQuo(halved, bfFromInt(k))) + sum.Add(sum, term) + if term.MantExp(nil) < -int(besselRefPrec)-10 { + break + } + } + for range squarings { + sum.Mul(sum, sum) + } + return sum +} + +// bfLn evaluates ln y for y > 0: the mantissa-exponent split turns the +// problem into ln m for m in [0.5, 1), which the doubled atanh series +// covers through t = (m−1)/(m+1) with |t| ≤ 1/3, plus e·ln 2 with ln 2 +// itself from the atanh series at t = 1/3. +func bfLn(y *big.Float) *big.Float { + m := new(big.Float).SetPrec(besselRefPrec) + e := y.MantExp(m) + t := bfQuo(bfSub(m, bf(1)), bfAdd(m, bf(1))) + ln2 := bfMul(bfFromInt(2), bfAtanh(bfQuo(bf(1), bfFromInt(3)))) + return bfAdd(bfMul(bfFromInt(2), bfAtanh(t)), bfMul(bfFromInt(int64(e)), ln2)) +} + +// bfAtanh sums atanh(t) = t + t³/3 + t⁵/5 + ... for |t| ≤ 1/3. +func bfAtanh(t *big.Float) *big.Float { + power := new(big.Float).SetPrec(besselRefPrec).Set(t) // t^k for odd k + sum := bf(0) + for k := int64(1); ; k += 2 { + term := bfQuo(power, bfFromInt(k)) + sum.Add(sum, term) + power = bfMul(power, bfMul(t, t)) + if term.MantExp(nil) < -int(besselRefPrec)-10 { + break + } + } + return sum +} + +// bfAtan sums atan(t) = t − t³/3 + t⁵/5 − ... for |t| ≤ 1/5. +func bfAtan(t *big.Float) *big.Float { + power := new(big.Float).SetPrec(besselRefPrec).Set(t) // t^k for odd k + sum := bf(0) + for k := int64(1); ; k += 2 { + term := bfQuo(power, bfFromInt(k)) + if (k/2)%2 == 1 { + term = bfSub(bf(0), term) + } + sum.Add(sum, term) + power = bfMul(power, bfMul(t, t)) + if term.MantExp(nil) < -int(besselRefPrec)-10 { + break + } + } + return sum +} + +// bfPi evaluates π by Machin's formula π = 16·atan(1/5) − 4·atan(1/239), +// so the reference carries no decimal constant beyond the rationals the +// formula itself names. +func bfPi() *big.Float { + fifth := bfQuo(bf(1), bf(5)) + two39 := bfQuo(bf(1), bf(239)) + return bfSub(bfMul(bfFromInt(16), bfAtan(fifth)), bfMul(bfFromInt(4), bfAtan(two39))) +} + +// stirlingTerms generates the coefficients cₙ = B₂ₙ/(2n(2n−1)) of the +// Stirling series exactly, from the Bernoulli recurrence +// Bₘ = −Σ C(m,k)·Bₖ/(m−k+1) over exact rationals. +func stirlingTerms(count int) []*big.Rat { + bernoulli := []*big.Rat{big.NewRat(1, 1)} + for m := 1; len(bernoulli) <= 2*count; m++ { + b := new(big.Rat) + for k := 0; k < m; k++ { + t := new(big.Rat).SetInt(new(big.Int).Binomial(int64(m), int64(k))) + t.Mul(t, bernoulli[k]) + t.Quo(t, new(big.Rat).SetInt64(int64(m-k+1))) + b.Sub(b, t) + } + bernoulli = append(bernoulli, b) + } + terms := make([]*big.Rat, count) + for n := 1; n <= count; n++ { + den := new(big.Rat).SetInt64(int64(2*n) * int64(2*n-1)) + terms[n-1] = new(big.Rat).Quo(bernoulli[2*n], den) + } + return terms +} + +// besselRefStirling carries 64 coefficients, far past the truncation +// point the series reaches at the lifted argument 64. +var besselRefStirling = stirlingTerms(64) + +// lnGammaBig evaluates ln Γ(z) for z > 0: the recurrence +// ln Γ(z) = ln Γ(z+1) − ln z lifts z above 64, the Stirling series runs +// to its optimal truncation (the terms turn around and grow), and the +// lifts unwind. The whole answer stays in logarithms, which is all the +// series reference consumes. +func lnGammaBig(zIn float64) *big.Float { + z := bf(zIn) + lifts := bf(0) + for z.MantExp(nil) < 7 { // z < 64 + lifts = bfAdd(lifts, bfLn(z)) + z = bfAdd(z, bf(1)) + } + lnZ := bfLn(z) + power := bf(1) // z^(2n−2), so the n-th term is cₙ/(power·z) = cₙ/z^(2n−1) + sum := bf(0) + prevAbs := bf(1) + for n := 1; n <= len(besselRefStirling); n++ { + coef := new(big.Float).SetPrec(besselRefPrec).SetRat(besselRefStirling[n-1]) + term := bfQuo(coef, bfMul(power, z)) + absTerm := new(big.Float).SetPrec(besselRefPrec).Abs(term) + if absTerm.Cmp(prevAbs) > 0 { + break // past the optimal truncation + } + prevAbs = absTerm + sum.Add(sum, term) + power = bfMul(power, bfMul(z, z)) + } + // ln Γ(z) = (z − ½)·ln z − z + ½·ln(2π) + Σ cₙ/z^(2n−1). + halfLnTwoPi := bfMul(bf(0.5), bfLn(bfMul(bf(2), bfPi()))) + total := bfMul(bfSub(z, bf(0.5)), lnZ) + total = bfAdd(bfSub(total, z), halfLnTwoPi) + total = bfAdd(total, sum) + return bfSub(total, lifts) +} + +// besselJRef evaluates J_ν(x) by the Frobenius series in extended +// precision: +// +// J_ν(x) = Σ_{k≥0} (−1)^k (x/2)^{2k+ν} / (k!·Γ(k+ν+1)), +// +// the k = 0 term from the exponential of ν·ln(x/2) − ln Γ(ν+1) and every +// later term stepping along the ratio −(x/2)²/(k(k+ν)). The second +// result reports whether the series' own cancellation guard passed: the +// peak term against the result must stay inside the working headroom, +// past which the reference would print noise. +func besselJRef(nu, x float64) (j *big.Float, ok bool) { + half := bfQuo(bf(x), bf(2)) + q := bfMul(half, half) + t := bfExp(bfSub(bfMul(bf(nu), bfLn(half)), lnGammaBig(nu+1))) + sum := bf(0) + sum.Add(sum, t) + peak := new(big.Float).SetPrec(besselRefPrec).Abs(t) + for k := int64(1); k < 4000; k++ { + t = bfQuo(bfMul(t, q), bfMul(bfFromInt(k), bfAdd(bfFromInt(k), bf(nu)))) + t.Neg(t) + sum.Add(sum, t) + if abs := new(big.Float).SetPrec(besselRefPrec).Abs(t); abs.Cmp(peak) > 0 { + peak = abs + } + // The terms fall monotonically once k passes x/2; the gap + // check ends the run when the term drops below the grid. + if t.MantExp(nil) < sum.MantExp(nil)-int(besselRefPrec)-20 { + break + } + } + limit := new(big.Float).SetPrec(besselRefPrec).SetMantExp(bf(1), int(besselRefPrec)-120) + if peak.Cmp(bfMul(limit, new(big.Float).SetPrec(besselRefPrec).Abs(sum))) > 0 { + return sum, false + } + return sum, true +} + +// TestBesselJRefSelfCheck holds the extended-precision series against +// the exact half-integer closed forms, so the sweep below stands on a +// reference proven at the points where a closed form exists. +func TestBesselJRefSelfCheck(t *testing.T) { + for _, x := range []float64{0.5, 3, 12, 20, 64} { + ref, ok := besselJRef(0.5, x) + if !ok { + t.Fatalf("the guard tripped at x=%g", x) + } + want, _ := ref.Float64() + if d := math.Abs(jHalf(0.5, x)-want) / math.Abs(want); d > 1e-13 { + t.Fatalf("the reference series gives %.17g at (0.5, %g) against the closed form %.17g (relative %.3g)", + want, x, jHalf(0.5, x), d) + } + ref15, _ := besselJRef(1.5, 20) + w15, _ := ref15.Float64() + if d := math.Abs(jHalf(1.5, 20)-w15) / math.Abs(w15); d > 1e-13 { + t.Fatalf("the reference series gives %.17g at (1.5, 20) against the closed form %.17g (relative %.3g)", + w15, jHalf(1.5, 20), d) + } + } +} + +// TestBesselJRealOrderAgainstBigFloat sweeps the whole regime map against +// the extended-precision series: the small-argument series branch, the +// crossover region where the series pays its cancellation and the +// expansion truncates, the climb above the crossover, and the downward +// Miller walk past the argument. +func TestBesselJRealOrderAgainstBigFloat(t *testing.T) { + xs := []float64{0.02, 0.1, 0.5, 1, 3, 7, 11, 12, 12.5, 13, 14, 15, 16, 18, 20, 25, 32, 48, 64} + nus := []float64{0.1, 0.3, 0.5, 0.7, 0.9, 1.1, 1.7, 2.3, 2.9, 3.5, 4.9, 7.3, 11.5, 12.1, 15.7, 20.5, 27.9, 33.3, 45.5, 72.5} + type worst struct { + nu, x, err float64 + } + var top [5]worst + skipped := 0 + for _, nu := range nus { + for _, x := range xs { + ref, ok := besselJRef(nu, x) + if !ok { + skipped++ + continue + } + got, err := BesselJRealOrder(nu, x) + if err != nil { + t.Fatalf("BesselJRealOrder(%g, %g): %v", nu, x, err) + } + want, _ := ref.Float64() + d := math.Abs(got-want) / math.Abs(want) + if d > top[0].err { + top[0] = worst{nu, x, d} + for k := 1; k < len(top) && top[k-1].err > top[k].err; k++ { + top[k-1], top[k] = top[k], top[k-1] + } + } + } + } + if skipped > 0 { + t.Logf("reference guard skipped %d points", skipped) + } + for i := len(top) - 1; i >= 0; i-- { + if top[i].err == 0 { + continue + } + t.Logf("worst %d: nu=%g x=%g relative %.3e", i+1, top[i].nu, top[i].x, top[i].err) + } + // The measured ceiling sits at the crossover band, where the + // expansion seeds truncate: the sweep's worst point lands at + // x = 12 within a few parts in 1e-11, an order below it everywhere + // else, and the series side of the crossover measures no better + // there (both sides sit at 5e-11 worst case at x = 12, so the + // boundary stays where it is). + if top[0].err > 1e-10 { + t.Errorf("the worst relative error %.3e at (nu=%g, x=%g) is past the measured crossover ceiling", + top[0].err, top[0].nu, top[0].x) + } +} + +// TestBesselJRealOrderMillerCaptures pins the fractional Miller walk's +// captures: the walk descends in floats, and a capture tested by float +// equality misses by an ulp for generic fractions, answering from an +// unseeded slot. These orders sit exactly where that used to bite: the +// answers must stay finite and accurate against the extended-precision +// series. +func TestBesselJRealOrderMillerCaptures(t *testing.T) { + for _, c := range [][2]float64{{12.1, 12}, {33.3, 12}, {45.5, 16}, {72.5, 25}, {13.7, 12.5}, {1.9, 20}} { + nu, x := c[0], c[1] + ref, ok := besselJRef(nu, x) + if !ok { + t.Fatalf("the reference guard tripped at (%g, %g)", nu, x) + } + want, _ := ref.Float64() + got, err := BesselJRealOrder(nu, x) + if err != nil { + t.Fatalf("BesselJRealOrder(%g, %g): %v", nu, x, err) + } + if math.IsNaN(got) || math.IsInf(got, 0) { + t.Fatalf("BesselJRealOrder(%g, %g) = %v, want a finite value near %.17g", nu, x, got, want) + } + if d := math.Abs(got-want) / math.Abs(want); d > 1e-10 { + t.Fatalf("BesselJRealOrder(%g, %g) = %.17g, want %.17g (relative %.3g)", nu, x, got, want, d) + } + } +} diff --git a/internal/core/besseljreal_test.go b/internal/core/besseljreal_test.go new file mode 100644 index 0000000..7a12d05 --- /dev/null +++ b/internal/core/besseljreal_test.go @@ -0,0 +1,180 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// jHalf evaluates the closed forms of the half-integer orders +// J₁/₂, J₃/₂ and J₅/₂, the exact referents no other Bessel test here +// enjoys. +func jHalf(nu, x float64) float64 { + s := math.Sqrt(2 / (math.Pi * x)) + sin, cos := math.Sincos(x) + switch nu { + case 0.5: + return s * sin + case 1.5: + return s * (sin/x - cos) + case 2.5: + return s * ((3/(x*x)-1)*sin - 3*cos/x) + } + panic("jHalf: unsupported order") +} + +// climbHalf climbs the closed forms from orders 1/2 and 3/2 up to +// 0.5+m by the exact three-term recurrence, the referent for orders +// the closed forms themselves do not cover. +func climbHalf(m int, x float64) float64 { + jm := jHalf(0.5, x) + j := jHalf(1.5, x) + for k := 1; k < m; k++ { + jm, j = j, 2*(0.5+float64(k))/x*j-jm + } + return j +} + +// TestBesselJRealOrderHalfIntegers holds every regime against the +// closed forms: the series below the crossover, the climb seeded from +// the expansion above it, and a phase reduction at an argument whose +// plain float64 phase would already be losing digits. +func TestBesselJRealOrderHalfIntegers(t *testing.T) { + // The sweep starts at 0.5: below that the closed forms themselves + // cancel catastrophically and stop being referents. The small-x + // behaviour is pinned separately, against the series' own leading + // terms. + xs := []float64{0.5, 1, 2, 7, 12.1, 14.9, 15, 20, 40, 100, 1000} + for _, nu := range []float64{0.5, 1.5, 2.5} { + for _, x := range xs { + got, err := BesselJRealOrder(nu, x) + if err != nil { + t.Fatalf("BesselJRealOrder(%g, %g): %v", nu, x, err) + } + want := jHalf(nu, x) + if d := math.Abs(got-want) / math.Abs(want); d > 5e-12 { + t.Fatalf("BesselJRealOrder(%g, %g) = %.17g, want %.17g (relative %.3g)", nu, x, got, want, d) + } + } + } + // At x = 0.05 the true J₅/₂ agrees with the series' first three + // terms to five parts in 1e10, the third term being the first one + // the referent omits. + const x, nu = 0.05, 2.5 + got, err := BesselJRealOrder(nu, x) + if err != nil { + t.Fatal(err) + } + half := 0.5 * x + t0 := math.Exp(nu*math.Log(half) - lnGammaReal(nu+1)) + q := half * half + referent := t0 * (1 - q/(nu+1)*(1-q/(2*(nu+2)))) + if d := math.Abs(got-referent) / t0; d > 1e-10 { + t.Fatalf("BesselJRealOrder(%g, %g) = %.17g, want %.17g (relative %.3g)", nu, x, got, referent, d) + } +} + +// TestBesselJRealOrderMiller pins the fractional Miller walk: an order +// past the argument at an argument past the crossover, against the +// exact closed forms climbed up by the recurrence. +func TestBesselJRealOrderMiller(t *testing.T) { + const x = 15 + got, err := BesselJRealOrder(20.5, x) + if err != nil { + t.Fatalf("BesselJRealOrder: %v", err) + } + want := climbHalf(20, x) + if d := math.Abs(got-want) / math.Abs(want); d > 1e-10 { + t.Fatalf("BesselJRealOrder(20.5, 15) = %.17g, want %.17g (relative %.3g)", got, want, d) + } +} + +// TestBesselJRealOrderRecurrence checks the three-term recurrence +// across the map of regimes, the generic verifier no closed form can +// replace: every pair of orders the test touches lives on one walk. +func TestBesselJRealOrderRecurrence(t *testing.T) { + cases := [][2]float64{ + {1.3, 15}, {2.3, 15}, {7.7, 40}, {12.3, 100}, + {25.5, 15}, {60.5, 40}, {1.7, 1000}, + } + for _, c := range cases { + nu, x := c[0], c[1] + jm, err := BesselJRealOrder(nu-1, x) + if err != nil { + t.Fatalf("BesselJRealOrder(%g, %g): %v", nu-1, x, err) + } + j, err := BesselJRealOrder(nu, x) + if err != nil { + t.Fatalf("BesselJRealOrder(%g, %g): %v", nu, x, err) + } + jp, err := BesselJRealOrder(nu+1, x) + if err != nil { + t.Fatalf("BesselJRealOrder(%g, %g): %v", nu+1, x, err) + } + scale := math.Max(math.Abs(jm), math.Max(math.Abs(jp), math.Abs(2*nu/x*j))) + residual := math.Abs(jm + jp - 2*nu/x*j) + // Near the crossover the expansion's truncation caps the seeds + // at about eight digits; further out it truncates past the + // rounding floor and the residual follows it down. + tol := 1e-11 + if x < 20 { + tol = 5e-8 + } + if residual > tol*scale { + t.Fatalf("the recurrence residual at (ν = %g, x = %g) is %.3g against scale %.3g", nu, x, residual, scale) + } + } +} + +// TestBesselJRealOrderIntegerDelegate pins the integer delegation: an +// exact integer order returns the integer algorithm's own bits, and an +// order a whisper away from it lands within a whisper of the same +// value through the general route. +func TestBesselJRealOrderIntegerDelegate(t *testing.T) { + for _, x := range []float64{2, 20} { + for _, n := range []int{0, 3, 17} { + got, err := BesselJRealOrder(float64(n), x) + if err != nil { + t.Fatalf("BesselJRealOrder(%d, %g): %v", n, x, err) + } + if want := BesselJ(n, x); got != want { + t.Fatalf("BesselJRealOrder(%d, %g) = %.17g, want the integer %.17g", n, x, got, want) + } + } + } + near, err := BesselJRealOrder(3+1e-12, 20) + if err != nil { + t.Fatal(err) + } + if d := math.Abs(near - BesselJ(3, 20)); d > 1e-11 { + t.Fatalf("an order 1e-12 off the integer moved J by %g", d) + } +} + +// TestBesselJRealOrderSeriesVsClimb crosses the two routes against +// each other just above the crossover, where both still carry roughly +// nine digits: the series pushed past its regime and the climb from +// the expansion must agree to that shared quality. +func TestBesselJRealOrderSeriesVsClimb(t *testing.T) { + const nu, x = 2.3, 12.1 + got, err := BesselJRealOrder(nu, x) + if err != nil { + t.Fatal(err) + } + want := besselJSeriesReal(nu, x) + if d := math.Abs(got-want) / math.Abs(want); d > 1e-8 { + t.Fatalf("the climb gives %.17g against the series %.17g (relative %.3g)", got, want, d) + } +} + +// TestBesselJRealOrderErrors pins the domain: a negative or NaN order +// and a non-positive or NaN argument are errors naming themselves. +func TestBesselJRealOrderErrors(t *testing.T) { + for _, c := range [][2]float64{{-1, 2}, {math.NaN(), 2}, {1, 0}, {1, -2}, {math.NaN(), math.NaN()}} { + if _, err := BesselJRealOrder(c[0], c[1]); err == nil { + t.Fatalf("BesselJRealOrder(%g, %g): want an error", c[0], c[1]) + } + } +} diff --git a/internal/core/besselmod.go b/internal/core/besselmod.go new file mode 100644 index 0000000..781d3cb --- /dev/null +++ b/internal/core/besselmod.go @@ -0,0 +1,300 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "math" + +// Modified Bessel functions of the first and second kind, the workhorses +// of heat conduction, waveguides and filtered noise. I₀ and I₁ use the +// polynomial fits below 3.75 (Abramowitz & Stegun 9.8.1 and 9.8.2) and +// their integral representations beyond, K₀ and K₁ integrate their +// kernels directly, and integer orders follow by the stable direction +// of the recurrence: upward for I, downward for K. +// All are element-wise over real arrays; ints and float32 promote. + +// BesselI0 returns the modified Bessel function of the first kind of +// order zero, I₀(x), of each element. I₀ is even with I₀(0) = 1. +func BesselI0(a *Array) (*Array, error) { + return a.realFunc("BesselI0", besselI0) +} + +// BesselI1 returns the modified Bessel function of the first kind of +// order one, I₁(x), of each element. I₁ is odd with I₁(0) = 0. +func BesselI1(a *Array) (*Array, error) { + return a.realFunc("BesselI1", besselI1) +} + +// BesselK0 returns the modified Bessel function of the second kind of +// order zero, K₀(x), of each element, defined for x > 0. K₀ diverges +// logarithmically at 0 and returns +Inf there. +func BesselK0(a *Array) (*Array, error) { + return a.realFunc("BesselK0", besselK0) +} + +// BesselK1 returns the modified Bessel function of the second kind of +// order one, K₁(x), of each element, defined for x > 0. K₁ diverges +// like 1/x at 0 and returns +Inf there. +func BesselK1(a *Array) (*Array, error) { + return a.realFunc("BesselK1", besselK1) +} + +// BesselIn returns the modified Bessel function of the first kind of +// integer order n, Iₙ(x), of each element. The recurrence runs in its +// stable direction: orders well below |x| climb upward from the +// quadrature I₀ and I₁ at O(n) cost, and orders comparable to or above +// it run Miller's downward algorithm, which is where the upward climb +// would amplify the parasitic K component the seeds carry. The +// downward walk starts above the turning point at order |x|, so it +// costs O(n + |x|) steps: for an order far below a large argument that +// is the difference between a microsecond and minutes, which is why +// the climb exists. +func BesselIn(n int, a *Array) (*Array, error) { + if n < 0 { + return nil, errf("BesselIn: order must be ≥ 0, got %d", n) + } + return a.realFunc("BesselIn", func(x float64) float64 { + if n == 0 { + return besselI0(x) + } + if n == 1 { + return besselI1(x) + } + // The relative contamination the climb suffers grows roughly + // like e^{n²/x}, and the seeds it climbs from are only as good + // as besselI0 and besselI1 are there: below 3.75 those are + // polynomial fits worth about 2.5e-8, so a short climb over a + // small argument multiplies a weak seed by the condition number + // and lands coarser than the downward walk, which renormalises + // through the same seed once. The climb is therefore taken only + // where its seeds are strong, from besselIUpwardMinX up, and the + // split follows n² ≤ 4x inside that: above it the downward walk + // is both accurate and affordable, since the order is then + // comparable to the argument and O(n + |x|) steps is the + // answer's own price. + ax := math.Abs(x) + if ax >= besselIUpwardMinX && float64(n)*float64(n) <= 4*ax { + return besselIUpward(n, x) + } + return besselInMiller(n, x) + }) +} + +// besselIUpwardMinX is the argument from which the upward climb beats +// Miller's downward walk on accuracy as well as on cost: the asymptotic +// I₀ and I₁ it climbs from are near full precision there, while the +// polynomial pair below 3.75 is not. Miller's own cost is O(n + |x|), so +// staying with it under this bound costs nothing that matters. +const besselIUpwardMinX = 12 + +// besselIUpward climbs I₀ and I₁ to order n by the upward recurrence +// Iₖ₊₁ = Iₖ₋₁ − (2k/x)·Iₖ (the modified kind carries the minus sign +// where J carries a plus), the stable direction while the order stays +// well below |x|: the parasitic K component the seeds carry decays +// relative to I there, whereas near and above the argument the same +// recurrence amplifies it, which is what besselInMiller's downward walk +// avoids. The quadrature seeds carry the absolute scale and the parity, +// so a negative argument needs no extra handling. The caller guarantees +// x ≠ 0, 2 ≤ n, n² ≤ 4|x| and |x| ≥ besselIUpwardMinX. +func besselIUpward(n int, x float64) float64 { + i0, i1 := besselI0(x), besselI1(x) + // The quadrature overflows to Inf once |x| passes about 714, where + // the true value overflows too: report that instead of letting the + // recurrence fold Inf - Inf into NaN. The sign follows the parity + // Iₙ(−x) = (−1)ⁿ Iₙ(x). + if math.IsInf(i0, 0) || math.IsInf(i1, 0) { + if x < 0 && n%2 == 1 { + return math.Inf(-1) + } + return math.Inf(1) + } + for k := 1; k < n; k++ { + i0, i1 = i1, i0-2*float64(k)/x*i1 + } + return i1 +} + +// besselInMiller evaluates I_n by the downward Miller recurrence. The +// upward recurrence for I is unstable once the order passes the +// argument (the terms subtract nearly equal neighbours), while the +// downward direction amplifies no rounding. The arbitrary seed scale +// is removed against the exact I_0 at the end. +func besselInMiller(n int, x float64) float64 { + if x == 0 { + if n == 0 { + return 1 + } + return 0 + } + // |x| keeps the start above n for negative arguments too: a start + // at or below n would never pass the order on the way down and + // renormalise against a garbage seed. + start := n + int(math.Abs(x)) + 40 + jp, j := 0.0, 1.0 // j_{k+1}, j_k, seeded at k = start + jn := 0.0 + for k := start; k >= 1; k-- { + if k == n { + jn = j + } + // Downward recurrence: j_{k−1} = j_{k+1} + 2k/x·j_k. + jp, j = j, jp+2*float64(k)/x*j + if aj := math.Abs(j); aj > 1e200 { + // I grows monotonically down the walk, and the unscaled + // seed overflows for tiny x or high orders whose true + // value is representable; a common rescale cancels in the + // final ratio. + jp /= aj + j /= aj + jn /= aj + } + } + // j now holds the unscaled I_0. The quotient is formed before the + // multiplication: above about x = 300 both walk values pass 1e100, + // and the product I₀·jₙ overflows even though the answer, which is + // the quotient times the seed scale, is an ordinary number. Taking + // the product first turned every such evaluation into +Inf. + return besselI0(x) * (jn / j) +} + +// BesselKn returns the modified Bessel function of the second kind of +// integer order n, Kₙ(x), of each element, by the upward recurrence +// Kₙ₊₁ = Kₙ₋₁ + 2n/x·Kₙ from K₀ and K₁, which is the stable direction +// for the second-kind functions. K is defined for x > 0 at every +// order: outside that domain the answer is not order-dependent but +// undefined, so any element that is NaN or non-positive is an error +// naming the element, the same contract BesselY enforces at one +// point, never a silent NaN or a spurious +Inf. +func BesselKn(n int, a *Array) (*Array, error) { + if n < 0 { + return nil, errf("BesselKn: order must be ≥ 0, got %d", n) + } + if a.dt == Complex { + return nil, errf("BesselKn: complex arrays are not supported") + } + for i := range a.Len() { + if x := a.floatAt(i); math.IsNaN(x) || x <= 0 { + return nil, errf("BesselKn: the argument must be positive, got %g at element %d", x, i) + } + } + return a.realFunc("BesselKn", func(x float64) float64 { + k0, k1 := besselK0(x), besselK1(x) + if n == 0 { + return k0 + } + if n == 1 { + return k1 + } + km := k0 + k := k1 + for j := 1; j < n; j++ { + km, k = k, km+2*float64(j)/x*k + } + return k + }) +} + +// besselI0 evaluates I₀: the standard 3.75-piecewise polynomial fit +// below, and the exact integral representation I₀ = (1/π)∫₀^π +// e^(x·cos θ) dθ above, where the polynomial's accuracy degrades. +func besselI0(x float64) float64 { + ax := math.Abs(x) + if ax < 3.75 { + t := x / 3.75 + t2 := t * t + return 1 + t2*(3.5156229+t2*(3.0899424+t2*(1.2067492+t2*(0.2659732+t2*(0.0360768+t2*0.0045813))))) + } + return quadratureBesselI(x, false) +} + +// besselI1 evaluates I₁: the polynomial fit below 3.75, the integral +// representation I₁ = (1/π)∫₀^π e^(x·cos θ)·cos θ dθ above. The odd +// parity comes out of the quadrature by itself; the small branch +// applies the sign of x explicitly. +func besselI1(x float64) float64 { + ax := math.Abs(x) + var r float64 + if ax < 3.75 { + t := x / 3.75 + t2 := t * t + r = ax * (0.5 + t2*(0.87890594+t2*(0.51498869+t2*(0.15084934+t2*(0.02658733+t2*(0.00301532+t2*0.00032411)))))) + if x < 0 { + r = -r + } + return r + } + return quadratureBesselI(x, true) +} + +// quadratureBesselI evaluates I₀ (weight false) or I₁ (weight true) +// through the exact integral representations over θ ∈ [0, π] by +// composite Simpson with 2048 intervals, which resolves the peak at +// θ = 0 for any practical |x|. +func quadratureBesselI(x float64, weightCos bool) float64 { + const intervals = 2048 + h := math.Pi / float64(intervals) + kern := func(th float64) float64 { + v := math.Exp(x * math.Cos(th)) + if weightCos { + v *= math.Cos(th) + } + return v + } + sum := kern(0) + kern(math.Pi) + for i := 1; i < intervals; i++ { + w := 2.0 + if i%2 == 1 { + w = 4.0 + } + sum += w * kern(float64(i)*h) + } + return sum * h / (3 * math.Pi) +} + +// besselK0 evaluates K₀(x) = ∫₀^∞ e^(−x·cosh t) dt by composite +// Simpson quadrature over a truncated domain. The tail beyond the +// cutoff falls below 1e-100 for every x > 0, and the integrand is +// smooth, so the quadrature carries roughly ten significant digits. +// That is the deliberate trade: a provably correct entry-level K +// without hand-typed approximation coefficients. +func besselK0(x float64) float64 { + if x <= 0 { + return math.Inf(1) + } + return integrateCoshKernel(x, func(_ float64, coshU float64) float64 { + return math.Exp(-x * coshU) + }) +} + +// besselK1 evaluates K₁(x) = ∫₀^∞ e^(−x·cosh t)·cosh t dt, the +// derivative partner of K₀, by the same quadrature. +func besselK1(x float64) float64 { + if x <= 0 { + return math.Inf(1) + } + return integrateCoshKernel(x, func(_ float64, coshU float64) float64 { + return math.Exp(-x*coshU) * coshU + }) +} + +// integrateCoshKernel integrates ∫₀^∞ kernel(u; x) du where the +// kernel carries the factor e^(−x·cosh u). The cutoff follows from +// requiring the tail e^(−x·e^U/2) below 1e-100; for large x the +// result legitimately underflows to zero like the true value. +func integrateCoshKernel(x float64, kernel func(float64, float64) float64) float64 { + const intervals = 2001 + u := math.Log(200/x) + 2 + if u < 2 { + u = 2 + } + h := u / float64(intervals) + sum := kernel(0, 1) + kernel(u, math.Cosh(u)) + for i := 1; i < intervals; i++ { + tu := float64(i) * h + w := 2.0 + if i%2 == 1 { + w = 4.0 + } + sum += w * kernel(tu, math.Cosh(tu)) + } + return sum * h / 3 +} diff --git a/internal/core/besselmod_test.go b/internal/core/besselmod_test.go new file mode 100644 index 0000000..6e4dee9 --- /dev/null +++ b/internal/core/besselmod_test.go @@ -0,0 +1,152 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// TestModifiedBesselTabulated pins I and K against tabulated values +// at x = 1, the standard reference point. +func TestModifiedBesselTabulated(t *testing.T) { + x := mustFloats(t, []float64{1}) + i0, _ := BesselI0(x) + i1, _ := BesselI1(x) + k0, _ := BesselK0(x) + k1, _ := BesselK1(x) + want := []struct { + name string + got float64 + ref float64 + }{ + {"I0(1)", i0.FloatAt(0), 1.2660658777520084}, + {"I1(1)", i1.FloatAt(0), 0.5651591039924850}, + {"K0(1)", k0.FloatAt(0), 0.4210244382407083}, + {"K1(1)", k1.FloatAt(0), 0.6019072301972346}, + } + for _, w := range want { + if math.Abs(w.got-w.ref) > 5e-7*(1+math.Abs(w.ref)) { + t.Fatalf("%s = %.16g, want %.16g", w.name, w.got, w.ref) + } + } +} + +// TestModifiedBesselSymmetry checks the parity and small-argument +// behaviour: I₀ even with I₀(0) = 1, I₁ odd with I₁(0) = 0, and the +// K divergence to +Inf at the origin. +func TestModifiedBesselSymmetry(t *testing.T) { + x := mustFloats(t, []float64{-1, 1}) + i0, _ := BesselI0(x) + if math.Abs(i0.FloatAt(0)-i0.FloatAt(1)) > 1e-14 { + t.Fatalf("I₀ must be even: %v vs %v", i0.FloatAt(0), i0.FloatAt(1)) + } + i1, _ := BesselI1(x) + if math.Abs(i1.FloatAt(0)+i1.FloatAt(1)) > 1e-14 { + t.Fatalf("I₁ must be odd: %v vs %v", i1.FloatAt(0), i1.FloatAt(1)) + } + i1At0, err := BesselI1(mustFloats(t, []float64{0})) + if err != nil { + t.Fatalf("BesselI1(0): %v", err) + } + if i1At0.FloatAt(0) != 0 { + t.Fatalf("I₁(0) = %v, want 0", i1At0.FloatAt(0)) + } + k0, err := BesselK0(mustFloats(t, []float64{0})) + if err != nil { + t.Fatalf("BesselK0(0): %v", err) + } + if !math.IsInf(k0.FloatAt(0), 1) { + t.Fatalf("K₀(0) = %v, want +Inf", k0.FloatAt(0)) + } +} + +// TestModifiedBesselWronskian checks the Wronskian identity +// I₀·K₁ + I₁·K₀ = 1/x, which couples the two kinds at every point. +func TestModifiedBesselWronskian(t *testing.T) { + x := mustFloats(t, []float64{0.1, 0.5, 1, 5, 20}) + i0, _ := BesselI0(x) + i1, _ := BesselI1(x) + k0, _ := BesselK0(x) + k1, _ := BesselK1(x) + for i := range x.Len() { + xv := x.FloatAt(i) + w := i0.FloatAt(i)*k1.FloatAt(i) + i1.FloatAt(i)*k0.FloatAt(i) + if math.Abs(w-1/xv) > 1e-6*(1/xv) { + t.Fatalf("x=%v: wronskian = %.12g, want %.12g", xv, w, 1/xv) + } + } +} + +// TestBesselRecurrences checks the three-term recurrences linking +// consecutive integer orders, for both kinds. +func TestBesselRecurrences(t *testing.T) { + x := mustFloats(t, []float64{1, 4, 9}) + const n = 3 + in, err := BesselIn(n, x) + if err != nil { + t.Fatalf("BesselIn: %v", err) + } + inm1, err := BesselIn(n-1, x) + if err != nil { + t.Fatalf("BesselIn(n−1): %v", err) + } + inm2, err := BesselIn(n-2, x) + if err != nil { + t.Fatalf("BesselIn(n−2): %v", err) + } + kn, err := BesselKn(n, x) + if err != nil { + t.Fatalf("BesselKn: %v", err) + } + knm1, err := BesselKn(n-1, x) + if err != nil { + t.Fatalf("BesselKn(n−1): %v", err) + } + knm2, err := BesselKn(n-2, x) + if err != nil { + t.Fatalf("BesselKn(n−2): %v", err) + } + for i := range x.Len() { + xv := x.FloatAt(i) + // I_{n−1} − I_{n+1} = 2n/x·I_n with n − 1 in the call. + got := inm2.FloatAt(i) - in.FloatAt(i) + want := 2 * float64(n-1) / xv * inm1.FloatAt(i) + if math.Abs(got-want) > 1e-6*(1+math.Abs(want)) { + t.Fatalf("I recurrence at %v: %v, want %v", xv, got, want) + } + // K_{n+1} = K_{n−1} + 2n/x·K_n with n − 1 in the call. + got = knm2.FloatAt(i) + 2*float64(n-1)/xv*knm1.FloatAt(i) + if math.Abs(got-kn.FloatAt(i)) > 1e-6*(1+math.Abs(kn.FloatAt(i))) { + t.Fatalf("K recurrence at %v: %v, want %v", xv, got, kn.FloatAt(i)) + } + } +} + +// TestModifiedBesselRejects pins the input contracts: negative orders +// are refused for both integer-order entry points, and K diverges to +// +Inf at the origin on the diagonal of its domain. +func TestModifiedBesselRejects(t *testing.T) { + x := mustFloats(t, []float64{1}) + if _, err := BesselIn(-1, x); err == nil { + t.Fatal("expected an error for a negative I order") + } + if _, err := BesselKn(-1, x); err == nil { + t.Fatal("expected an error for a negative K order") + } + k0, err := BesselK0(mustFloats(t, []float64{0})) + if err != nil { + t.Fatalf("BesselK0(0): %v", err) + } + if !math.IsInf(k0.FloatAt(0), 1) { + t.Fatalf("K₀(0) = %v, want +Inf", k0.FloatAt(0)) + } + k1, err := BesselK1(mustFloats(t, []float64{0})) + if err != nil { + t.Fatalf("BesselK1(0): %v", err) + } + if !math.IsInf(k1.FloatAt(0), 1) { + t.Fatalf("K₁(0) = %v, want +Inf", k1.FloatAt(0)) + } +} diff --git a/internal/core/broadcast.go b/internal/core/broadcast.go new file mode 100644 index 0000000..e8cbe31 --- /dev/null +++ b/internal/core/broadcast.go @@ -0,0 +1,244 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/engine" + +// Explicit broadcasting. Silent broadcasting is banned; the +// escape hatch is a named operation: BroadcastTo expands one array to a +// target shape, BroadcastWith expands two operands to their common shape. +// Both copy: results never alias their receiver. + +// broadcastMinPerWorker is the element count every spawned broadcast +// worker must carry. Its inner loop is one run fill or one element write +// per step, cheaper per element than an arithmetic map, but a broadcast +// below a few thousand elements still measures several times faster on +// the calling goroutine than fanned out: a 256-element expansion paid +// more in thirty-two spawns than the whole fill cost. +const broadcastMinPerWorker = 1024 + +// BroadcastTo returns a new array expanded to the target shape: size-1 +// dimensions replicate, missing leading dimensions prepend, and +// anything else is a loud error naming both shapes. +func BroadcastTo(a *Array, shape ...int) (*Array, error) { + target, _, err := checkedDims(shape) + if err != nil { + return nil, err + } + if err := broadcastable(a.shape, shape); err != nil { + return nil, err + } + + sh := make([]int, len(shape)) + copy(sh, shape) + out := &Array{shape: sh, dt: a.dt} + out.alloc(target) + if target == 0 { + return out, nil + } + + // BroadcastTo used to walk the target one element at a time: every + // element recomputed its source index over the whole rank and took + // an odometer step. Profiling the MLP training step put that walk at + // the top of the profile, so the output is filled in runs instead. + // + // For each target dimension, srcStride is the stride the source + // advances by (0 where the source dimension is 1 or prepended). A + // trailing block of zero strides means the source element is + // constant across the whole block, so each run is a plain fill: the + // bias-shaped broadcast (1, C, 1, 1) to (N, C, H, W) collapses to one + // fill per channel rather than an odometer step per cell. + off := len(shape) - len(a.shape) + srcStride := make([]int, len(shape)) + srcDense := denseStrides(a.shape) + for d := range shape { + if d >= off && a.shape[d-off] != 1 { + srcStride[d] = srcDense[d-off] + } + } + tail := 0 + for d := len(shape) - 1; d >= 0 && srcStride[d] == 0; d-- { + tail++ + } + run := 1 + for d := len(shape) - tail; d < len(shape); d++ { + run *= shape[d] + } + outerShape := shape[:len(shape)-tail] + outer := target / run + + // The target offsets are disjoint, so the walk parallelises. + // A worker needs broadcastMinPerWorker output elements before its + // spawn pays for itself, and one outer step writes a whole run, so + // the item floor is the element floor over the run length. + parMin := max(1, broadcastMinPerWorker/run) + engine.ParallelMin(outer, parMin, func(s, e int) { + coord := make([]int, len(outerShape)) + src := 0 + rem := s + for d := len(outerShape) - 1; d >= 0; d-- { + if outerShape[d] == 0 { + coord[d] = 0 + continue + } + c := rem % outerShape[d] + rem /= outerShape[d] + coord[d] = c + src += c * srcStride[d] + } + if run == 1 { + // No constant run to fill; keep the plain element write so a + // pure prefix broadcast does not pay for the run machinery. + for o := s; o < e; o++ { + out.setFrom(o, a, src) + src += advanceStride(coord, outerShape, srcStride) + } + return + } + for o := s; o < e; o++ { + out.fillRun(o*run, run, a, src) + src += advanceStride(coord, outerShape, srcStride) + } + }) + return out, nil +} + +// advanceStride steps the odometer one position and returns the change +// in Σ coord[d]·stride[d], so a caller can carry the source index +// forward instead of rebuilding it from the coordinates each element. +func advanceStride(coord, shape, stride []int) int { + delta := 0 + for d := len(coord) - 1; d >= 0; d-- { + coord[d]++ + if coord[d] < shape[d] { + return delta + stride[d] + } + // Rolled from shape[d]−1 back to 0: that dim's whole + // contribution leaves the sum. + delta -= (shape[d] - 1) * stride[d] + coord[d] = 0 + } + return delta +} + +// fillRun writes n copies of element srcFlat of src starting at dst of a. +// The destination is freshly allocated and contiguous; the source may be +// a view. +func (a *Array) fillRun(dst, n int, src *Array, srcFlat int) { + p := src.physIndex(srcFlat) + switch a.dt { + case Int: + fillRepeat(a.ints[dst:dst+n], src.ints[p]) + case Bool: + fillRepeat(a.bools[dst:dst+n], src.bools[p]) + case Int8: + fillRepeat(a.i8s[dst:dst+n], src.i8s[p]) + case Uint8: + fillRepeat(a.u8s[dst:dst+n], src.u8s[p]) + case Int16: + fillRepeat(a.i16s[dst:dst+n], src.i16s[p]) + case Uint16: + fillRepeat(a.u16s[dst:dst+n], src.u16s[p]) + case Int32: + fillRepeat(a.i32s[dst:dst+n], src.i32s[p]) + case Uint32: + fillRepeat(a.u32s[dst:dst+n], src.u32s[p]) + case Float16: + fillRepeat(a.halves[dst:dst+n], src.halves[p]) + case Float32: + fillRepeat(a.floats32[dst:dst+n], src.floats32[p]) + case Float: + fillRepeat(a.floats[dst:dst+n], src.floats[p]) + default: + fillRepeat(a.complexes[dst:dst+n], src.complexes[p]) + } +} + +// fillRepeat writes v into every slot of dst. The first element seeds the +// run and each further round doubles the stretched prefix with one copy, +// so a long constant run is a logarithmic number of memmoves rather than +// one store per element. Source and destination never overlap in the +// direction of travel: the copied prefix ends exactly where the +// destination starts. +func fillRepeat[T any](dst []T, v T) { + if len(dst) == 0 { + return + } + dst[0] = v + for i := 1; i < len(dst); i *= 2 { + copy(dst[i:], dst[:i]) + } +} + +// BroadcastWith broadcasts both operands to their common shape, so +// `aw, bw, err := tensor.BroadcastWith(a, b)` can precede an ordinary +// element-wise operation. +func BroadcastWith(a, b *Array) (*Array, *Array, error) { + common, err := commonShape(a.shape, b.shape) + if err != nil { + return nil, nil, err + } + aw, err := BroadcastTo(a, common...) + if err != nil { + return nil, nil, err + } + bw, err := BroadcastTo(b, common...) + if err != nil { + return nil, nil, err + } + return aw, bw, nil +} + +// broadcastable reports whether from can broadcast to target, erroring +// with both shapes otherwise. +func broadcastable(from, target []int) error { + if len(from) > len(target) { + return errf("BroadcastTo: cannot broadcast %s to %s", shapeText(from), shapeText(target)) + } + off := len(target) - len(from) + for d := range from { + if from[d] != 1 && from[d] != target[off+d] { + return errf("BroadcastTo: cannot broadcast %s to %s", shapeText(from), shapeText(target)) + } + } + return nil +} + +// commonShape returns the shape two arrays broadcast to together. +func commonShape(a, b []int) ([]int, error) { + n := max(len(b), len(a)) + out := make([]int, n) + for d := range n { + av, bv := 1, 1 + if d >= n-len(a) { + av = a[d-(n-len(a))] + } + if d >= n-len(b) { + bv = b[d-(n-len(b))] + } + switch { + case av == bv: + out[d] = av + case av == 1: + out[d] = bv + case bv == 1: + out[d] = av + default: + return nil, errf("BroadcastWith: shapes %s and %s do not meet", + shapeText(a), shapeText(b)) + } + } + return out, nil +} + +// advanceOdometer increments a row-major coordinate over shape. +func advanceOdometer(coord, shape []int) { + for d := len(coord) - 1; d >= 0; d-- { + coord[d]++ + if coord[d] < shape[d] { + return + } + coord[d] = 0 + } +} diff --git a/internal/core/broadcast_test.go b/internal/core/broadcast_test.go new file mode 100644 index 0000000..b0faf5e --- /dev/null +++ b/internal/core/broadcast_test.go @@ -0,0 +1,234 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "strings" + "testing" +) + +func TestBroadcastTo(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3}, 3, 1) + + b, err := BroadcastTo(a, 3, 2) + if err != nil { + t.Fatalf("BroadcastTo: %v", err) + } + want := mustFromInts(t, []int64{1, 1, 2, 2, 3, 3}, 3, 2) + if !Equal(want, b) { + t.Fatalf("BroadcastTo: %s", b) + } + + // Prepending a leading dimension. + c, err := BroadcastTo(a, 2, 3, 1) + if err != nil { + t.Fatalf("BroadcastTo prepend: %v", err) + } + if c.NDim() != 3 || c.Shape()[0] != 2 { + t.Fatalf("BroadcastTo prepend shape: %v", c.Shape()) + } + if v, _ := IntAt(c, 1, 2, 0); v != 3 { + t.Fatalf("BroadcastTo prepend value: %d", v) + } + + // An identical shape is a copy. + same, err := BroadcastTo(a, 3, 1) + if err != nil || !Equal(a, same) { + t.Fatalf("BroadcastTo same: %s %v", same, err) + } + + if _, err := BroadcastTo(a, 4, 1); err == nil || !strings.Contains(err.Error(), "cannot broadcast") { + t.Fatalf("BroadcastTo incompatible: %v", err) + } + tall := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + if _, err := BroadcastTo(tall, 2, 3); err == nil || !strings.Contains(err.Error(), "cannot broadcast") { + t.Fatalf("BroadcastTo rank-2 mismatch: %v", err) + } + if _, err := BroadcastTo(a, -1); err == nil { + t.Fatalf("BroadcastTo negative must error") + } +} + +func TestBroadcastWith(t *testing.T) { + a := mustFromInts(t, []int64{10, 20, 30}, 3, 1) + b := mustFromInts(t, []int64{1, 2}, 1, 2) + + aw, bw, err := BroadcastWith(a, b) + if err != nil { + t.Fatalf("BroadcastWith: %v", err) + } + if shape := aw.Shape(); shape[0] != 3 || shape[1] != 2 { + t.Fatalf("BroadcastWith shape: %v", shape) + } + // (3,2): a column repeats across, b row repeats down. + wantA := mustFromInts(t, []int64{10, 10, 20, 20, 30, 30}, 3, 2) + wantB := mustFromInts(t, []int64{1, 2, 1, 2, 1, 2}, 3, 2) + if !Equal(wantA, aw) || !Equal(wantB, bw) { + t.Fatalf("BroadcastWith: %s / %s", aw, bw) + } + + // The broadcast operands now satisfy the strict element-wise ops. + sum, err := Add(aw, bw) + if err != nil { + t.Fatalf("Add after broadcast: %v", err) + } + wantSum := mustFromInts(t, []int64{11, 12, 21, 22, 31, 32}, 3, 2) + if !Equal(wantSum, sum) { + t.Fatalf("Add after broadcast: %s", sum) + } + + x := mustFromInts(t, []int64{1, 2}, 2) + y := mustFromInts(t, []int64{1, 2, 3}, 3) + if _, _, err := BroadcastWith(x, y); err == nil || !strings.Contains(err.Error(), "do not meet") { + t.Fatalf("BroadcastWith incompatible: %v", err) + } +} + +// broadcastReference is a deliberately naive row-major walk: for every +// target position it derives the source coordinate, so it shares no code +// with the run-fill implementation. Any disagreement is a bug in one of +// them. +func broadcastReference(a *Array, shape []int) *Array { + total := 1 + for _, d := range shape { + total *= d + } + out := &Array{shape: append([]int(nil), shape...), dt: a.Dtype()} + out.alloc(total) + coord := make([]int, len(shape)) + for i := range total { + src := 0 + for d := range a.Shape() { + td := len(shape) - len(a.Shape()) + d + c := coord[td] + if a.Shape()[d] == 1 { + c = 0 + } + src = src*a.Shape()[d] + c + } + out.setFrom(i, a, src) + advanceOdometer(coord, shape) + // advanceOdometer is the helper under test elsewhere; the + // independent rebuild above is what makes this a reference. + } + return out +} + +// TestBroadcastToMatchesReference drives the run-fill rewrite against +// the naive walk over the shape families it specialises: a constant +// trailing run (bias layouts), a prefixed rank, a size-1 dim mid-shape, +// and a full no-op broadcast. +func TestBroadcastToMatchesReference(t *testing.T) { + for _, tc := range []struct { + src []int + dst []int + name string + }{ + {[]int{3}, []int{3}, "identity"}, + {[]int{1, 3, 1, 1}, []int{2, 3, 4, 5}, "bias NCHW"}, + {[]int{1, 4}, []int{6, 4}, "bias NC"}, + {[]int{2, 3}, []int{5, 2, 3}, "prefix prepend"}, + {[]int{2, 1, 3}, []int{2, 4, 3}, "middle size-1"}, + {[]int{1, 1, 1}, []int{2, 3, 4}, "scalar-ish"}, + {[]int{5, 1}, []int{5, 7}, "trailing replicate"}, + {[]int{1}, []int{4, 1}, "single prepend"}, + {[]int{2, 3, 1}, []int{1, 2, 3, 6}, "mixed"}, + } { + n := 1 + for _, d := range tc.src { + n *= d + } + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(i)*1.5 - 3 + } + src, err := FromFloats(vals, tc.src...) + if err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + got, err := BroadcastTo(src, tc.dst...) + if err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + want := broadcastReference(src, tc.dst) + if got.Len() != want.Len() { + t.Fatalf("%s: length %d, want %d", tc.name, got.Len(), want.Len()) + } + for i := range want.Len() { + if got.FloatAt(i) != want.FloatAt(i) { + t.Fatalf("%s: element %d = %v, want %v", tc.name, i, got.FloatAt(i), want.FloatAt(i)) + } + } + } +} + +// TestBroadcastToSplitMatchesSerial pins the parallel outer walk's chunk +// cursor. The fill splits the outer positions across workers and each +// worker rebuilds the source coordinate of its own first position, so +// the result must not depend on where the chunk boundaries fall. Every +// outer position below reads a different source element, and the split +// is asserted before it is compared so a shrunken workload cannot +// quietly fall back to the single-chunk fill. +func TestBroadcastToSplitMatchesSerial(t *testing.T) { + const workers = 4 + cases := []struct { + src, target []int + run int // trailing target elements one outer position fills + }{ + {[]int{4096, 1}, []int{4096, 8}, 8}, + {[]int{512, 1, 1}, []int{512, 4, 4}, 16}, + } + prev := NumWorkers() + defer SetNumCPU(prev) + for _, tc := range cases { + n := 1 + for _, d := range tc.src { + n *= d + } + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(i) + 0.5 + } + src := mustFromFloats(t, vals, tc.src...) + + outer := 1 + for _, d := range tc.target { + outer *= d + } + outer /= tc.run + parMin := max(1, broadcastMinPerWorker/tc.run) + if chunk := (outer + workers - 1) / workers; chunk < parMin { + t.Fatalf("src %v to %v: %d outer positions no longer split at %d workers: chunk %d below the %d-element floor", + tc.src, tc.target, outer, workers, chunk, parMin) + } + + SetNumCPU(1) + want, err := BroadcastTo(src, tc.target...) + if err != nil { + t.Fatalf("BroadcastTo %v: %v", tc.target, err) + } + SetNumCPU(workers) + got, err := BroadcastTo(src, tc.target...) + if err != nil { + t.Fatalf("BroadcastTo %v split: %v", tc.target, err) + } + wf, gf := want.RawFloats(), got.RawFloats() + if len(gf) != len(wf) { + t.Fatalf("BroadcastTo %v: split result has %d elements, the serial fill %d", + tc.target, len(gf), len(wf)) + } + for i := range wf { + if gf[i] != wf[i] { + t.Fatalf("BroadcastTo %v: split element %d = %v, the serial fill %v", + tc.target, i, gf[i], wf[i]) + } + } + // The last outer position is the one a moved chunk boundary + // reads from the wrong source slot: it must be its own row. + if last := float64(outer-1) + 0.5; gf[len(gf)-1] != last { + t.Fatalf("BroadcastTo %v: last element = %v, want %v", + tc.target, gf[len(gf)-1], last) + } + } +} diff --git a/internal/core/canonical_partition_pins_test.go b/internal/core/canonical_partition_pins_test.go new file mode 100644 index 0000000..87e21e9 --- /dev/null +++ b/internal/core/canonical_partition_pins_test.go @@ -0,0 +1,132 @@ +package core + +import "testing" + +// The canonical partition's block boundaries are c·n/foldParts(n); the +// pins here hold the multi-block and remainder shapes against the +// global one-dimensional folds bit for bit, the consistency the shape +// consistency family promises for every length. + +func TestSumAxisRemainderAgainstGlobal(t *testing.T) { + for _, n := range []int{8191, 131073} { + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(i%17) - 8 + float64(i)*1e-9 + } + a, err := FromFloats(vals, 1, n) + if err != nil { + t.Fatal(err) + } + got, err := SumAxis(a, 1) + if err != nil { + t.Fatal(err) + } + want := Sum(a) + if got.FloatAt(0) != want.Float() { + t.Fatalf("SumAxis over length %d is not the global Sum's bits: %v vs %v", n, got.FloatAt(0), want.Float()) + } + mean, err := MeanAxis(a, 1) + if err != nil { + t.Fatal(err) + } + wantMean := want.Float() / float64(n) + if mean.FloatAt(0) != wantMean { + t.Fatalf("MeanAxis over length %d is not Mean's bits: %v vs %v", n, mean.FloatAt(0), wantMean) + } + } +} + +func TestNormLineMultiBlock(t *testing.T) { + n := 131073 + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(i%11) - 5 + float64(i%7)*1e-7 + } + a, err := FromFloats(vals, 1, n) + if err != nil { + t.Fatal(err) + } + flat, err := FromFloats(vals, n) + if err != nil { + t.Fatal(err) + } + for _, p := range []float64{1, 2, 3.5} { + got, err := Norm(a, p, 1, false) + if err != nil { + t.Fatal(err) + } + want, err := Norm(flat, p, 0, false) + if err != nil { + t.Fatal(err) + } + if got.FloatAt(0) != want.FloatAt(0) { + t.Fatalf("Norm p=%v over a multi-block line is not the 1-D Norm's bits: %v vs %v", p, got.FloatAt(0), want.FloatAt(0)) + } + } +} + +func TestProdLineMultiBlock(t *testing.T) { + n := 131073 + vals := make([]float64, n) + for i := range vals { + vals[i] = 1 + float64(i%5)*0.125 + } + a, err := FromFloats(vals, 1, n) + if err != nil { + t.Fatal(err) + } + flat, err := FromFloats(vals, n) + if err != nil { + t.Fatal(err) + } + got, err := Prod(a, 1, false) + if err != nil { + t.Fatal(err) + } + want, err := Prod(flat, 0, false) + if err != nil { + t.Fatal(err) + } + if got.FloatAt(0) != want.FloatAt(0) { + t.Fatalf("Prod over a multi-block line is not the 1-D Prod's bits: %v vs %v", got.FloatAt(0), want.FloatAt(0)) + } + + // The half payload walks the same partition with its own narrowing; + // an odd length puts a remainder into the last block. + halves := make([]float64, n) + for i := range halves { + halves[i] = 1 + float64(i%3)*0.25 + } + h, err := FromFloats(halves, 1, n) + if err != nil { + t.Fatal(err) + } + h16, err := Astype(h, Float16) + if err != nil { + t.Fatal(err) + } + hflat, err := Astype(mustFloats1D(t, halves), Float16) + if err != nil { + t.Fatal(err) + } + hgot, err := Prod(h16, 1, false) + if err != nil { + t.Fatal(err) + } + hwant, err := Prod(hflat, 0, false) + if err != nil { + t.Fatal(err) + } + if hgot.FloatAt(0) != hwant.FloatAt(0) { + t.Fatalf("half Prod over a multi-block line is not the 1-D half Prod's bits: %v vs %v", hgot.FloatAt(0), hwant.FloatAt(0)) + } +} + +func mustFloats1D(t *testing.T, vals []float64) *Array { + t.Helper() + a, err := FromFloats(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + return a +} diff --git a/internal/core/centralsum_kernels_test.go b/internal/core/centralsum_kernels_test.go new file mode 100644 index 0000000..15b04e3 --- /dev/null +++ b/internal/core/centralsum_kernels_test.go @@ -0,0 +1,364 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + mrand "math/rand/v2" + "testing" +) + +// The central-sum evidence file: the accessor-chain Covariance, +// Correlation and Integrate walks against their canonical-block + +// treeSum candidates, both measured against exact big.Float referents +// (512 bits) and timed in one binary. + +// covarianceLegacy is the walk the entry point keeps: two accessor +// reads per element folded into one chain. +func covarianceLegacy(a, b []float64, ma, mb float64) float64 { + var cov float64 + for i := range a { + cov += (a[i] - ma) * (b[i] - mb) + } + return cov +} + +// covarianceBlocks is the candidate: the same per-element arithmetic +// cut into the canonical blocks, one chain partial per block, the +// partials combined through the balanced tree. +func covarianceBlocks(a, b []float64, ma, mb float64) float64 { + n := len(a) + parts := foldParts(n) + if parts == 1 { + return covarianceLegacy(a, b, ma, mb) + } + partials := make([]float64, parts) + for c := range parts { + lo, hi := c*n/parts, (c+1)*n/parts + var acc float64 + for i := lo; i < hi; i++ { + acc += (a[i] - ma) * (b[i] - mb) + } + partials[c] = acc + } + return treeSum(partials) +} + +// correlationLegacy is the accessor-chain correlation walk. +func correlationLegacy(a, b []float64, ma, mb float64) (num, da2, db2 float64) { + for i := range a { + da := a[i] - ma + dbv := b[i] - mb + num += da * dbv + da2 += da * da + db2 += dbv * dbv + } + return num, da2, db2 +} + +// correlationBlocks is the candidate: canonical blocks and a tree over +// the partials of each of the three sums. +func correlationBlocks(a, b []float64, ma, mb float64) (num, da2, db2 float64) { + n := len(a) + parts := foldParts(n) + if parts == 1 { + return correlationLegacy(a, b, ma, mb) + } + pn, px, py := make([]float64, parts), make([]float64, parts), make([]float64, parts) + for c := range parts { + lo, hi := c*n/parts, (c+1)*n/parts + var sn, sx, sy float64 + for i := lo; i < hi; i++ { + da := a[i] - ma + dbv := b[i] - mb + sn += da * dbv + sx += da * da + sy += dbv * dbv + } + pn[c], px[c], py[c] = sn, sx, sy + } + return treeSum(pn), treeSum(px), treeSum(py) +} + +// integrateLegacy is the plain trapezoid chain. +func integrateLegacy(y []float64) float64 { + var total float64 + for i := 1; i < len(y); i++ { + total += (y[i-1] + y[i]) / 2 + } + return total +} + +// integrateBlocks is the candidate: the trapezoid areas cut into the +// canonical blocks over the area indices (n−1 of them), the block +// chains combined through the tree. +func integrateBlocks(y []float64) float64 { + m := len(y) - 1 + parts := foldParts(m) + if parts == 1 { + return integrateLegacy(y) + } + partials := make([]float64, parts) + for c := range parts { + lo, hi := c*m/parts, (c+1)*m/parts + var acc float64 + for i := lo; i < hi; i++ { + acc += (y[i] + y[i+1]) / 2 + } + partials[c] = acc + } + return treeSum(partials) +} + +// spreadPair builds a deterministic correlated pair whose magnitudes +// span thirty orders, the shape that cancels a central sum fastest. +func spreadPair(n int, seed uint64) (x, y []float64) { + rng := mrand.New(mrand.NewPCG(seed, seed)) + x = make([]float64, n) + y = make([]float64, n) + for i := range x { + mag := math.Pow(10, -15+30*rng.Float64()) + if rng.Float64() < 0.5 { + mag = -mag + } + x[i] = mag + y[i] = 3*mag + math.Pow(10, -15+30*rng.Float64()) + } + return x, y +} + +// exactMeans returns the exact arithmetic means at the given precision. +func exactMeans(x, y []float64, prec uint) (*big.Float, *big.Float) { + sx := new(big.Float).SetPrec(prec) + sy := new(big.Float).SetPrec(prec) + one := new(big.Float).SetPrec(prec) + for i := range x { + sx.Add(sx, new(big.Float).SetPrec(prec).SetFloat64(x[i])) + sy.Add(sy, new(big.Float).SetPrec(prec).SetFloat64(y[i])) + } + one.SetInt(new(big.Int).SetInt64(int64(len(x)))) + return new(big.Float).SetPrec(prec).Quo(sx, one), new(big.Float).SetPrec(prec).Quo(sy, one) +} + +// TestCentralSumsAccuracy measures the legacy and candidate walks of +// the three central-sum kernels against exact referents. +func TestCentralSumsAccuracy(t *testing.T) { + const prec = 512 + n := 1 << 20 + x, y := spreadPair(n, 0xBEEF) + exma, exmb := exactMeans(x, y, prec) + maf, _ := exma.Float64() + mbf, _ := exmb.Float64() + + // Covariance referent: Σ(x−x̄)(y−ȳ)/(n−1) at full precision. + ref := new(big.Float).SetPrec(prec) + for i := range x { + dx := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(x[i]), exma) + dy := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(y[i]), exmb) + ref.Add(ref, new(big.Float).SetPrec(prec).Mul(dx, dy)) + } + ref.Quo(ref, new(big.Float).SetPrec(prec).SetInt(new(big.Int).SetInt64(int64(n-1)))) + refF, _ := ref.Float64() + + oldCov := covarianceLegacy(x, y, maf, mbf) / float64(n-1) + newCov := covarianceBlocks(x, y, maf, mbf) / float64(n-1) + scale := math.Max(math.Abs(refF), 1) + t.Logf("covariance n=2^20: exact=%.17g legacy=%.17g (%.3e) blocks=%.17g (%.3e)", + refF, oldCov, math.Abs(oldCov-refF)/scale, newCov, math.Abs(newCov-refF)/scale) + + // Correlation referent. + var rnum, rda2, rdb2 = new(big.Float).SetPrec(prec), new(big.Float).SetPrec(prec), new(big.Float).SetPrec(prec) + for i := range x { + dx := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(x[i]), exma) + dy := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(y[i]), exmb) + rnum.Add(rnum, new(big.Float).SetPrec(prec).Mul(dx, dy)) + rda2.Add(rda2, new(big.Float).SetPrec(prec).Mul(dx, dx)) + rdb2.Add(rdb2, new(big.Float).SetPrec(prec).Mul(dy, dy)) + } + root := new(big.Float).SetPrec(prec).Mul(rda2, rdb2) + root.Sqrt(root) + refR, _ := new(big.Float).SetPrec(prec).Quo(rnum, root).Float64() + + on, oda2, odb2 := correlationLegacy(x, y, maf, mbf) + nn, nda2, ndb2 := correlationBlocks(x, y, maf, mbf) + oldR := on / math.Sqrt(oda2*odb2) + newR := nn / math.Sqrt(nda2*ndb2) + t.Logf("correlation n=2^20: exact=%.17g legacy=%.17g (%.3e) blocks=%.17g (%.3e)", + refR, oldR, math.Abs(oldR-refR), newR, math.Abs(newR-refR)) + + // Integrate referent: the exact trapezoid sum of the same sample. + rngI := mrand.New(mrand.NewPCG(0xDADA, 99)) + yv := make([]float64, n) + for i := range yv { + yv[i] = math.Pow(10, -15+30*rngI.Float64()) + } + tri := new(big.Float).SetPrec(prec) + for i := 1; i < n; i++ { + a := new(big.Float).SetPrec(prec).SetFloat64(yv[i-1]) + b := new(big.Float).SetPrec(prec).SetFloat64(yv[i]) + tri.Add(tri, new(big.Float).SetPrec(prec).Quo(new(big.Float).SetPrec(prec).Add(a, b), new(big.Float).SetPrec(prec).SetFloat64(2))) + } + refI, _ := tri.Float64() + oldI := integrateLegacy(yv) + newI := integrateBlocks(yv) + scaleI := math.Max(math.Abs(refI), 1) + t.Logf("integrate n=2^20: exact=%.17g legacy=%.17g (%.3e) blocks=%.17g (%.3e)", + refI, oldI, math.Abs(oldI-refI)/scaleI, newI, math.Abs(newI-refI)/scaleI) +} + +// TestCentralSumsProductionPins pins the shipped entry points on the +// spread sample against the exact referent, at tolerances the block +// walk clears by two orders and the plain chain fails: the covariance +// chain erred 1.111e-11 relative here, the block walk 1.642e-14; the +// correlation chain 1.858e-11, the block walk 2.032e-14; the integrate +// chain 5.341e-13, the block walk 3.046e-14. +func TestCentralSumsProductionPins(t *testing.T) { + const prec = 512 + n := 1 << 20 + x, y := spreadPair(n, 0xBEEF) + xa, _ := FromFloats(x, n) + ya, _ := FromFloats(y, n) + exma, exmb := exactMeans(x, y, prec) + + ref := new(big.Float).SetPrec(prec) + for i := range x { + dx := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(x[i]), exma) + dy := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(y[i]), exmb) + ref.Add(ref, new(big.Float).SetPrec(prec).Mul(dx, dy)) + } + ref.Quo(ref, new(big.Float).SetPrec(prec).SetInt(new(big.Int).SetInt64(int64(n-1)))) + refF, _ := ref.Float64() + cov, err := Covariance(xa, ya) + if err != nil { + t.Fatal(err) + } + if rel := math.Abs(cov-refF) / math.Abs(refF); rel > 1e-12 { + t.Errorf("covariance relative error %g exceeds 1e-12 (got %v, want %v)", rel, cov, refF) + } + + rnum := new(big.Float).SetPrec(prec) + rda2 := new(big.Float).SetPrec(prec) + rdb2 := new(big.Float).SetPrec(prec) + for i := range x { + dx := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(x[i]), exma) + dy := new(big.Float).SetPrec(prec).Sub(new(big.Float).SetPrec(prec).SetFloat64(y[i]), exmb) + rnum.Add(rnum, new(big.Float).SetPrec(prec).Mul(dx, dy)) + rda2.Add(rda2, new(big.Float).SetPrec(prec).Mul(dx, dx)) + rdb2.Add(rdb2, new(big.Float).SetPrec(prec).Mul(dy, dy)) + } + root := new(big.Float).SetPrec(prec).Mul(rda2, rdb2) + root.Sqrt(root) + refR, _ := new(big.Float).SetPrec(prec).Quo(rnum, root).Float64() + corr, err := Correlation(xa, ya) + if err != nil { + t.Fatal(err) + } + if math.Abs(corr-refR) > 1e-12 { + t.Errorf("correlation error %g exceeds 1e-12 (got %v, want %v)", math.Abs(corr-refR), corr, refR) + } + + yv := make([]float64, n) + rngI := mrand.New(mrand.NewPCG(0xDADA, 99)) + for i := range yv { + yv[i] = math.Pow(10, -15+30*rngI.Float64()) + } + tri := new(big.Float).SetPrec(prec) + for i := 1; i < n; i++ { + a := new(big.Float).SetPrec(prec).SetFloat64(yv[i-1]) + b := new(big.Float).SetPrec(prec).SetFloat64(yv[i]) + tri.Add(tri, new(big.Float).SetPrec(prec).Quo(new(big.Float).SetPrec(prec).Add(a, b), new(big.Float).SetPrec(prec).SetFloat64(2))) + } + refI, _ := tri.Float64() + ya2, _ := FromFloats(yv, n) + gotI, err := Integrate(ya2, 1) + if err != nil { + t.Fatal(err) + } + if rel := math.Abs(gotI-refI) / math.Abs(refI); rel > 1e-13 { + t.Errorf("integrate relative error %g exceeds 1e-13 (got %v, want %v)", rel, gotI, refI) + } +} + +// TestCentralSumsShortInputsPinEqualBits pins the contract that the +// canonical partition leaves short inputs' bits alone: at and below one +// fold block the block walk is the chain walk, sample for sample. +func TestCentralSumsShortInputsPinEqualBits(t *testing.T) { + rng := mrand.New(mrand.NewPCG(5, 5)) + n := 3000 + x := make([]float64, n) + y := make([]float64, n) + for i := range x { + x[i] = rng.NormFloat64() + y[i] = rng.NormFloat64() + } + ma, mb := 0.37, -1.2 + if covarianceLegacy(x, y, ma, mb) != covarianceBlocks(x, y, ma, mb) { + t.Error("covariance: short-input bits moved") + } + on, oda, odb := correlationLegacy(x, y, ma, mb) + nn, nda, ndb := correlationBlocks(x, y, ma, mb) + if on != nn || oda != nda || odb != ndb { + t.Error("correlation: short-input bits moved") + } + yy := make([]float64, n+1) + for i := range yy { + yy[i] = rng.NormFloat64() + } + if integrateLegacy(yy) != integrateBlocks(yy) { + t.Error("integrate: short-input bits moved") + } +} + +// BenchmarkCentralSumsTimes the two walks of each kernel over one +// pair of 2^20 samples. +func BenchmarkCentralSums(b *testing.B) { + n := 1 << 20 + x, y := spreadPair(n, 0xBEEF) + ma, mb := 0.37, -1.2 + b.Run("covariance/legacy", func(b *testing.B) { + for b.Loop() { + covarianceLegacy(x, y, ma, mb) + } + b.SetBytes(int64(n) * 8) + }) + b.Run("covariance/blocks", func(b *testing.B) { + for b.Loop() { + covarianceBlocks(x, y, ma, mb) + } + b.SetBytes(int64(n) * 8) + }) + b.Run("correlation/legacy", func(b *testing.B) { + for b.Loop() { + correlationLegacy(x, y, ma, mb) + } + b.SetBytes(int64(n) * 8) + }) + b.Run("correlation/blocks", func(b *testing.B) { + for b.Loop() { + correlationBlocks(x, y, ma, mb) + } + b.SetBytes(int64(n) * 8) + }) + b.Run("integrate/legacy", func(b *testing.B) { + yy := make([]float64, n) + for i := range yy { + yy[i] = float64(i % 977) + } + for b.Loop() { + integrateLegacy(yy) + } + b.SetBytes(int64(n) * 8) + }) + b.Run("integrate/blocks", func(b *testing.B) { + yy := make([]float64, n) + for i := range yy { + yy[i] = float64(i % 977) + } + for b.Loop() { + integrateBlocks(yy) + } + b.SetBytes(int64(n) * 8) + }) +} diff --git a/internal/core/chunk_precision_pins_test.go b/internal/core/chunk_precision_pins_test.go new file mode 100644 index 0000000..095c65d --- /dev/null +++ b/internal/core/chunk_precision_pins_test.go @@ -0,0 +1,342 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Regression pins for precision and routing: the worker chunk that +// aborted on a divergent element, the argmax int64 comparisons, einsum +// mixed-dtype promotion and label validation, and the asymptotic +// special functions (Fresnel, elliptic F). + +// TestSphericalBesselYChunkAbort pins the divergence handling: a zero +// in the input used to abort the whole worker chunk, leaving every +// later element of that chunk at zero. +func TestSphericalBesselYChunkAbort(t *testing.T) { + // Two workers make the chunk multi-element on any machine: with one + // element per chunk the abandoned-chunk defect is invisible, and a + // 64-core host would miss it. + defer engine.SetNumWorkers(engine.SetNumWorkers(2)) + x := make([]float64, 64) + for i := range x { + x[i] = 2 + } + x[0] = 0 + a, err := FromFloats(x, 64) + if err != nil { + t.Fatal(err) + } + out, err := SphericalBesselY(0, a) + if err != nil { + t.Fatalf("SphericalBesselY: %v", err) + } + if got := out.FloatAt(0); !math.IsInf(got, -1) { + t.Fatalf("y_0(0) = %v, want -Inf", got) + } + want := -math.Cos(2) / 2 + for i := 1; i < 64; i++ { + if got := out.FloatAt(i); math.Abs(got-want) > 1e-15 { + t.Fatalf("y_0(2) at index %d = %v, want %v: the zero aborted the chunk", i, got, want) + } + } +} + +// TestArgMaxInt64Precision pins the ordering of int arrays above the +// float64 integer range: widening to float64 rounds neighbours equal. +func TestArgMaxInt64Precision(t *testing.T) { + const big = int64(1) << 53 + a, err := FromInts([]int64{big, big + 1}, 2) + if err != nil { + t.Fatal(err) + } + got, err := ArgMax(a) + if err != nil { + t.Fatalf("ArgMax: %v", err) + } + if got != 1 { + t.Fatalf("ArgMax = %d, want 1", got) + } + gotMin, err := ArgMin(a) + if err != nil { + t.Fatalf("ArgMin: %v", err) + } + if gotMin != 0 { + t.Fatalf("ArgMin = %d, want 0", gotMin) + } +} + +// TestArgMaxAxisInt64Precision covers the axis path, which widened the +// elements the same way. +func TestArgMaxAxisInt64Precision(t *testing.T) { + const big = int64(1) << 53 + a, err := FromInts([]int64{big, big + 1, big + 1, big}, 2, 2) + if err != nil { + t.Fatal(err) + } + out, err := ArgMaxAxis(a, 1) + if err != nil { + t.Fatalf("ArgMaxAxis: %v", err) + } + if got := out.RawInts(); got[0] != 1 || got[1] != 0 { + t.Fatalf("ArgMaxAxis = %v, want [1 0]", got) + } +} + +// TestEinsumMixedDtypes pins the general engine on a real operand +// under a complex result: it used to enter the product as zero. +func TestEinsumMixedDtypes(t *testing.T) { + cv := make([]complex128, 2*3*4) + for i := range cv { + cv[i] = 1 + 1i + } + a, err := FromComplexes(cv, 2, 3, 4) + if err != nil { + t.Fatal(err) + } + b, err := FromFloats([]float64{2, 2, 2, 2}, 4) + if err != nil { + t.Fatal(err) + } + out, err := Einsum("aij,j->ai", a, b) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + want := complex(8, 8) // four terms of 2·(1+1i) + for i, got := range out.RawComplexes() { + if got != want { + t.Fatalf("Einsum[%d] = %v, want %v", i, got, want) + } + } + // The reverse order, and an int operand, must widen the same way. + iv, err := FromInts([]int64{3, 3, 3, 3}, 4) + if err != nil { + t.Fatal(err) + } + out, err = Einsum("aij,j->ai", a, iv) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + if got := out.RawComplexes()[0]; got != complex(12, 12) { + t.Fatalf("Einsum[0] = %v, want 12+12i", got) + } +} + +// TestEinsumLabelMismatch pins the labelled-axis rule: a written label +// must agree in length, size 1 included. +func TestEinsumLabelMismatch(t *testing.T) { + a, err := FromFloats(make([]float64, 2*3*4), 2, 3, 4) + if err != nil { + t.Fatal(err) + } + b, err := FromFloats(make([]float64, 5), 1, 5) + if err != nil { + t.Fatal(err) + } + if _, err := Einsum("aij,jk->aik", a, b); err == nil { + t.Fatal("expected an error for label j being 4 in one operand and 1 in another") + } else if !strings.Contains(err.Error(), "label") { + t.Fatalf("error = %v, want a label-mismatch refusal", err) + } + // The ellipsis keeps its broadcasting, written where the engine + // reads it: leading the subscripts. Both operands put a single + // label on the last axis and let the ellipsis carry the first. + s, err := FromFloats(make([]float64, 1*4), 1, 4) + if err != nil { + t.Fatal(err) + } + m, err := FromFloats(make([]float64, 3*4), 3, 4) + if err != nil { + t.Fatal(err) + } + out, err := Einsum("...i,...i->...i", s, m) + if err != nil { + t.Fatalf("Einsum with broadcasting ellipsis: %v", err) + } + if got := out.Shape(); got[0] != 3 || got[1] != 4 { + t.Fatalf("Einsum shape = %v, want [3 4]", got) + } + // An ellipsis written after a label is refused: the engine maps the + // ellipsis onto the leading axes, so accepting it would compute a + // different contraction than the one written. + if _, err := Einsum("i...,i...->i...", s, m); err == nil { + t.Fatal("expected an error for an ellipsis that does not lead") + } else if !strings.Contains(err.Error(), "must lead") { + t.Fatalf("error = %v, want an ellipsis-placement refusal", err) + } +} + +// TestFresnelAsymptotic pins the asymptotic branch: it used to run +// the power series in extended precision with a working size growing +// with x² and a fixed term cap, so arguments beyond ~505 returned +// noise and larger ones cost gigabytes per element. References are +// mpmath values at 40 digits (mp.dps = 40). +func TestFresnelAsymptotic(t *testing.T) { + // The tolerance grows with x: the reconstruction turns on the + // phase πx²/2, whose float64 representation is only good to + // πx²·eps/2, and that error enters the integrals scaled by + // f + g ≈ 2/(πx). + cases := []struct { + x float64 + wantC, wantS float64 + tol float64 + }{ + {11, 0.52893666175579644963, 0.49992388379375106274, 1e-14}, + {20, 0.49998733497234438819, 0.48408453592595389271, 1e-14}, + {40, 0.49999841685744546488, 0.49204225379027309718, 1e-14}, + {100, 0.49999989867881789756, 0.49681690114783755327, 1e-14}, + {1000, 0.49999999989867881636, 0.49968169011381630608, 1e-12}, + {100000, 0.49999999999999989868, 0.49999681690113816209, 5e-12}, + } + for _, tc := range cases { + if gotC := fresnelC(tc.x); math.Abs(gotC-tc.wantC) > tc.tol { + t.Errorf("C(%v) = %.17g, mpmath says %.17g", tc.x, gotC, tc.wantC) + } + if gotS := fresnelS(tc.x); math.Abs(gotS-tc.wantS) > tc.tol { + t.Errorf("S(%v) = %.17g, mpmath says %.17g", tc.x, gotS, tc.wantS) + } + } + // The asymptotic and the extended-precision series must agree + // across the crossover, which is the in-repo independent check. + for _, x := range []float64{10, 12, 15, 20} { + for _, sine := range []bool{false, true} { + got := fresnelAsym(x, sine) + want := fresnelBig(x, sine) + if math.Abs(got-want) > 1e-13*math.Max(1, math.Abs(want)) { + t.Errorf("x=%v sine=%v: asymptotic %.17g, series %.17g", x, sine, got, want) + } + } + } + // The derivative identity, C' = cos(πx²/2) and S' = sin(πx²/2). + const h = 1e-5 + for _, x := range []float64{11, 20} { + dC := (fresnelC(x+h) - fresnelC(x-h)) / (2 * h) + dS := (fresnelS(x+h) - fresnelS(x-h)) / (2 * h) + if math.Abs(dC-math.Cos(math.Pi*x*x/2)) > 1e-6 { + t.Errorf("C'(%v) = %v, want %v", x, dC, math.Cos(math.Pi*x*x/2)) + } + if math.Abs(dS-math.Sin(math.Pi*x*x/2)) > 1e-6 { + t.Errorf("S'(%v) = %v, want %v", x, dS, math.Sin(math.Pi*x*x/2)) + } + } + // Odd symmetry through the dispatch, and the series below the + // crossover. + for _, x := range []float64{0.5, 3, 11, 1000} { + if got, want := fresnelC(-x), -fresnelC(x); got != want { + t.Errorf("C(-%v) = %v, want %v", x, got, want) + } + if got, want := fresnelS(-x), -fresnelS(x); got != want { + t.Errorf("S(-%v) = %v, want %v", x, got, want) + } + } +} + +// TestEllipticFAccuracy pins the incomplete integral near the +// singularity: the 64-point Gauss-Legendre rule it used to run lost +// accuracy as m tends to 1 (1e-7 relative at m = 1−1e-4, 3e-3 at +// 1−1e-6, growing from there), because a product rule cannot follow a +// square-root endpoint singularity. Carlson's R_F is exact there, so +// what is left is the conditioning of the float64 evaluation: +// F(π/2_float, m) differs from K(m) by the π/2 rounding amplified by +// 1/√(1−m). The AGM is the independent in-repo reference; the fixed +// points come from mpmath's ellipf at 30 digits on the same float64 +// inputs (both m = k²). +func TestEllipticFAccuracy(t *testing.T) { + const agmTol = 1e-10 + for _, m := range []float64{0, 0.1, 0.5, 0.9, 0.99, 0.999, 0.9999, 1 - 1e-4, 1 - 1e-6, 1 - 1e-8, 1 - 1e-12} { + got := ellipticF(math.Pi/2, m) + want := EllipticKScalar(m) + if rel := math.Abs(got-want) / want; rel > agmTol { + t.Errorf("F(π/2, %v) = %.17g, AGM K says %.17g (relative %.2g)", m, got, want, rel) + } + } + for _, tc := range []struct { + phi, m, want float64 + }{ + {1.0, 0.999999, 1.2261907568130584099}, + {2.0, 0.9999999999, 24.274987126565337874}, + {0.7, 0.5, 0.72877030571819021448}, + {2.0, 0.9, 3.7114432264647169119}, + } { + got := ellipticF(tc.phi, tc.m) + if rel := math.Abs(got-tc.want) / tc.want; rel > 1e-12 { + t.Errorf("F(%v, %v) = %.17g, mpmath says %.17g (relative %.2g)", + tc.phi, tc.m, got, tc.want, rel) + } + } + // Near the singular end of the parameter range the agreement with + // an exact evaluation is bounded by the conditioning: the + // tolerance below is eps/√(1−m), the same factor that separates + // F(π/2_float, m) from K(m). + for _, tc := range []struct { + phi, m, want float64 + }{ + {math.Pi / 2, 1 - 1e-6, 8.2940514636010009696}, + {math.Pi / 2, 1 - 1e-8, 10.59663475457466848}, + {math.Pi / 2, 1 - 1e-12, 15.201815980008887263}, + } { + got := ellipticF(tc.phi, tc.m) + tol := 1e-13 / math.Sqrt(1-tc.m) + if rel := math.Abs(got-tc.want) / tc.want; rel > tol { + t.Errorf("F(%v, %v) = %.17g, mpmath says %.17g (relative %.2g, tolerance %.2g)", + tc.phi, tc.m, got, tc.want, rel, tol) + } + } + // The argument reduction: the integrand has period π, so F gains + // 2K over a period and is odd. + const m = 0.7 + k := EllipticKScalar(m) + base := ellipticF(0.4, m) + if got, want := ellipticF(0.4+math.Pi, m), base+2*k; math.Abs(got-want) > 1e-12*k { + t.Errorf("F(0.4+π, %v) = %v, want %v", m, got, want) + } + if got, want := ellipticF(-0.4, m), -base; math.Abs(got-want) > 1e-15 { + t.Errorf("F(−0.4, %v) = %v, want %v", m, got, want) + } + // Carlson's R_F itself, against mpmath elliprf at 40 digits on the + // same float64 arguments, including the extreme dynamic ranges the + // Fresnel-style series cannot handle. + for _, tc := range []struct { + x, y, z, want float64 + }{ + {1, 1, 1, 1.0}, + {0.25, 0.5, 1, 1.3701716332668719479}, + {0, 1e-12, 1, 15.201804919087715184}, + {3.75e-33, 1e-12, 1, 15.201804919026477941}, + {1e-30, 0.5, 1, 1.8540746773013705042}, + {1, 1e-14, 1e14, 1.7504389912078256668e-6}, + } { + got := carlsonRF(tc.x, tc.y, tc.z) + if rel := math.Abs(got-tc.want) / tc.want; rel > 1e-13 { + t.Errorf("R_F(%v, %v, %v) = %.17g, mpmath says %.17g (relative %.2g)", + tc.x, tc.y, tc.z, got, tc.want, rel) + } + } +} + +// TestEinsumLabelName pins the label spelling in the mismatch message: +// the ids run past z into the upper case letters, and a rune arithmetic +// slip printed punctuation instead. +func TestEinsumLabelName(t *testing.T) { + a, err := FromFloats(make([]float64, 6), 2, 3) + if err != nil { + t.Fatal(err) + } + b, err := FromFloats(make([]float64, 8), 2, 4) + if err != nil { + t.Fatal(err) + } + _, err = Einsum("aB,aB->aB", a, b) + if err == nil { + t.Fatal("expected a label mismatch") + } + if !strings.Contains(err.Error(), `label "B"`) { + t.Fatalf("error = %v, want the label spelled B", err) + } +} diff --git a/internal/core/complex_test.go b/internal/core/complex_test.go new file mode 100644 index 0000000..f22425e --- /dev/null +++ b/internal/core/complex_test.go @@ -0,0 +1,269 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "strings" + "testing" +) + +func mustFromComplexes(t *testing.T, vals []complex128, shape ...int) *Array { + t.Helper() + a, err := FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err) + } + return a +} + +func TestComplexConstructors(t *testing.T) { + a := mustFromComplexes(t, []complex128{complex(1, 2), complex(3, -4)}, 2) + if a.Dtype() != Complex || a.Len() != 2 { + t.Fatalf("complex array: %s len %d", a.Dtype(), a.Len()) + } + if v, err := ComplexAt(a, 0); err != nil || v != complex(1, 2) { + t.Fatalf("ComplexAt: %v %v", v, err) + } + if _, err := IntAt(a, 0); err == nil || !strings.Contains(err.Error(), "not int") { + t.Fatalf("IntAt on complex: %v", err) + } + if _, err := FloatAt(a, 0); err == nil || !strings.Contains(err.Error(), "not float") { + t.Fatalf("FloatAt on complex: %v", err) + } + + z, err := Zeros(Complex, 2) + if err != nil { + t.Fatalf("Zeros complex: %v", err) + } + if v, _ := ComplexAt(z, 1); v != 0 { + t.Fatalf("Zeros complex value: %v", v) + } + + one, _ := Ones(Complex, 2) + if v, _ := ComplexAt(one, 0); v != 1 { + t.Fatalf("Ones complex value: %v", v) + } + + full, _ := FullC(complex(0.5, -0.5), 2) + if v, _ := ComplexAt(full, 1); v != complex(0.5, -0.5) { + t.Fatalf("FullC value: %v", v) + } + + id, _ := Identity(Complex, 2) + if v, _ := ComplexAt(id, 1, 1); v != 1 { + t.Fatalf("Identity complex: %v", v) + } + if v, _ := ComplexAt(id, 0, 1); v != 0 { + t.Fatalf("Identity complex off-diagonal: %v", v) + } +} + +func TestComplexPromotion(t *testing.T) { + c := mustFromComplexes(t, []complex128{complex(1, 1)}, 1) + i := mustFromInts(t, []int64{2}, 1) + f := mustFromFloats(t, []float64{0.5}, 1) + + sum, err := Add(c, i) + if err != nil || sum.Dtype() != Complex { + t.Fatalf("complex + int: %s %v", sum.Dtype(), err) + } + if v, _ := ComplexAt(sum, 0); v != complex(3, 1) { + t.Fatalf("complex + int value: %v", v) + } + + // float + complex promotes too. + sum2, _ := Add(f, c) + if sum2.Dtype() != Complex { + t.Fatalf("float + complex dtype: %s", sum2.Dtype()) + } + if v, _ := ComplexAt(sum2, 0); v != complex(1.5, 1) { + t.Fatalf("float + complex value: %v", v) + } + + // Division on complex stays complex. + q, _ := Div(c, f) + if q.Dtype() != Complex { + t.Fatalf("complex division dtype: %s", q.Dtype()) + } + if v, _ := ComplexAt(q, 0); v != complex(1, 1)/complex(0.5, 0) { + t.Fatalf("complex division value: %v", v) + } + + // Quo rejects complex. + if _, err := Quo(c, c); err == nil || !strings.Contains(err.Error(), "needs int arrays") { + t.Fatalf("Quo complex: %v", err) + } +} + +func TestComplexScalarMirrors(t *testing.T) { + c := mustFromComplexes(t, []complex128{complex(1, 2)}, 1) + i := mustFromInts(t, []int64{1}, 1) + + // Int scalars keep complex complex. + if v, _ := ComplexAt(AddI(c, 5), 0); v != complex(6, 2) { + t.Fatalf("AddI on complex: %v", v) + } + if v, _ := ComplexAt(MulI(c, 2), 0); v != complex(2, 4) { + t.Fatalf("MulI on complex: %v", v) + } + if v, _ := ComplexAt(DivF(c, 2), 0); v != complex(0.5, 1) { + t.Fatalf("DivF on complex: %v", v) + } + + // Complex scalars promote everything. + if v, _ := ComplexAt(AddC(i, complex(0, 1)), 0); v != complex(1, 1) { + t.Fatalf("AddC on int: %v", v) + } + out := MulC(i, complex(0, 2)) + if out.Dtype() != Complex { + t.Fatalf("MulC dtype: %s", out.Dtype()) + } + if v, _ := ComplexAt(out, 0); v != complex(0, 2) { + t.Fatalf("MulC value: %v", v) + } + if v, _ := ComplexAt(SubC(c, complex(1, 2)), 0); v != 0 { + t.Fatalf("SubC value: %v", v) + } + if v, _ := ComplexAt(DivC(c, complex(0, 1)), 0); v != complex(2, -1) { + t.Fatalf("DivC value: %v", v) + } + if _, err := QuoI(i, 1); err != nil { + t.Fatalf("QuoI still works on int: %v", err) + } +} + +func TestComplexSumAndDot(t *testing.T) { + c := mustFromComplexes(t, []complex128{complex(1, 1), complex(2, -1)}, 2) + s := Sum(c) + if !s.IsComplex() || s.Complex() != complex(3, 0) { + t.Fatalf("Sum complex: %s", s) + } + + d, err := Dot(c, c) + if err != nil { + t.Fatalf("Dot complex: %v", err) + } + // (1+i)(1+i) + (2-i)(2-i) = 2i + 3-4i = 3-2i + if !d.IsComplex() || d.Complex() != complex(3, -2) { + t.Fatalf("Dot complex: %s", d) + } + + // Mixed Dot promotes to complex. + dm, _ := Dot(mustFromInts(t, []int64{1, 1}, 2), c) + if !dm.IsComplex() { + t.Fatalf("Dot mixed: %s", dm) + } +} + +func TestComplexUnorderedErrors(t *testing.T) { + c := mustFromComplexes(t, []complex128{complex(1, 1)}, 1) + + if _, err := Min(c); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("Min complex: %v", err) + } + if _, err := Max(c); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("Max complex: %v", err) + } + if _, err := Mean(c); err == nil || !strings.Contains(err.Error(), "no float mean") { + t.Fatalf("Mean complex: %v", err) + } + if _, err := Lt(c, c); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("Lt complex: %v", err) + } + if _, err := GtI(c, 1); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("GtI complex: %v", err) + } + if _, err := LtF(c, 1); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("LtF complex: %v", err) + } +} + +func TestComplexEqNe(t *testing.T) { + a := mustFromComplexes(t, []complex128{complex(1, 2), complex(3, 4)}, 2) + b := mustFromComplexes(t, []complex128{complex(1, 2), complex(4, 4)}, 2) + + eq, err := Eq(a, b) + if err != nil { + t.Fatalf("Eq complex: %v", err) + } + if got := maskOf(t, eq); got[0] != 1 || got[1] != 0 { + t.Fatalf("Eq complex: %v", got) + } + ne, _ := Ne(a, b) + if got := maskOf(t, ne); got[0] != 0 || got[1] != 1 { + t.Fatalf("Ne complex: %v", got) + } + + // Complex vs float compares in complex space; only the element whose + // imaginary part is zero can match. + realish := mustFromComplexes(t, []complex128{1, complex(3, 4)}, 2) + f := mustFromFloats(t, []float64{1, 4}, 2) + eqf, _ := Eq(realish, f) + if got := maskOf(t, eqf); got[0] != 1 || got[1] != 0 { + t.Fatalf("Eq complex-float: %v", got) + } +} + +func TestComplexMatMulAndMachinery(t *testing.T) { + m := mustFromComplexes(t, []complex128{complex(0, 1), 0, 0, complex(0, 1)}, 2, 2) + v := mustFromComplexes(t, []complex128{1, 1}, 2, 1) + + out, err := MatMul2D(m, v) + if err != nil { + t.Fatalf("MatMul complex: %v", err) + } + if out.Dtype() != Complex || out.Shape()[0] != 2 { + t.Fatalf("MatMul complex shape: %s", out) + } + if w, _ := ComplexAt(out, 0, 0); w != complex(0, 1) { + t.Fatalf("MatMul complex value: %v", w) + } + + // Transpose, Slice, Row, Mask and WithComplex all carry the dtype. + tt := Transpose(m) + if tt.Dtype() != Complex { + t.Fatalf("Transpose complex dtype: %s", tt.Dtype()) + } + s, _ := Slice(m, 0, 0, 1) + if s.Dtype() != Complex || s.Shape()[0] != 1 { + t.Fatalf("Slice complex: %s", s) + } + r, _ := Row(m, 1) + if v2, _ := ComplexAt(r, 1); v2 != complex(0, 1) { + t.Fatalf("Row complex: %v", v2) + } + upd, err := WithComplex(m, complex(9, 9), 0, 0) + if err != nil { + t.Fatalf("WithComplex: %v", err) + } + if v2, _ := ComplexAt(upd, 0, 0); v2 != complex(9, 9) { + t.Fatalf("WithComplex value: %v", v2) + } + if v2, _ := ComplexAt(m, 0, 0); v2 != complex(0, 1) { + t.Fatalf("WithComplex receiver mutated: %v", v2) + } + + // Where promotes to complex and Mask selects complex elements. + cond := mustFromInts(t, []int64{1, 0}, 2) + w, _ := Where(cond, mustFromComplexes(t, []complex128{1, 1}, 2), mustFromComplexes(t, []complex128{0, 0}, 2)) + if w.Dtype() != Complex { + t.Fatalf("Where complex dtype: %s", w.Dtype()) + } + if v2, _ := ComplexAt(w, 1); v2 != 0 { + t.Fatalf("Where complex value: %v", v2) + } + mask := cond + sel, _ := Select(mustFromComplexes(t, []complex128{complex(5, 5), complex(6, 6)}, 2), mask) + if sel.Len() != 1 { + t.Fatalf("Mask complex len: %d", sel.Len()) + } + if v2, _ := ComplexAt(sel, 0); v2 != complex(5, 5) { + t.Fatalf("Mask complex value: %v", v2) + } + + // String renders complex elements. + if got := mustFromComplexes(t, []complex128{complex(1, 2)}, 1).String(); !strings.Contains(got, "(1+2i)") { + t.Fatalf("String complex: %q", got) + } +} diff --git a/internal/core/concat.go b/internal/core/concat.go new file mode 100644 index 0000000..28b5858 --- /dev/null +++ b/internal/core/concat.go @@ -0,0 +1,285 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +import "math" + +// Concatenation and stacking. Concat joins along an existing +// dimension, Stack along a new one; both copy, the result dtype promotes +// per the ladder, and every shape disagreement names both shapes. + +// Concat joins b after a along the given existing dimension. Every other +// dimension must agree. +func Concat(a, b *Array, dim int) (*Array, error) { + if a.NDim() != b.NDim() { + return nil, errf("Concat: ranks differ, %s and %s", shapeText(a.shape), shapeText(b.shape)) + } + if dim < 0 || dim >= a.NDim() { + return nil, errf("Concat: dimension %d is out of range for shape %s", dim, shapeText(a.shape)) + } + for d := range a.shape { + if d != dim && a.shape[d] != b.shape[d] { + return nil, errf("Concat: shapes %s and %s disagree outside dimension %d", + shapeText(a.shape), shapeText(b.shape), dim) + } + } + + newShape := a.Shape() + // Bound the joined length before adding it: a wrapped sum would pair + // a negative dimension with an unvalidated allocation. + if b.shape[dim] > math.MaxInt-a.shape[dim] { + return nil, errf("Concat: joining along dimension %d overflows: %d + %d", dim, a.shape[dim], b.shape[dim]) + } + newShape[dim] = a.shape[dim] + b.shape[dim] + total, _, terr := checkedDims(newShape) + if terr != nil { + return nil, base.WrapErr("Concat", terr) + } + dt := promote(a.dt, b.dt) + out := &Array{shape: newShape, dt: dt} + out.alloc(total) + if dt == a.dt && dt == b.dt && a.isContiguous() && b.isContiguous() { + // No promotion: every outer position contributes one run of each + // operand, so the join is two block copies per position rather + // than a converted element per slot. copyRun dispatches every + // dtype the package stores, the narrow payloads included. + tail := 1 + for d := dim + 1; d < len(newShape); d++ { + tail *= newShape[d] + } + head := 1 + for d := range dim { + head *= newShape[d] + } + na, nb := a.shape[dim], b.shape[dim] + for h := range head { + base := h * newShape[dim] * tail + copyRun(out, a, base, h*na*tail, na*tail) + copyRun(out, b, base+na*tail, h*nb*tail, nb*tail) + } + return out, nil + } + if a.isContiguous() && b.isContiguous() { + // Promotion: every outer position still contributes one run of + // each operand, exactly as the no-promotion path above lays the + // join out, so the join is two bulk conversions per position + // rather than a setConverted dispatch per element. The + // conversions are the ones setConverted performs, value for + // value, and the run layout is the same walk's. + tail := 1 + for d := dim + 1; d < len(newShape); d++ { + tail *= newShape[d] + } + head := 1 + for d := range dim { + head *= newShape[d] + } + na, nb := a.shape[dim], b.shape[dim] + for h := range head { + base := h * newShape[dim] * tail + concatConvert(out, a, base, h*na*tail, na*tail) + concatConvert(out, b, base+na*tail, h*nb*tail, nb*tail) + } + return out, nil + } + coord := make([]int, len(newShape)) + for i := range total { + src := a + c := coord[dim] + if c >= a.shape[dim] { + src = b + c -= a.shape[dim] + } + // Flat offset into the source: same walk with the local dim + // coordinate and the source's own dimension sizes. + off := 0 + for d := range coord { + cd := coord[d] + if d == dim { + cd = c + } + off = off*src.shape[d] + cd + } + out.setConverted(i, src, off) + advanceOdometer(coord, newShape) + } + return out, nil +} + +// concatConvert converts n elements of a contiguous src, read from +// srcOff, into dst's dtype at dstOff, the exact conversions setConverted +// performs for the same pair: the source's payload is read where it +// lies and the promoted store is a plain loop, so no element pays a +// second dispatch or a stride resolution. The conversion arms mirror +// setConverted arm for arm. +func concatConvert(dst, src *Array, dstOff, srcOff, n int) { + switch dst.dt { + case Float: + d := dst.floats[dstOff : dstOff+n] + switch src.dt { + case Float: + copy(d, src.floats[srcOff:srcOff+n]) + case Float32: + p := src.floats32[srcOff : srcOff+n] + for i := range d { + d[i] = float64(p[i]) + } + case Float16: + p := src.halves[srcOff : srcOff+n] + for i := range d { + d[i] = HalfToFloat64(p[i]) + } + case Int: + p := src.ints[srcOff : srcOff+n] + for i := range d { + d[i] = float64(p[i]) + } + default: + // Bool and the narrow integers: floatAt widens each exactly. + for i := range d { + d[i] = src.floatAt(srcOff + i) + } + } + case Complex: + d := dst.complexes[dstOff : dstOff+n] + if src.dt == Complex { + copy(d, src.complexes[srcOff:srcOff+n]) + return + } + for i := range d { + d[i] = src.complexAt(srcOff + i) + } + case Int: + d := dst.ints[dstOff : dstOff+n] + switch src.dt { + case Int: + // Exact: the float64 detour rounds above 2^53. + copy(d, src.ints[srcOff:srcOff+n]) + case Float: + p := src.floats[srcOff : srcOff+n] + for i := range d { + d[i] = int64(p[i]) + } + case Float16: + p := src.halves[srcOff : srcOff+n] + for i := range d { + d[i] = int64(HalfToFloat64(p[i])) + } + case Float32: + p := src.floats32[srcOff : srcOff+n] + for i := range d { + d[i] = int64(p[i]) + } + default: + for i := range d { + d[i] = src.intAt(srcOff + i) + } + } + case Float32: + d := dst.floats32[dstOff : dstOff+n] + if src.dt == Float32 { + copy(d, src.floats32[srcOff:srcOff+n]) + return + } + for i := range d { + d[i] = src.float32At(srcOff + i) + } + case Float16: + d := dst.halves[dstOff : dstOff+n] + for i := range d { + d[i] = HalfFromFloat64(src.floatAt(srcOff + i)) + } + default: + // The narrow integer destinations: setConverted casts the + // widened intAt read down, the implicit-store semantics every + // promoted join target has always carried. + for i := range n { + dst.setConverted(dstOff+i, src, srcOff+i) + } + } +} + +// Stack joins b after a along a new dimension inserted at dim; the two +// arrays must have identical shapes. +func Stack(a, b *Array, dim int) (*Array, error) { + if !sameShape(a.shape, b.shape) { + return nil, errf("Stack: shapes %s and %s must be identical", + shapeText(a.shape), shapeText(b.shape)) + } + if dim < 0 || dim > a.NDim() { + return nil, errf("Stack: dimension %d is out of range for inserting into shape %s", + dim, shapeText(a.shape)) + } + + newShape := make([]int, 0, a.NDim()+1) + newShape = append(newShape, a.shape[:dim]...) + newShape = append(newShape, 2) + newShape = append(newShape, a.shape[dim:]...) + total := 1 + for _, d := range newShape { + total *= d + } + dt := promote(a.dt, b.dt) + out := &Array{shape: newShape, dt: dt} + out.alloc(total) + tail := 1 + for d := dim; d < a.NDim(); d++ { + tail *= a.shape[d] + } + head := 1 + for d := range dim { + head *= a.shape[d] + } + if dt == a.dt && dt == b.dt && a.isContiguous() && b.isContiguous() { + // The stacked axis is a factor of two: each outer position holds + // one run of each operand, so the join is two block copies. + // copyRun dispatches every dtype the package stores, the narrow + // payloads included. + for h := range head { + copyRun(out, a, h*2*tail, h*tail, tail) + copyRun(out, b, h*2*tail+tail, h*tail, tail) + } + return out, nil + } + if a.isContiguous() && b.isContiguous() { + // Promotion: each outer position still holds one run of each + // operand, so the join is two bulk conversions per position + // rather than a setConverted dispatch per element, the + // conversions concatConvert keeps. + for h := range head { + concatConvert(out, a, h*2*tail, h*tail, tail) + concatConvert(out, b, h*2*tail+tail, h*tail, tail) + } + return out, nil + } + coord := make([]int, len(newShape)) + for i := range total { + src := a + if coord[dim] == 1 { + src = b + } + // Flat offset into the source: drop the stacked coordinate. + off := 0 + for d := range coord { + if d == dim { + continue + } + off = off*src.shape[stackDim(d, dim)] + coord[d] + } + out.setConverted(i, src, off) + advanceOdometer(coord, newShape) + } + return out, nil +} + +// stackDim maps a new-shape dimension to the source-shape dimension, +// skipping the inserted axis. +func stackDim(d, dim int) int { + if d > dim { + return d - 1 + } + return d +} diff --git a/internal/core/concat_test.go b/internal/core/concat_test.go new file mode 100644 index 0000000..f2c6121 --- /dev/null +++ b/internal/core/concat_test.go @@ -0,0 +1,232 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "slices" + "strings" + "testing" +) + +func TestConcat(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + b := mustFromInts(t, []int64{5, 6, 7, 8}, 2, 2) + + // Along columns: rows grow wider. + wide, err := Concat(a, b, 1) + if err != nil { + t.Fatalf("Concat: %v", err) + } + want := mustFromInts(t, []int64{1, 2, 5, 6, 3, 4, 7, 8}, 2, 4) + if !Equal(want, wide) { + t.Fatalf("Concat dim 1: %s", wide) + } + + // Along rows: the array grows taller. + tall, err := Concat(a, b, 0) + if err != nil { + t.Fatalf("Concat dim 0: %v", err) + } + wantTall := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 4, 2) + if !Equal(wantTall, tall) { + t.Fatalf("Concat dim 0: %s", tall) + } + + // Promotion across operands. + f := mustFromFloats(t, []float64{0.5, 0.5}, 1, 2) + mixed, err := Concat(a, f, 0) + if err != nil || mixed.Dtype() != Float { + t.Fatalf("Concat promote: %s %v", mixed, err) + } + if v, _ := FloatAt(mixed, 2, 0); v != 0.5 { + t.Fatalf("Concat promote value: %v", v) + } + + if _, err := Concat(a, b, 2); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("Concat dim range: %v", err) + } + other := mustFromInts(t, []int64{1, 2, 3}, 1, 3) + if _, err := Concat(a, other, 0); err == nil || !strings.Contains(err.Error(), "disagree outside dimension") { + t.Fatalf("Concat shape: %v", err) + } + v := mustFromInts(t, []int64{1}, 1) + if _, err := Concat(a, v, 0); err == nil || !strings.Contains(err.Error(), "ranks differ") { + t.Fatalf("Concat rank: %v", err) + } +} + +func TestStack(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + b := mustFromInts(t, []int64{5, 6, 7, 8}, 2, 2) + + s, err := Stack(a, b, 0) + if err != nil { + t.Fatalf("Stack: %v", err) + } + if shape := s.Shape(); shape[0] != 2 || shape[1] != 2 || shape[2] != 2 { + t.Fatalf("Stack shape: %v", shape) + } + if v, _ := IntAt(s, 0, 0, 0); v != 1 { + t.Fatalf("Stack[0,0,0]: %d", v) + } + if v, _ := IntAt(s, 1, 1, 1); v != 8 { + t.Fatalf("Stack[1,1,1]: %d", v) + } + + // Inserting a middle axis keeps the source order. + mid, err := Stack(a, b, 1) + if err != nil { + t.Fatalf("Stack mid: %v", err) + } + if shape := mid.Shape(); shape[0] != 2 || shape[1] != 2 || shape[2] != 2 { + t.Fatalf("Stack mid shape: %v", shape) + } + // mid[0][1][1] is b[0][1]; the inserted axis selects the operand. + if v, _ := IntAt(mid, 0, 1, 1); v != 6 { + t.Fatalf("Stack mid value: %d", v) + } + + // Promotion and errors. + f := mustFromFloats(t, []float64{9, 9, 9, 9}, 2, 2) + sf, err := Stack(a, f, 2) + if err != nil || sf.Dtype() != Float { + t.Fatalf("Stack promote: %s %v", sf, err) + } + if _, err := Stack(a, mustFromInts(t, []int64{1}, 1), 0); err == nil || !strings.Contains(err.Error(), "must be identical") { + t.Fatalf("Stack shape: %v", err) + } + if _, err := Stack(a, b, 3); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("Stack dim: %v", err) + } +} + +// TestConcatStackUnequalParts pins the run copies of Concat and Stack +// when the operands differ in length along the joined axis and several +// outer positions precede it: the second operand's run starts at a +// different source offset per position and its length is its own. +func TestConcatStackUnequalParts(t *testing.T) { + av := make([]int64, 12) + for i := range av { + av[i] = int64(i + 1) + } + bv := make([]int64, 18) + for i := range bv { + bv[i] = int64(101 + i) + } + a := mustFromInts(t, av, 2, 2, 3) + b := mustFromInts(t, bv, 2, 3, 3) + + c, err := Concat(a, b, 1) + if err != nil { + t.Fatalf("Concat unequal dim 1: %v", err) + } + if sh := c.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 5 || sh[2] != 3 { + t.Fatalf("Concat unequal dim 1 shape: %v, want [2 5 3]", sh) + } + for i := range 2 { + for j := range 5 { + for k := range 3 { + var want int64 + if j < 2 { + want, _ = IntAt(a, i, j, k) + } else { + want, _ = IntAt(b, i, j-2, k) + } + if got, _ := IntAt(c, i, j, k); got != want { + t.Fatalf("Concat unequal dim 1 [%d %d %d] = %d, want %d", i, j, k, got, want) + } + } + } + } + + // The last axis, where every leading position is its own head: two + // elements of one operand and one of the other. + av2 := make([]int64, 12) + for i := range av2 { + av2[i] = int64(i + 1) + } + a2 := mustFromInts(t, av2, 2, 3, 2) + b2 := mustFromInts(t, []int64{201, 202, 203, 204, 205, 206}, 2, 3, 1) + c2, err := Concat(a2, b2, 2) + if err != nil { + t.Fatalf("Concat unequal dim 2: %v", err) + } + if sh := c2.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 3 || sh[2] != 3 { + t.Fatalf("Concat unequal dim 2 shape: %v, want [2 3 3]", sh) + } + for i := range 2 { + for j := range 3 { + for k := range 3 { + var want int64 + if k < 2 { + want, _ = IntAt(a2, i, j, k) + } else { + want, _ = IntAt(b2, i, j, k-2) + } + if got, _ := IntAt(c2, i, j, k); got != want { + t.Fatalf("Concat unequal dim 2 [%d %d %d] = %d, want %d", i, j, k, got, want) + } + } + } + } + + // Stack inserts an axis of two: each outer position contributes one + // run of each operand. + s1 := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + s2 := mustFromInts(t, []int64{11, 12, 13, 14, 15, 16}, 2, 3) + st, err := Stack(s1, s2, 1) + if err != nil { + t.Fatalf("Stack: %v", err) + } + if sh := st.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 2 || sh[2] != 3 { + t.Fatalf("Stack shape: %v, want [2 2 3]", sh) + } + for i := range 2 { + for k := range 3 { + if got, _ := IntAt(st, i, 0, k); got != int64(i*3+k+1) { + t.Fatalf("Stack [%d 0 %d] = %d, want %d", i, k, got, i*3+k+1) + } + if got, _ := IntAt(st, i, 1, k); got != int64(11+i*3+k) { + t.Fatalf("Stack [%d 1 %d] = %d, want %d", i, k, got, 11+i*3+k) + } + } + } +} + +// TestConcatNarrowDtypes pins that the narrow payloads ride the +// converted walk until the run copy dispatches them: same-dtype joins +// keep their dtype and mixed joins promote through the containment +// table. +func TestConcatNarrowDtypes(t *testing.T) { + a, err := FromInt8s([]int8{1, 2}, 2) + if err != nil { + t.Fatal(err) + } + b, err := FromInt8s([]int8{3}, 1) + if err != nil { + t.Fatal(err) + } + joined, err := Concat(a, b, 0) + if err != nil { + t.Fatalf("Concat int8: %v", err) + } + if joined.Dtype() != Int8 { + t.Fatalf("Concat int8 answered %s, want int8", joined.Dtype()) + } + if want := []int8{1, 2, 3}; !slices.Equal(joined.RawInt8s(), want) { + t.Fatalf("Concat int8 = %v, want %v", joined.RawInt8s(), want) + } + // Mixed signedness promotes: int8 with uint8 answers int16. + u, err := FromUint8s([]uint8{200}, 1) + if err != nil { + t.Fatal(err) + } + mix, err := Concat(a, u, 0) + if err != nil { + t.Fatalf("Concat int8 with uint8: %v", err) + } + if mix.Dtype() != Int16 || !slices.Equal(mix.RawInt16s(), []int16{1, 2, 200}) { + t.Fatalf("Concat int8 with uint8 = %s %v", mix.Dtype(), mix.RawInt16s()) + } +} diff --git a/internal/core/cosm1_test.go b/internal/core/cosm1_test.go new file mode 100644 index 0000000..b6ab3cd --- /dev/null +++ b/internal/core/cosm1_test.go @@ -0,0 +1,102 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + "testing" +) + +// cosm1Big evaluates cos(x) − 1 by the Taylor series in 256-bit +// arithmetic, the exact referent the float64 implementation is held +// against across the crossover. +func cosm1Big(x float64) *big.Float { + const prec = 256 + xf := new(big.Float).SetPrec(prec).SetFloat64(x) + sq := new(big.Float).SetPrec(prec).Mul(xf, xf) + sum := new(big.Float).SetPrec(prec).SetInt64(1) + term := new(big.Float).SetPrec(prec).SetInt64(1) + for k := 1; k <= 60; k++ { + term.Mul(term, sq) + term.Quo(term, new(big.Float).SetPrec(prec).SetInt64(int64(2*k-1))) + term.Quo(term, new(big.Float).SetPrec(prec).SetInt64(int64(2*k))) + term.Neg(term) + sum.Add(sum, term) + } + return sum.Sub(sum, new(big.Float).SetPrec(prec).SetInt64(1)) +} + +// TestCosm1AgainstSeries holds the implementation against the exact +// referent on both sides of the crossover, from arguments whose answer +// is −x²/2 as far as the format can see up to ones where the direct +// subtraction carries it alone. +func TestCosm1AgainstSeries(t *testing.T) { + points := []float64{ + 1e-300, 1e-200, 5e-12, 1e-10, 1e-9, 1e-8, 1e-6, 1e-4, 0.01, 0.1, + 0.3, 0.5, 0.78, math.Pi / 4, 0.79, 1, 2, -0.3, -0.78, -2, + } + for _, x := range points { + got, err := Cosm1(mustFloats(t, []float64{x})) + if err != nil { + t.Fatalf("Cosm1(%v): %v", x, err) + } + want, _ := cosm1Big(x).Float64() + v := got.FloatAt(0) + if want == 0 { + if v != 0 { + t.Fatalf("Cosm1(%v) = %v, want 0", x, v) + } + continue + } + if d := math.Abs(v-want) / math.Abs(want); d > 6e-16 { + t.Fatalf("Cosm1(%v) = %.17g, want %.17g (relative %.3g)", x, v, want, d) + } + } +} + +// TestCosm1Pins pins the behaviour the referent cannot speak for: the +// exact zero at the origin, the underflowed answer keeping the minus +// sign the function's range promises, the bit-equality with the direct +// subtraction above the crossover, and the integer promotion. +func TestCosm1Pins(t *testing.T) { + zero, err := Cosm1(mustFloats(t, []float64{0})) + if err != nil { + t.Fatal(err) + } + if zero.FloatAt(0) != 0 { + t.Fatalf("Cosm1(0) = %v, want 0", zero.FloatAt(0)) + } + tiny, err := Cosm1(mustFloats(t, []float64{1e-300})) + if err != nil { + t.Fatal(err) + } + if tiny.FloatAt(0) != 0 || !math.Signbit(tiny.FloatAt(0)) { + t.Fatalf("Cosm1(1e-300) = %v, want the negative zero the true answer underflows to", tiny.FloatAt(0)) + } + for _, x := range []float64{0.79, 1, 2, 10, 100, 1e6, -3.5} { + got, err := Cosm1(mustFloats(t, []float64{x})) + if err != nil { + t.Fatalf("Cosm1(%v): %v", x, err) + } + if want := math.Cos(x) - 1; got.FloatAt(0) != want { + t.Fatalf("Cosm1(%v) = %.17g, want the direct %.17g", x, got.FloatAt(0), want) + } + } + ints, err := FromInts([]int64{0, 1}, 2) + if err != nil { + t.Fatal(err) + } + promoted, err := Cosm1(ints) + if err != nil { + t.Fatalf("Cosm1 over ints: %v", err) + } + if promoted.FloatAt(1) != math.Cos(1)-1 { + t.Fatalf("Cosm1(int 1) = %.17g, want %.17g", promoted.FloatAt(1), math.Cos(1)-1) + } + cx, _ := FromComplexes([]complex128{1}, 1) + if _, err := Cosm1(cx); err == nil { + t.Fatal("Cosm1 over complex: want an error") + } +} diff --git a/internal/core/cumprod_half_test.go b/internal/core/cumprod_half_test.go new file mode 100644 index 0000000..25b2b32 --- /dev/null +++ b/internal/core/cumprod_half_test.go @@ -0,0 +1,29 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// TestCumProdHalfNarrowsCarryPerStep pins the float16 cumulative +// product's carry discipline: every step narrows the running product to +// half before combining, exactly as the stored element narrows, so each +// step multiplies the rounded carry. The input 67, 65, 6 makes the +// discipline visible: 67 times 65 is 4355 in float64, which narrows to +// 4356, and 4356 times 6 rounds to 26144; a carry left in float64 would +// give 4355 times 6 = 26130, which rounds to 26128. +func TestCumProdHalfNarrowsCarryPerStep(t *testing.T) { + a := mustFromHalves(t, []uint16{ + HalfFromFloat64(67), HalfFromFloat64(65), HalfFromFloat64(6), + }, 3) + out, err := CumProd(a, 0) + if err != nil { + t.Fatalf("CumProd: %v", err) + } + want := []float64{67, 4356, 26144} + for i, w := range want { + if got := HalfToFloat64(out.RawHalves()[i]); got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } +} diff --git a/internal/core/data_test.go b/internal/core/data_test.go new file mode 100644 index 0000000..5228d3e --- /dev/null +++ b/internal/core/data_test.go @@ -0,0 +1,42 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "testing" +) + +func TestAllAnyCountNonzero(t *testing.T) { + z, _ := FromInts([]int64{0, 0, 0}, 3) + a, err := All(z) + if err != nil || a { + t.Errorf("All zeros: %v %v", a, err) + } + any, err := Any(z) + if err != nil || any { + t.Errorf("Any zeros: %v %v", any, err) + } + m, _ := FromInts([]int64{0, 5, 0, 7}, 2, 2) + all, err := All(m) + if err != nil || all { + t.Errorf("All mixed: %v %v", all, err) + } + anyM, err := Any(m) + if err != nil || !anyM { + t.Errorf("Any mixed: %v %v", anyM, err) + } + cnt, err := CountNonzero(m) + if err != nil || cnt != 2 { + t.Errorf("CountNonzero: %d %v, want 2", cnt, err) + } + allOnes, _ := FromInts([]int64{1, 2, 3}, 3) + allT, err := All(allOnes) + if err != nil || !allT { + t.Errorf("All non-zero: %v %v", allT, err) + } + c, _ := FromComplexes([]complex128{1 + 2i}, 1) + if _, err := All(c); err == nil { + t.Error("All complex: expected error") + } +} diff --git a/internal/core/defect_family_pins_test.go b/internal/core/defect_family_pins_test.go new file mode 100644 index 0000000..be74220 --- /dev/null +++ b/internal/core/defect_family_pins_test.go @@ -0,0 +1,504 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Regression pins for a defect family: rebased views whose +// payload runs past Len, the trigamma reflection sign, Miller recurrence +// seeds for negative arguments, empty-dimension guards, complex-dtype +// rejection, NaN-first axis extrema and sparse matmul validation. + +// floatTailView returns the 1-D view of f's first n elements. The view's +// payload keeps every slot of f, longer than the view's own Len, +// exactly as Slice's payload rebasing leaves it: the tail slots belong +// to the parent and must stay invisible to every kernel. +func floatTailView(t *testing.T, f *Array, n int) *Array { + t.Helper() + v, err := Slice(f, 0, 0, n) + if err != nil { + t.Fatalf("Slice: %v", err) + } + if len(v.RawFloats()) <= v.Len() { + t.Fatalf("setup: expected a rebased view with tail slots, payload %d, len %d", len(v.RawFloats()), v.Len()) + } + return v +} + +// intTailView is the int twin of floatTailView. +func intTailView(t *testing.T, f *Array, n int) *Array { + t.Helper() + v, err := Slice(f, 0, 0, n) + if err != nil { + t.Fatalf("Slice: %v", err) + } + if len(v.RawInts()) <= v.Len() { + t.Fatalf("setup: expected a rebased view with tail slots, payload %d, len %d", len(v.RawInts()), v.Len()) + } + return v +} + +// TestViewSumMeanIgnoresTailSlots pins that Sum and Mean read exactly +// the view's own elements: the tail slots of a rebased payload hold the +// parent's values and must not contribute. +func TestViewSumMeanIgnoresTailSlots(t *testing.T) { + base := mustFromFloats(t, []float64{1, 2, 3, 4, 100, 100, 100, 100}, 8) + v := floatTailView(t, base, 4) + if got := Sum(v).Float(); got != 10 { + t.Errorf("Sum(view) = %v, want 10", got) + } + mean, err := Mean(v) + if err != nil || mean != 2.5 { + t.Errorf("Mean(view) = %v, %v, want 2.5", mean, err) + } + // The 2-D shape of the same bug: the first row of a matrix is a + // view whose payload still spans both rows. + mat := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 4) + row, err := Slice(mat, 0, 0, 1) + if err != nil { + t.Fatal(err) + } + if len(row.RawFloats()) <= row.Len() { + t.Fatalf("setup: row view payload %d, len %d", len(row.RawFloats()), row.Len()) + } + if got := Sum(row).Float(); got != 10 { + t.Errorf("Sum(row view) = %v, want 10", got) + } +} + +// TestViewIntReduceAndDotIgnoreTailSlots pins the int paths of Min, Max +// and Dot on rebased views: the old payload-length loops either read +// invisible tail slots or indexed past the other operand. +func TestViewIntReduceAndDotIgnoreTailSlots(t *testing.T) { + base := mustFromInts(t, []int64{9, 1, 3, 2, -5, 6, -7, 8}, 8) + v := intTailView(t, base, 4) + mn, err := Min(v) + if err != nil || mn.Int() != 1 { + t.Errorf("Min(view) = %s, %v, want 1", mn, err) + } + mx, err := Max(v) + if err != nil || mx.Int() != 9 { + t.Errorf("Max(view) = %s, %v, want 9", mx, err) + } + other := mustFromInts(t, []int64{1, 1, 1, 1}, 4) + d, err := Dot(v, other) + if err != nil || d.Int() != 15 { + t.Errorf("Dot(view, other) = %s, %v, want 15", d, err) + } +} + +// TestViewQuoPowValidateOwnElements pins the pre-validation loops of +// Quo and Pow: a zero divisor or negative exponent in the parent's tail +// slots must not reject an operation on the view's own elements. +func TestViewQuoPowValidateOwnElements(t *testing.T) { + divBase := mustFromInts(t, []int64{2, 3, 0, 0}, 8/2) + divisor := intTailView(t, divBase, 2) + got, err := Quo(mustFromInts(t, []int64{8, 9}, 2), divisor) + if err != nil { + t.Fatalf("Quo with a view divisor: %v", err) + } + if got.RawInts()[0] != 4 || got.RawInts()[1] != 3 { + t.Errorf("Quo = %v, want [4 3]", got.RawInts()) + } + expBase := mustFromInts(t, []int64{2, 3, -1, -1}, 4) + exponent := intTailView(t, expBase, 2) + pw, err := Pow(mustFromInts(t, []int64{2, 3}, 2), exponent) + if err != nil { + t.Fatalf("Pow with a view exponent: %v", err) + } + if pw.RawInts()[0] != 4 || pw.RawInts()[1] != 27 { + t.Errorf("Pow = %v, want [4 27]", pw.RawInts()) + } +} + +// TestViewOneHotAndSelectValidateOwnElements pins OneHot's code check +// and Select's hit count on rebased views: an out-of-range code or a +// hit living only in the tail slots must not change the result. +func TestViewOneHotAndSelectValidateOwnElements(t *testing.T) { + codesBase := mustFromInts(t, []int64{0, 2, 9, 9}, 4) + codes := intTailView(t, codesBase, 2) + hot, err := OneHot(codes, 3) + if err != nil { + t.Fatalf("OneHot on a view: %v", err) + } + if hot.Shape()[0] != 2 || hot.Shape()[1] != 3 { + t.Fatalf("OneHot shape: %v", hot.Shape()) + } + want := [][]float64{{1, 0, 0}, {0, 0, 1}} + for i := range 2 { + for j := range 3 { + if g := float64(hot.RawFloat32s()[i*3+j]); g != want[i][j] { + t.Errorf("OneHot[%d][%d] = %v, want %v", i, j, g, want[i][j]) + } + } + } + maskBase := mustFromInts(t, []int64{1, 0, 1, 1, 1, 1}, 6) + mask := intTailView(t, maskBase, 3) + sel, err := Select(mustFromFloats(t, []float64{10, 20, 30}, 3), mask) + if err != nil { + t.Fatalf("Select with a view mask: %v", err) + } + if sel.Len() != 2 || sel.RawFloats()[0] != 10 || sel.RawFloats()[1] != 30 { + t.Errorf("Select = %s, want the two hits [10, 30]", sel) + } +} + +// TestTrigammaReflection pins the reflection sign: for negative +// non-integer x, psi'(x) + psi'(1-x) = pi^2/sin^2(pi x), and the closed +// form psi'(-1/2) = pi^2/2 + 4 must come out near 8.9348. +func TestTrigammaReflection(t *testing.T) { + got := trigamma(-0.5) + if want := math.Pi*math.Pi/2 + 4; math.Abs(got-want) > 1e-10 { + t.Errorf("trigamma(-0.5) = %.16g, want %.16g", got, want) + } + for _, x := range []float64{-0.3, -1.7, -4.2} { + lhs := trigamma(x) + trigamma(1-x) + want := math.Pi * math.Pi / (math.Sin(math.Pi*x) * math.Sin(math.Pi*x)) + if math.Abs(lhs-want) > 1e-9*(1+math.Abs(want)) { + t.Errorf("reflection at %v: psi'(x)+psi'(1-x) = %.16g, want %.16g", x, lhs, want) + } + } + // The public entry point agrees. + arr, err := Trigamma(mustFloats(t, []float64{-0.5})) + if err != nil { + t.Fatal(err) + } + if v := arr.FloatAt(0); math.Abs(v-8.934802200544679) > 1e-10 { + t.Errorf("Trigamma(-0.5) = %.16g, want 8.934802200544679", v) + } +} + +// TestBesselNegativeArgumentMiller pins the Miller seed fix: the start +// order must sit above the requested order for negative arguments too, +// so j_l(-x) and I_n(-x) keep their parity values instead of +// collapsing to zero. +func TestBesselNegativeArgumentMiller(t *testing.T) { + neg := mustFloats(t, []float64{-100}) + pos := mustFloats(t, []float64{100}) + j2n, err := SphericalBesselJ(2, neg) + if err != nil { + t.Fatal(err) + } + if j2n.FloatAt(0) == 0 { + t.Fatal("SphericalBesselJ(2, -100) collapsed to 0") + } + j2p, _ := SphericalBesselJ(2, pos) + // Parity: j_2(-x) = j_2(x). + if r := math.Abs(j2n.FloatAt(0)-j2p.FloatAt(0)) / math.Abs(j2p.FloatAt(0)); r > 1e-9 { + t.Errorf("j_2(-100) = %v, want %v", j2n.FloatAt(0), j2p.FloatAt(0)) + } + in3n, err := BesselIn(3, neg) + if err != nil { + t.Fatal(err) + } + if in3n.FloatAt(0) == 0 { + t.Fatal("BesselIn(3, -100) collapsed to 0") + } + in3p, _ := BesselIn(3, pos) + // Parity: I_3(-x) = -I_3(x). + if r := math.Abs(in3n.FloatAt(0)+in3p.FloatAt(0)) / math.Abs(in3p.FloatAt(0)); r > 1e-9 { + t.Errorf("I_3(-100) = %v, want %v", in3n.FloatAt(0), -in3p.FloatAt(0)) + } +} + +// TestPadRejectsEmptyDimension pins the guard that keeps the folding +// modes from hanging or panicking on a zero-size dimension: reflect, +// replicate and circular are rejected up front, while constant mode +// simply fills. +func TestPadRejectsEmptyDimension(t *testing.T) { + empty := mustFromFloats(t, nil, 0, 3) + for _, mode := range []string{"reflect", "replicate", "circular"} { + if _, err := Pad(empty, []int{1, 1, 0, 0}, mode, 0); err == nil { + t.Errorf("Pad mode %q on an empty dimension: expected an error", mode) + } + } + // 2-D pad = (left, right, top, bottom): the pair (1, 1) extends the + // empty leading dimension. Constant mode has nothing to mirror, so + // it simply fills. + got, err := Pad(empty, []int{0, 0, 1, 1}, "constant", 7) + if err != nil { + t.Fatalf("Pad constant on an empty dimension: %v", err) + } + if got.Shape()[0] != 2 || got.Shape()[1] != 3 { + t.Fatalf("Pad constant shape: %v", got.Shape()) + } + for i := range got.Len() { + if v, _ := FloatAt(got, i/3, i%3); v != 7 { + t.Fatalf("Pad constant [%d]: got %v, want 7", i, v) + } + } +} + +// TestEmptyDimensionReductions pins the up-front extent validation on +// the dim-walking reductions: an empty dimension is an error naming the +// dimension, not a division by zero. +func TestEmptyDimensionReductions(t *testing.T) { + emptyTail := mustFromFloats(t, nil, 2, 0) + if _, err := CumSum(emptyTail, 1); err == nil || !strings.Contains(err.Error(), "empty") { + t.Errorf("CumSum on an empty dimension: %v", err) + } + if _, err := CumProd(emptyTail, 1); err == nil || !strings.Contains(err.Error(), "empty") { + t.Errorf("CumProd on an empty dimension: %v", err) + } + if _, _, err := TopK(emptyTail, 1, 1); err == nil || !strings.Contains(err.Error(), "empty") { + t.Errorf("TopK on an empty dimension: %v", err) + } + emptyHead := mustFromFloats(t, nil, 0, 3) + if _, err := CumSum(emptyHead, 0); err == nil || !strings.Contains(err.Error(), "empty") { + t.Errorf("CumSum on an empty leading dimension: %v", err) + } + // Norm reports an out-of-range dimension instead of panicking in + // reduceShape. + b := mustFromFloats(t, []float64{3, 4}, 2) + if _, err := Norm(b, 2, 2, false); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Errorf("Norm dim 2 of a 1-D array: %v", err) + } + if _, err := Norm(b, 2, -1, false); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Errorf("Norm dim -1: %v", err) + } +} + +// TestRealOnlyEntriesRejectComplex pins the explicit complex rejection +// of the polynomial and binning entries that used to fall into +// floatAt's default branch and panic. +func TestRealOnlyEntriesRejectComplex(t *testing.T) { + c := mustFromComplexes(t, []complex128{complex(1, 1)}, 1) + edges := mustFloats(t, []float64{0, 1, 2}) + calls := []struct { + name string + fn func() (*Array, error) + }{ + {"Hermite", func() (*Array, error) { return Hermite(1, c) }}, + {"Laguerre n=0", func() (*Array, error) { return Laguerre(0, 0, c) }}, + {"Laguerre n=2", func() (*Array, error) { return Laguerre(2, 0.5, c) }}, + {"ChebyshevT", func() (*Array, error) { return ChebyshevT(1, c) }}, + {"ChebyshevU", func() (*Array, error) { return ChebyshevU(1, c) }}, + } + for _, tc := range calls { + if _, err := tc.fn(); err == nil { + t.Errorf("%s: expected a complex rejection", tc.name) + } + } + if _, err := AssignBins(c, edges); err == nil { + t.Error("AssignBins: expected a complex rejection") + } +} + +// TestAxisExtremaNaNNeverWins pins the NaN rule across the axis +// extrema: a line starting with NaN takes its first finite element, the +// arg variants skip the leading NaN index, and TopK never elects a NaN +// candidate. +func TestAxisExtremaNaNNeverWins(t *testing.T) { + a := mustFromFloats(t, []float64{math.NaN(), 5, -1, math.NaN()}, 2, 2) + mn, err := MaxAxis(a, 1) + if err != nil { + t.Fatal(err) + } + if v, _ := FloatAt(mn, 0); v != 5 { + t.Errorf("MaxAxis leading NaN: got %v, want 5", v) + } + if v, _ := FloatAt(mn, 1); v != -1 { + t.Errorf("MaxAxis trailing NaN: got %v, want -1", v) + } + mi, _ := MinAxis(a, 1) + if v, _ := FloatAt(mi, 0); v != 5 { + t.Errorf("MinAxis leading NaN: got %v, want 5", v) + } + if v, _ := FloatAt(mi, 1); v != -1 { + t.Errorf("MinAxis trailing NaN: got %v, want -1", v) + } + am, err := ArgMaxAxis(a, 1) + if err != nil { + t.Fatal(err) + } + for i, w := range []int64{1, 0} { + if v, _ := IntAt(am, i); v != w { + t.Errorf("ArgMaxAxis [%d]: got %d, want %d", i, v, w) + } + } + ai, err := ArgMinAxis(a, 1) + if err != nil { + t.Fatal(err) + } + for i, w := range []int64{1, 0} { + if v, _ := IntAt(ai, i); v != w { + t.Errorf("ArgMinAxis [%d]: got %d, want %d", i, v, w) + } + } + // A line with no finite element has no extreme: an error, mirroring + // the 1-D ArgMax. + allNaN := mustFromFloats(t, []float64{math.NaN(), math.NaN()}, 1, 2) + if _, err := ArgMaxAxis(allNaN, 1); err == nil || !strings.Contains(err.Error(), "every element") { + t.Errorf("ArgMaxAxis all-NaN line: %v", err) + } + // TopK: the NaN at the head of a line never ranks. + vals, idxs, err := TopK(a, 2, 1) + if err != nil { + t.Fatal(err) + } + if v, _ := FloatAt(vals, 0, 0); v != 5 { + t.Errorf("TopK leading NaN vals[0]: got %v, want 5", v) + } + if v, _ := FloatAt(vals, 0, 1); !math.IsNaN(v) { + t.Errorf("TopK leading NaN vals[1]: got %v, want NaN", v) + } + if j, _ := IntAt(idxs, 0, 0); j != 1 { + t.Errorf("TopK leading NaN idxs[0]: got %d, want 1", j) + } +} + +// TestAxisExtremaNaNSplitAcrossWorkers repeats the leading-NaN check +// with several workers, so a line is seeded in one worker and folded in +// another: the merge must not resurrect a NaN. +func TestAxisExtremaNaNSplitAcrossWorkers(t *testing.T) { + prev := engine.SetNumWorkers(4) + defer engine.SetNumWorkers(prev) + vals := make([]float64, 2*64) + for i := range vals { + vals[i] = 1 + } + vals[0] = math.NaN() // row 0, col 0 + vals[40] = 9 // row 0, col 40, on the far side of the worker split + vals[64] = math.NaN() // row 1, col 0 + vals[64+33] = 0.5 // row 1, col 33 + a := mustFromFloats(t, vals, 2, 64) + mx, err := MaxAxis(a, 1) + if err != nil { + t.Fatal(err) + } + if v, _ := FloatAt(mx, 0); v != 9 { + t.Errorf("MaxAxis split line 0: got %v, want 9", v) + } + if v, _ := FloatAt(mx, 1); v != 1 { + t.Errorf("MaxAxis split line 1: got %v, want 1", v) + } + mn, _ := MinAxis(a, 1) + if v, _ := FloatAt(mn, 0); v != 1 { + t.Errorf("MinAxis split line 0: got %v, want 1", v) + } + if v, _ := FloatAt(mn, 1); v != 0.5 { + t.Errorf("MinAxis split line 1: got %v, want 0.5", v) + } +} + +// TestSpMatMulValidatesIndicesAndPromotes pins the SpMatMul contract: +// an out-of-range coordinate is an error from the shared index check, +// and the result dtype follows the promotion ladder instead of pinning +// to the sparse values. +func TestSpMatMulValidatesIndicesAndPromotes(t *testing.T) { + badIndices := mustFromInts(t, []int64{0, 0, 5, 1}, 2, 2) + badValues := mustFloats(t, []float64{1, 2}, 2) + bad := &SparseCOO{Indices: badIndices, Values: badValues, Shape: []int{2, 2}} + dense := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + if _, err := SpMatMul(bad, dense); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Errorf("SpMatMul with an out-of-range index: %v", err) + } + // int sparse times float dense promotes to float, without the old + // truncation through an int cast. + iIndices := mustFromInts(t, []int64{0, 0}, 1, 2) + iValues := mustFromInts(t, []int64{2}, 1) + si := &SparseCOO{Indices: iIndices, Values: iValues, Shape: []int{2, 2}} + frac := mustFromFloats(t, []float64{1.5, 2, 3, 4}, 2, 2) + out, err := SpMatMul(si, frac) + if err != nil { + t.Fatal(err) + } + if out.Dtype() != Float { + t.Errorf("SpMatMul int x float dtype: %s, want float", out.Dtype()) + } + if v, _ := FloatAt(out, 0, 0); v != 3 { + t.Errorf("SpMatMul int x float [0,0]: got %v, want 3", v) + } + // int with int keeps int, like SpMul. + denseI := mustFromInts(t, []int64{3, 4, 5, 6}, 2, 2) + outi, err := SpMatMul(si, denseI) + if err != nil { + t.Fatal(err) + } + if outi.Dtype() != Int { + t.Errorf("SpMatMul int x int dtype: %s, want int", outi.Dtype()) + } + if v, _ := IntAt(outi, 0, 0); v != 6 { + t.Errorf("SpMatMul int x int [0,0]: got %d, want 6", v) + } +} + +// TestTruncatedNormalValidatesShape pins that the shape goes through +// the checked constructor path: an invalid shape yields a nil array +// like the other no-error constructors, never a make panic. +func TestTruncatedNormalValidatesShape(t *testing.T) { + g := NewGenerator(7) + if got := TruncatedNormal(g, []int{-2}, 0, 1); got != nil { + t.Errorf("TruncatedNormal negative dim: got %s, want nil", got) + } + if got := TruncatedNormal(g, nil, 0, 1); got != nil { + t.Errorf("TruncatedNormal nil shape: got %s, want nil", got) + } + w := TruncatedNormal(g, []int{4, 5}, 0, 1) + if w == nil || w.Shape()[0] != 4 || w.Shape()[1] != 5 { + t.Fatalf("TruncatedNormal valid shape: %v", w.Shape()) + } + for _, v := range w.RawFloat32s() { + if v < -2 || v > 2 { + t.Fatalf("TruncatedNormal value %v outside [-2, 2]", v) + } + } +} + +// TestScanZeroTrailingDimension pins the zero trailing dimension on +// the scan kernels: the reduced dimension passes the empty guard, the +// stride folds to zero and the per-line count would divide zero by +// zero, so the scan must return the empty array instead. +func TestScanZeroTrailingDimension(t *testing.T) { + cases := []struct { + name string + shape []int + dim int + }{ + {"CumSum 2x0 dim 0", []int{2, 0}, 0}, + {"CumProd 3x2x0 dim 1", []int{3, 2, 0}, 1}, + } + for _, tc := range cases { + a := mustFromFloats(t, nil, tc.shape...) + var out *Array + var err error + if tc.name[:6] == "CumSum" { + out, err = CumSum(a, tc.dim) + } else { + out, err = CumProd(a, tc.dim) + } + if err != nil { + t.Errorf("%s: %v", tc.name, err) + continue + } + if out.Len() != 0 { + t.Errorf("%s: len %d, want 0", tc.name, out.Len()) + } + } +} + +// TestPadRepeatRejectWrappingArithmetic pins the hostile-arithmetic +// guards: pad and repeat values that wrap the extents must be errors, +// not wrapped shapes handed to alloc. +func TestPadRepeatRejectWrappingArithmetic(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 3, 1) + big := math.MaxInt64/2 + 1 + if _, err := Pad(a, []int{big, big, 0, 0}, "constant", 0); err == nil { + t.Error("Pad accepted pad values that wrap the extent") + } + if _, err := Repeat(a, big, 0); err == nil { + t.Error("Repeat accepted a repeat that wraps the extent") + } + // CopyRows refuses an out-of-range row instead of panicking like + // slice indexing. + if _, err := a.CopyRows([]int{0, 3}); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Errorf("CopyRows with row 3 of 3: %v", err) + } +} diff --git a/internal/core/diff.go b/internal/core/diff.go new file mode 100644 index 0000000..2c525cc --- /dev/null +++ b/internal/core/diff.go @@ -0,0 +1,103 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +// Discrete differences: the workhorse behind finite-difference +// derivatives and signal detrending, along any single axis. + +// Diff takes the successive differences along one axis, order times: +// order 1 is out[i] = a[i+1] − a[i] along the axis, order 2 applies +// it again, and so on. The axis shrinks by order; the axis must +// therefore hold more elements than the order, and axis must name one +// of the array's axes. Complex arrays are fine: differences carry no +// ordering assumption. Int arrays keep their dtype, because the +// difference of two int64 is the int64 difference; everything else +// produces float. +func Diff(a *Array, order, axis int) (*Array, error) { + if order < 1 { + return nil, errf("Diff: the order must be at least 1, got %d", order) + } + if axis < 0 || axis >= a.NDim() { + return nil, errf("Diff: axis %d is outside the %d axes of shape %s", axis, a.NDim(), shapeText(a.Shape())) + } + if a.Shape()[axis] <= order { + return nil, errf("Diff: axis %d holds %d elements, more than the order %d is needed", + axis, a.Shape()[axis], order) + } + cur := a + for range order { + next, err := diffOnce(cur, axis) + if err != nil { + return nil, err + } + cur = next + } + return cur, nil +} + +// diffOnce applies one round of differences along the axis. Int +// differences stay int64 and complex ones stay complex; every other +// dtype produces float, the only route that used to widen the int side +// through float64 and round neighbours above 2^53 together. +// +// The output is walked run by run: for one position along the trailing +// dimensions the two neighbours sit a fixed stride apart, so a run of +// the output is a plain elementwise subtraction with the dtype +// dispatched once, not a per-element coordinate fold. +func diffOnce(a *Array, axis int) (*Array, error) { + if !a.isContiguous() { + // A strided view's payload window is not the run the walk needs, + // so reduce it to a dense copy first; the elements, and with + // them the differences, are the ones the accessors returned. + a = a.materialise() + } + shape := a.Shape() + outShape := append([]int{}, shape...) + outShape[axis]-- + dt := Float + switch a.dt { + case Complex: + dt = Complex + case Int: + dt = Int + } + out := &Array{shape: outShape, dt: dt} + out.alloc(out.Len()) + tail := 1 + for d := axis + 1; d < len(shape); d++ { + tail *= shape[d] + } + head := 1 + for d := range axis { + head *= shape[d] + } + n := shape[axis] + switch dt { + case Int: + diffRuns(out.ints, a.ints[:a.Len()], head, n, tail) + case Complex: + diffRuns(out.complexes, a.complexes[:a.Len()], head, n, tail) + default: + // The source widens exactly as FloatAt widens it; a float64 + // source is read in place. + diffRuns(out.floats, floatPayload(a), head, n, tail) + } + return out, nil +} + +// diffRuns fills dst with the successive differences of src along an +// axis of n elements that steps by tail elements, for each of head +// outer positions. +func diffRuns[T int64 | float64 | complex128](dst, src []T, head, n, tail int) { + for h := range head { + base := h * n * tail + dstBase := h * (n - 1) * tail + for i := range tail { + s, d := base+i, dstBase+i + for k := range n - 1 { + dst[d+k*tail] = src[s+(k+1)*tail] - src[s+k*tail] + } + } + } +} diff --git a/internal/core/dtypes_promote_test.go b/internal/core/dtypes_promote_test.go new file mode 100644 index 0000000..eb7bd4e --- /dev/null +++ b/internal/core/dtypes_promote_test.go @@ -0,0 +1,345 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +// The promotion property pin for the narrow element types. The +// narrow-dtype contract: same dtype promotes to itself; integer-class +// pairs promote +// to the smallest dtype whose value range contains both operands; +// cross-class pairs promote to the higher ladder rung; and the rank +// order and the diagnostic renderings stay exactly as recorded here. + +// promoteDtypes lists the twelve dtypes the promote surface covers. +var promoteDtypes = [...]Dtype{Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int, Float16, Float32, Float, Complex} + +// intClassDtypes lists the integer class in containment order, the +// axis the promotion table is keyed by. +var intClassDtypes = [...]Dtype{Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int} + +// dtypeRange answers the exact value range of an integer-class dtype. +func dtypeRange(d Dtype) (lo, hi int64) { + switch d { + case Bool: + return 0, 1 + case Int8: + return math.MinInt8, math.MaxInt8 + case Uint8: + return 0, math.MaxUint8 + case Int16: + return math.MinInt16, math.MaxInt16 + case Uint16: + return 0, math.MaxUint16 + case Int32: + return math.MinInt32, math.MaxInt32 + case Uint32: + return 0, math.MaxUint32 + default: + return math.MinInt64, math.MaxInt64 + } +} + +// promoteProbeInts builds a one-element Int array carrying v, for the +// Astype boundary probes below. +func promoteProbeInts(t *testing.T, v int64) *Array { + t.Helper() + a, err := FromInts([]int64{v}, 1) + if err != nil { + t.Fatalf("FromInts(%d): %v", v, err) + } + return a +} + +// expectedPromote is the promotion table as the design records it, +// written out independently of intPromote so a change to the +// implementation table has to survive this copy. Same dtype answers +// itself; integer-class pairs take the hand-written containment cell; +// everything else takes the higher dtypeRank rung, ties going to a. +func expectedPromote(a, b Dtype) Dtype { + if a == b { + return a + } + // The integer-class containment matrix, symmetric, rows and + // columns in intClassDtypes order, diagonal cells unused because + // the same-dtype rule above answers them. + var containment = [len(intClassDtypes)][len(intClassDtypes)]Dtype{ + // Bool: contained by every integer dtype. + {Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int}, + // Int8 [-128,127]: with Uint8 needs int16. + {Int8, Int8, Int16, Int16, Int32, Int32, Int, Int}, + // Uint8 [0,255]. + {Uint8, Int16, Uint8, Int16, Uint16, Int32, Uint32, Int}, + // Int16: with Uint16 needs int32; with Uint32 needs int. + {Int16, Int16, Int16, Int16, Int32, Int32, Int, Int}, + // Uint16 [0,65535]. + {Uint16, Int32, Uint16, Int32, Uint16, Int32, Uint32, Int}, + // Int32: with Uint32 needs int. + {Int32, Int32, Int32, Int32, Int32, Int32, Int, Int}, + // Uint32 [0,2^32-1]: contained by int. + {Uint32, Int, Uint32, Int, Uint32, Int, Uint32, Int}, + // Int contains the whole class. + {Int, Int, Int, Int, Int, Int, Int, Int}, + } + ia, ib := -1, -1 + for i, c := range intClassDtypes { + if c == a { + ia = i + } + if c == b { + ib = i + } + } + if ia >= 0 && ib >= 0 { + return containment[ia][ib] + } + if dtypeRank(a) >= dtypeRank(b) { + return a + } + return b +} + +// TestDtypesPromoteSame pins rule (a): every dtype promotes to itself. +func TestDtypesPromoteSame(t *testing.T) { + for _, d := range promoteDtypes { + if got := promote(d, d); got != d { + t.Fatalf("promote(%s, %s) = %s, want %s", d, d, got, d) + } + } +} + +// TestDtypesPromoteTable pins the full 12x12 table against the design's +// independently written copy, and the symmetry the containment rule +// implies. +func TestDtypesPromoteTable(t *testing.T) { + for _, a := range promoteDtypes { + for _, b := range promoteDtypes { + want := expectedPromote(a, b) + if got := promote(a, b); got != want { + t.Fatalf("promote(%s, %s) = %s, want %s", a, b, got, want) + } + if got := promote(b, a); got != want { + t.Fatalf("promote(%s, %s) = %s, want %s (symmetry)", b, a, got, want) + } + } + } +} + +// TestDtypesPromoteNamedPairs pins the specific mixed pairs the design +// names, both operands' orders. +func TestDtypesPromoteNamedPairs(t *testing.T) { + cases := []struct { + a, b, want Dtype + }{ + {Int8, Uint8, Int16}, + {Int16, Uint16, Int32}, + {Int32, Uint32, Int}, + {Uint8, Int16, Int16}, + {Uint16, Int32, Int32}, + {Uint32, Int, Int}, + } + for _, c := range cases { + if got := promote(c.a, c.b); got != c.want { + t.Fatalf("promote(%s, %s) = %s, want %s", c.a, c.b, got, c.want) + } + if got := promote(c.b, c.a); got != c.want { + t.Fatalf("promote(%s, %s) = %s, want %s", c.b, c.a, got, c.want) + } + } +} + +// TestDtypesPromoteCrossClass pins rule (c): a cross-class pair +// promotes to the higher dtypeRank operand, and Bool with any numeric +// answers the numeric. +func TestDtypesPromoteCrossClass(t *testing.T) { + integers := []Dtype{Int8, Uint8, Int16, Uint16, Int32, Uint32, Int} + for _, d := range integers { + if got := promote(d, Float16); got != Float16 { + t.Fatalf("promote(%s, float16) = %s, want float16", d, got) + } + if got := promote(d, Float32); got != Float32 { + t.Fatalf("promote(%s, float32) = %s, want float32", d, got) + } + if got := promote(d, Float); got != Float { + t.Fatalf("promote(%s, float) = %s, want float", d, got) + } + if got := promote(d, Complex); got != Complex { + t.Fatalf("promote(%s, complex) = %s, want complex", d, got) + } + if got := promote(Bool, d); got != d { + t.Fatalf("promote(bool, %s) = %s, want %s", d, got, d) + } + } + for _, d := range []Dtype{Bool, Float16, Float32, Float} { + if got := promote(d, Complex); got != Complex { + t.Fatalf("promote(%s, complex) = %s, want complex", d, got) + } + } + if got := promote(Float16, Float32); got != Float32 { + t.Fatalf("promote(float16, float32) = %s, want float32", got) + } + if got := promote(Float16, Float); got != Float { + t.Fatalf("promote(float16, float) = %s, want float", got) + } + if got := promote(Float32, Float); got != Float { + t.Fatalf("promote(float32, float) = %s, want float", got) + } + if got := promote(Bool, Float16); got != Float16 { + t.Fatalf("promote(bool, float16) = %s, want float16", got) + } +} + +// TestDtypesPromoteIntegerContainment pins rule (b) by boundary +// probing: every integer-class pair promotes to a dtype that contains +// both operands' ranges (each operand's min, max and zero convert into +// it and read back exactly), it is the smallest such dtype (every +// candidate below it in containment order fails at least one operand +// boundary with the loud Astype wording), and a value outside both +// operands' ranges fails the narrower candidates the same way. +func TestDtypesPromoteIntegerContainment(t *testing.T) { + for _, a := range intClassDtypes { + for _, b := range intClassDtypes { + p := promote(a, b) + alo, ahi := dtypeRange(a) + blo, bhi := dtypeRange(b) + // Containment: each operand's boundaries survive Astype + // into the promoted dtype and read back exact. + for _, probe := range []struct { + src Dtype + lo int64 + hi int64 + }{{a, alo, ahi}, {b, blo, bhi}} { + for _, v := range []int64{probe.lo, probe.hi, 0} { + if v < probe.lo || v > probe.hi { + continue // zero only where the operand has it + } + conv, err := Astype(promoteProbeInts(t, v), p) + if err != nil { + t.Fatalf("Astype(%s value %d, %s) refused a boundary of %s: %v", probe.src, v, p, probe.src, err) + } + back, berr := IntAt(conv, 0) + if berr != nil || back != v { + t.Fatalf("Astype(%s value %d, %s) read back %d (err %v), want the exact boundary", probe.src, v, p, back, berr) + } + } + } + // Minimality: every integer-class candidate below the + // promoted dtype in containment order fails at least one + // operand boundary, loudly. + pIdx := -1 + for i, c := range intClassDtypes { + if c == p { + pIdx = i + } + } + if pIdx < 0 { + t.Fatalf("promote(%s, %s) = %s, which is not an integer-class dtype", a, b, p) + } + outside := max(ahi, bhi) + if outside < math.MaxInt64 { + outside++ + } + for i, c := range intClassDtypes { + if i >= pIdx { + continue + } + clo, chi := dtypeRange(c) + fails := false + for _, v := range []int64{alo, ahi, blo, bhi} { + if v < clo || v > chi { + fails = true + if c == Bool { + // Bool has no range to check by design: + // every source reaches it by the test + // against zero. Its non-containment is the + // range fact above, not an Astype refusal. + continue + } + if _, err := Astype(promoteProbeInts(t, v), c); err == nil { + t.Fatalf("Astype(value %d, %s) succeeded although %d lies outside %s's range", v, c, v, c) + } else if !strings.Contains(err.Error(), "does not fit") { + t.Fatalf("Astype(value %d, %s) refused with %q, want the range wording", v, c, err) + } + } + } + if !fails { + t.Fatalf("promote(%s, %s) = %s but the narrower candidate %s contains both ranges too", a, b, p, c) + } + } + // A value beyond both operands fails every candidate below + // the promoted dtype, the loud refusal naming the range. + // Bool is exempt for the reason above. + if outside > ahi && outside > bhi { + for i, c := range intClassDtypes { + if i >= pIdx || c == Bool { + continue + } + if _, err := Astype(promoteProbeInts(t, outside), c); err == nil { + t.Fatalf("Astype(outside value %d, %s) succeeded; want a refusal below the promoted %s", outside, c, p) + } + } + } + } + } +} + +// TestDtypesRankPin pins rule (e) first half: the dtypeRank ordering, +// including float16's slot between int and float32. +func TestDtypesRankPin(t *testing.T) { + want := map[Dtype]int{ + Bool: 0, Int8: 1, Uint8: 1, Int16: 2, Uint16: 2, + Int32: 3, Uint32: 3, Int: 4, Float16: 5, Float32: 6, + Float: 7, Complex: 8, + } + for d, w := range want { + if got := dtypeRank(d); got != w { + t.Fatalf("dtypeRank(%s) = %d, want %d", d, got, w) + } + } + // The ladder order the cross-class promotions walk. + ladder := []Dtype{Bool, Int8, Int16, Int32, Int, Float16, Float32, Float, Complex} + for i := 1; i < len(ladder); i++ { + if dtypeRank(ladder[i-1]) >= dtypeRank(ladder[i]) { + t.Fatalf("dtypeRank(%s) = %d must sit below dtypeRank(%s) = %d", + ladder[i-1], dtypeRank(ladder[i-1]), ladder[i], dtypeRank(ladder[i])) + } + } + // The same-width signedness pairs share a rung. + for _, pair := range [][2]Dtype{{Int8, Uint8}, {Int16, Uint16}, {Int32, Uint32}} { + if dtypeRank(pair[0]) != dtypeRank(pair[1]) { + t.Fatalf("%s and %s must share a rank, got %d and %d", pair[0], pair[1], dtypeRank(pair[0]), dtypeRank(pair[1])) + } + } + // intClass membership: the containment class holds Bool and the + // integers, and no float. + for _, d := range intClassDtypes { + if !intClass(d) { + t.Fatalf("intClass(%s) = false, want true", d) + } + } + for _, d := range []Dtype{Float16, Float32, Float, Complex} { + if intClass(d) { + t.Fatalf("intClass(%s) = true, want false", d) + } + } +} + +// TestDtypesStringPin pins rule (e) second half: the diagnostic +// renderings of all twelve dtypes. +func TestDtypesStringPin(t *testing.T) { + want := map[Dtype]string{ + Bool: "bool", Int8: "int8", Uint8: "uint8", Int16: "int16", + Uint16: "uint16", Int32: "int32", Uint32: "uint32", Int: "int", + Float16: "float16", Float32: "float32", Float: "float", Complex: "complex", + } + for _, d := range promoteDtypes { + if got := d.String(); got != want[d] { + t.Fatalf("Dtype(%d).String() = %q, want %q", d, got, want[d]) + } + } +} diff --git a/internal/core/einsum.go b/internal/core/einsum.go new file mode 100644 index 0000000..cff50f8 --- /dev/null +++ b/internal/core/einsum.go @@ -0,0 +1,1653 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "slices" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Einsum implements Einstein summation for a useful subset of +// operations. The spec is a string of the form +// "lhs[,lhs,...]->rhs". Each label (an ASCII letter a-z or A-Z) names +// a dimension; the same letter across operands must agree. Any other +// character is an error, so an invalid spec can never silently compute. +// +// Supported cases: +// +// "ij,jk->ik" matrix multiply +// "ij,ji->" Frobenius inner product of two matrices +// "ij,ij->" same: sum of element-wise products +// "ij,ij->ij" element-wise product +// "ij->ji" transpose +// "ii->i" diagonal +// "ii->" trace +// "ij->" full sum +// "i->i" identity (returns the operand) +// +// Other patterns, among them broadcasting over the ellipsis, +// reduction-only axes (e.g. "ij->i") and batched matmul +// ("bij,bjk->bik"), are handled by the general engine behind the same +// entry point; an ellipsis must lead the subscripts it appears in. +func Einsum(spec string, operands ...*Array) (*Array, error) { + for _, op := range operands { + // The einsum engines (table, batched and general) are built on + // the matmul kernels and are not offered for the half dtype; + // the refusal is loud and named rather than a silent payload + // misread, and the Astype conversion is cheap. The narrow + // element types ride the same refusal until their kernels + // arrive. + if op.Dtype() == Float16 { + return nil, errf("Einsum: float16 operands are not supported; convert with Astype") + } + if narrowRefused(op.Dtype()) { + return nil, errf("Einsum: dtype %s is not supported; convert with Astype", op.Dtype()) + } + } + lhsPart, rhsPart, hasArrow := strings.Cut(spec, "->") + lhsStrs := strings.Split(lhsPart, ",") + if len(lhsStrs) != len(operands) { + return nil, errf("Einsum: spec %q has %d operands, got %d", spec, len(lhsStrs), len(operands)) + } + // The general engine owns the patterns the table below does not + // table: ellipsis broadcasting, implicit output, arbitrary sums. + if !hasArrow || strings.Contains(spec, ".") { + return einsumGeneral(lhsStrs, rhsPart, hasArrow, operands) + } + lhs := make([][]int, len(operands)) + for i, s := range lhsStrs { + l, lerr := einsumLabels(s) + if lerr != nil { + return nil, lerr + } + lhs[i] = l + } + rhs, rerr := einsumLabels(rhsPart) + if rerr != nil { + return nil, rerr + } + out, err := einsumEval(lhs, rhs, operands) + if err == nil { + return out, nil + } + // The table declined; the general engine gets the last word so a + // valid pattern never errors twice. Its error is the one reported: + // it saw the whole pattern, while the table can only say that no + // kernel matched. A shape that would previously print a label table + // now names the label that disagrees. + if g, gerr := einsumGeneral(lhsStrs, rhsPart, true, operands); gerr == nil { + return g, nil + } else { + return nil, gerr + } +} + +// einsumLabels converts a label string like "ij" into a slice of int +// labels (a=0, b=1, …, z=25, A=26, …, Z=51). Empty string returns nil. +// Every character must be an ASCII letter, matching einsumParse in the +// general engine: digits and anything else are a named error quoting +// the offending character, never a silent misread. +func einsumLabels(s string) ([]int, error) { + if s == "" { + return nil, nil + } + out := make([]int, 0, len(s)) + for _, c := range s { + switch { + case c >= 'a' && c <= 'z': + out = append(out, int(c-'a')) + case c >= 'A' && c <= 'Z': + out = append(out, int(c-'A')+26) + default: + return nil, errf("Einsum: bad label %q in %q", string(c), s) + } + } + return out, nil +} + +// einsumEval is the dispatcher. It inspects the spec and routes to +// one of the concrete implementations. +func einsumEval(lhs [][]int, rhs []int, operands []*Array) (*Array, error) { + // An output label may appear once: "ii->ii" is a typo for the + // diagonal "ii->i", and the identity fast path below would silently + // return a copy of the whole operand instead. The general engine + // refuses a repeated output label as well. + for i, l := range rhs { + if slices.Contains(rhs[:i], l) { + return nil, errf("Einsum: output label %q repeats", einsumLabelName(l)) + } + } + // Validate operand shapes against labels. + for i, op := range operands { + if op.NDim() != len(lhs[i]) { + return nil, errf("Einsum: operand %d has shape %s but spec has %d labels", i, shapeText(op.shape), len(lhs[i])) + } + } + // Trace "ii->": sum of diagonal. Routed through Diagonal + Sum so the + // scalar keeps the operand's dtype (int stays int, exact above 2^53) + // and complex operands work, exactly as the general engine would + // answer them. + if len(lhs) == 1 && len(lhs[0]) == 2 && lhs[0][0] == lhs[0][1] && len(rhs) == 0 { + op := operands[0] + if op.shape[0] != op.shape[1] { + return nil, errf("Einsum: label %q has sizes %d and %d in operand 0", + einsumLabelName(lhs[0][0]), op.shape[0], op.shape[1]) + } + diag, derr := Diagonal(op, 0) + if derr != nil { + return nil, derr + } + return scalarArray(Sum(diag)) + } + // Diagonal "ii->i": extract main diagonal as 1-D. + if len(lhs) == 1 && len(lhs[0]) == 2 && lhs[0][0] == lhs[0][1] && + len(rhs) == 1 && rhs[0] == lhs[0][0] { + return Diagonal(operands[0], 0) + } + // Identity "xn->xn": a copy, not the operand; results never alias + // their inputs. + if len(lhs) == 1 && einsumEqLabels(lhs[0], rhs) { + return Copy(operands[0]), nil + } + // Full sum "...": "xN,..->" + if len(rhs) == 0 { + return einsumReduceAll(lhs, operands) + } + // Transpose "ij->ji" or other permutation. + if len(lhs) == 1 && len(rhs) == len(lhs[0]) && einsumIsPermutation(lhs[0], rhs) { + return einsumTranspose(operands[0], lhs[0], rhs) + } + // Reduction-only axes "ij->i", "ijk->ik": the axes the output drops + // are summed away, the surviving ones keep their order. + if len(lhs) == 1 && len(rhs) > 0 && len(rhs) < len(lhs[0]) { + if res, ok, err := einsumSumAxes(lhs[0], rhs, operands[0]); ok { + return res, err + } + } + // Matrix multiply "ij,jk->ik" (or equivalent with letters). + if len(lhs) == 2 && len(rhs) == 2 && + len(lhs[0]) == 2 && len(lhs[1]) == 2 && + lhs[0][1] == lhs[1][0] && lhs[0][0] == rhs[0] && lhs[1][1] == rhs[1] { + return MatMul2D(operands[0], operands[1]) + } + // Matrix-vector "ij,j->i". + if len(lhs) == 2 && len(rhs) == 1 && + len(lhs[0]) == 2 && len(lhs[1]) == 1 && + lhs[0][1] == lhs[1][0] && lhs[0][0] == rhs[0] { + return MatMul2D(operands[0], operands[1]) + } + // Vector-matrix "i,ij->j". + if len(lhs) == 2 && len(rhs) == 1 && + len(lhs[0]) == 1 && len(lhs[1]) == 2 && + lhs[0][0] == lhs[1][0] && lhs[1][1] == rhs[0] { + return MatMul2D(operands[0], operands[1]) + } + // Batched products "bij,bjk->bik", "bij,bj->bi" and "bi,bij->bj": + // the shared leading label multiplies matching contiguous batch + // slices into one shared payload. For a fixed output slot the + // walk's summed label visits ascending, which is MatMul2D's + // ascending-p accumulation per cell, so each batch is the walk's + // slot sum (the Float32 kernel keeps its documented narrow-once + // semantics, as on the 2-D path above). Only patterns with + // pairwise-distinct labels and exactly agreeing shared dimensions + // are taken; broadcast batches and repeated labels stay with the + // engine. + if res, ok, err := einsumBatched(lhs, rhs, operands); ok { + return res, err + } + // Dot product "i,i->". + if len(lhs) == 2 && len(rhs) == 0 && + len(lhs[0]) == 1 && len(lhs[1]) == 1 && lhs[0][0] == lhs[1][0] { + s, err := Dot(operands[0], operands[1]) + if err != nil { + return nil, err + } + return scalarArray(s) + } + // Outer product "i,j->ij". + if len(lhs) == 2 && len(rhs) == 2 && + len(lhs[0]) == 1 && len(lhs[1]) == 1 && lhs[0][0] != lhs[1][0] { + return einsumOuter(operands[0], operands[1], rhs, lhs) + } + // Element-wise product with output "ij,ij->ij". + if len(lhs) == 2 && len(rhs) == 2 && + einsumEqLabels(lhs[0], lhs[1]) && einsumEqLabels(lhs[0], rhs) { + return Mul(operands[0], operands[1]) + } + // Element-wise multiply, reduce-all "ij,ij->". + if len(lhs) == 2 && len(rhs) == 0 && + einsumEqLabels(lhs[0], lhs[1]) { + prod, err := Mul(operands[0], operands[1]) + if err != nil { + return nil, err + } + return scalarArray(Sum(prod)) + } + return nil, errf("Einsum: unsupported pattern lhs=%v rhs=%v", lhs, rhs) +} + +// einsumSumAxes folds away the axes an explicit output leaves out of a +// single operand, the patterns "ij->i" and "kji->k" reduce to. It +// declines, reporting ok=false, unless the axis fold reproduces the +// engine's walk bit for bit: +// +// - the output must read as the operand's labels with some of them +// dropped, in order: a reordered output or a label the operand does +// not have is not this pattern, and a repeated label is a diagonal +// the fold would not take; +// - the operand must be contiguous, as every kernel reading a payload +// window demands; +// - a float32 operand stays with the engine, which narrows into the +// float32 slot once per addend while the fold widens, sums in +// float64 and narrows once at the end; +// - a complex operand stays with the engine, whose product seed +// multiplies the first operand by 1+0i: that is the identity for +// finite values but turns an infinite component into a NaN. +// +// The fold runs the dropped axes from the largest label id down, which +// is what makes its nesting the engine's: the engine walks the sorted +// label union with the smallest id outermost, so the innermost loop +// carries the largest id and has to be closed first. +func einsumSumAxes(labels, rhs []int, a *Array) (*Array, bool, error) { + if a.strides != nil || a.dt == Float32 || a.dt == Complex { + return nil, false, nil + } + for i, l := range labels { + if slices.Contains(labels[:i], l) { + return nil, false, nil + } + } + // Which axes the output skips: every output label must be met in + // the operand's own order. + keep := make([]bool, len(labels)) + j := 0 + for _, r := range rhs { + for j < len(labels) && labels[j] != r { + j++ + } + if j == len(labels) { + return nil, false, nil + } + keep[j] = true + j++ + } + // The dropped axes as (position, label); a fold removes one axis and + // shifts the positions after it down by one. + rest := make([]einsumAxis, 0, len(labels)-len(rhs)) + for axis, l := range labels { + if !keep[axis] { + rest = append(rest, einsumAxis{pos: axis, label: l}) + } + } + cur := a + for len(rest) > 0 { + last := 0 + for i, r := range rest { + if r.label > rest[last].label { + last = i + } + } + next, err := SumAxis(cur, rest[last].pos) + if err != nil { + return nil, true, err + } + cur = next + removed := rest[last].pos + rest = slices.Delete(rest, last, last+1) + for i := range rest { + if rest[i].pos > removed { + rest[i].pos-- + } + } + } + return cur, true, nil +} + +// einsumAxis is one axis of an operand on its way to being folded: its +// current position and the label it carries. +type einsumAxis struct { + pos int + label int +} + +// einsumBatched recognises the batched product patterns whose shared +// leading label multiplies matching batch slices into one payload: +// "bij,bjk->bik", "bij,bj->bi" and "bi,bij->bj". It declines, +// reporting ok=false, unless every label in the pattern is pairwise +// distinct and the shared dimensions agree exactly; broadcast batches +// and repeated labels keep the general engine's semantics. ok=true +// means the returned result, or error, is final. +func einsumBatched(lhs [][]int, rhs []int, operands []*Array) (*Array, bool, error) { + if len(lhs) != 2 { + return nil, false, nil + } + x, y := lhs[0], lhs[1] + sa, sb := operands[0].Shape(), operands[1].Shape() + switch { + // Matrix times matrix: a (b,i,j), b (b,j,k), out (b,i,k). + case len(x) == 3 && len(y) == 3 && len(rhs) == 3 && + x[0] == y[0] && x[0] == rhs[0] && x[1] == rhs[1] && y[2] == rhs[2] && x[2] == y[1] && + x[0] != x[1] && x[0] != x[2] && x[0] != y[2] && + x[1] != x[2] && x[1] != y[2] && x[2] != y[2] && + sa[0] == sb[0] && sa[2] == sb[1]: + // Matrix times vector: a (b,i,j), b (b,j), out (b,i). + case len(x) == 3 && len(y) == 2 && len(rhs) == 2 && + x[0] == y[0] && x[0] == rhs[0] && x[1] == rhs[1] && x[2] == y[1] && + x[0] != x[1] && x[0] != x[2] && x[1] != x[2] && + sa[0] == sb[0] && sa[2] == sb[1]: + // Vector times matrix: a (b,i), b (b,i,j), out (b,j). + case len(x) == 2 && len(y) == 3 && len(rhs) == 2 && + x[0] == y[0] && x[0] == rhs[0] && x[1] == y[1] && y[2] == rhs[1] && + x[0] != x[1] && x[0] != y[2] && x[1] != y[2] && + sa[0] == sb[0] && sa[1] == sb[1]: + default: + return nil, false, nil + } + res, err := einsumBatchedProduct(operands[0], operands[1]) + if err != nil { + return nil, true, err + } + return res, true, nil +} + +// einsumBatchedProduct multiplies matching batches of a and b straight +// into one shared output payload. Both operands are materialised +// first: the per-batch slices must be contiguous payload windows. Each +// batch runs the panel kernels MatMul2D runs for its shape over the +// same ascending-p walk, so the payload holds exactly the stacked +// MatMul2D results; assembling through the kernels rather than through +// MatMul2D only removes the per-batch output array and the per-batch +// goroutine fan-out. Mixed real operands convert once for the whole +// tensor, which reads back the values MatMul2D's per-batch conversions +// produced. +func einsumBatchedProduct(a, b *Array) (*Array, error) { + a = a.materialise() + b = b.materialise() + sa, sb := a.shape, b.shape + batches := sa[0] + // Per-batch result shape: a's batch dims minus the summed axis, + // then b's trailing dims. + outTail := append(append([]int{}, sa[1:len(sa)-1]...), sb[2:]...) + tailTotal := 1 + for _, d := range outTail { + tailTotal *= d + } + out := &Array{shape: append([]int{batches}, outTail...), dt: promote(a.dt, b.dt)} + out.alloc(batches * tailTotal) + if batches == 0 { + return out, nil + } + // One batch's geometry, named as MatMul2D names its dimensions: + // "bij,bjk->bik" multiplies an n×k batch by a k×m batch, + // "bij,bj->bi" an n×k batch by a k batch, and "bi,bij->bj" a k + // batch by a k×m batch. nOut is the batch's output row count. + bmm, bmv, vbm := false, false, false + var n, k, m int + switch { + case len(sa) == 3 && len(sb) == 3: + bmm = true + n, k, m = sa[1], sa[2], sb[2] + case len(sa) == 3: + bmv = true + n, k = sa[1], sa[2] + default: + vbm = true + k, m = sb[1], sb[2] + } + nOut := n + if vbm { + nOut = m + } + perA := a.Len() / batches + perB := b.Len() / batches + // Workers own disjoint output row blocks, so the kernels need no + // locks. Two layouts: a worker takes each batch whole, running the + // batch's kernel alone where the wide float64 panels win, or the + // batches' output rows share one global split over the full pool + // with narrow panels. Measured on the batched benchmarks, whole + // batches win from two batches up: each kernel then streams its b + // once per wide panel on a private core, instead of narrow panels + // sharing L1 under SMT and paying a goroutine fan-out per split. + // A lone batch cannot feed the pool with whole batches, so it + // keeps MatMul2D's own row split. + whole := batches > 1 + drive := func(body func(walk func(visit func(bt, r0, r1 int)))) { + if whole { + engine.Parallel(batches, func(ks, ke int) { + body(func(visit func(bt, r0, r1 int)) { + for bt := ks; bt < ke; bt++ { + visit(bt, 0, nOut) + } + }) + }) + return + } + engine.Parallel(batches*nOut, func(rs, re int) { + body(func(visit func(bt, r0, r1 int)) { + for bt := rs / nOut; bt*nOut < re; bt++ { + r0 := max(rs, bt*nOut) - bt*nOut + r1 := min(re, (bt+1)*nOut) - bt*nOut + visit(bt, r0, r1) + } + }) + }) + } + switch out.dt { + case Int: + ai, bi, oi := a.ints, b.ints, out.ints + drive(func(walk func(visit func(bt, r0, r1 int))) { + walk(func(bt, r0, r1 int) { + aK := ai[bt*perA : bt*perA+perA] + bK := bi[bt*perB : bt*perB+perB] + oK := oi[bt*tailTotal : bt*tailTotal+tailTotal] + switch { + case bmm: + matMulIntRows(aK, bK, oK, r0, r1, k, m) + case bmv: + matVecRows(aK, bK, oK, r0, r1, k) + default: + vecMatCols(aK, bK, oK[r0:r1], m, r0, r1) + } + }) + }) + case Float32: + // Accumulate in float64 and round once, exactly as + // MatMul2D's Float32 path: the payloads read directly, every + // widening is exact, and each finished row narrows once. The + // pooled scratch serves the matmul panels (four m-wide rows) + // or the vecMat column walk (one accumulator row); each + // kernel clears the scratch it uses, so one buffer outlasts + // the worker's whole chunk. + aIs32, bIs32 := a.dt == Float32, b.dt == Float32 + drive(func(walk func(visit func(bt, r0, r1 int))) { + var scratch []float64 + switch { + case bmm: + scratch = engine.GetFloat64Buf(4 * m) + case vbm: + scratch = engine.GetFloat64Buf(m) + } + if scratch != nil { + defer engine.PutFloat64Buf(scratch) + } + walk(func(bt, r0, r1 int) { + oK := out.floats32[bt*tailTotal : bt*tailTotal+tailTotal] + var aK32, bK32 []float32 + var aKi, bKi []int64 + if aIs32 { + aK32 = a.floats32[bt*perA : bt*perA+perA] + } else { + aKi = a.ints[bt*perA : bt*perA+perA] + } + if bIs32 { + bK32 = b.floats32[bt*perB : bt*perB+perB] + } else { + bKi = b.ints[bt*perB : bt*perB+perB] + } + switch { + case bmm: + switch { + case aIs32 && bIs32: + matMulF32Rows(aK32, bK32, oK, scratch, r0, r1, k, m) + case aIs32: + matMulF32Rows(aK32, bKi, oK, scratch, r0, r1, k, m) + case bIs32: + matMulF32Rows(aKi, bK32, oK, scratch, r0, r1, k, m) + default: + matMulF32Rows(aKi, bKi, oK, scratch, r0, r1, k, m) + } + case bmv: + switch { + case aIs32 && bIs32: + matVecF32Rows(aK32, bK32, oK, r0, r1, k) + case aIs32: + matVecF32Rows(aK32, bKi, oK, r0, r1, k) + case bIs32: + matVecF32Rows(aKi, bK32, oK, r0, r1, k) + default: + matVecF32Rows(aKi, bKi, oK, r0, r1, k) + } + default: + switch { + // The column walk narrows over the accumulator's + // own length, so the scratch must be exactly the + // window's width, as MatMul2D sizes it. + case aIs32 && bIs32: + vecMatF32Cols(aK32, bK32, oK[r0:r1], scratch[:r1-r0], m, r0, r1) + case aIs32: + vecMatF32Cols(aK32, bKi, oK[r0:r1], scratch[:r1-r0], m, r0, r1) + case bIs32: + vecMatF32Cols(aKi, bK32, oK[r0:r1], scratch[:r1-r0], m, r0, r1) + default: + vecMatF32Cols(aKi, bKi, oK[r0:r1], scratch[:r1-r0], m, r0, r1) + } + } + }) + }) + case Float: + // Mixed real operands convert once at entry, in the same + // element order MatMul2D's per-batch conversions read, so the + // kernel streams plain rows; a float64 operand's payload + // already is that stream and skips the copy. + af := a.floats + if a.dt != Float { + af = denseFloatsLocal(a, a.Len(), 1) + } + bf := b.floats + if b.dt != Float { + bf = denseFloatsLocal(b, b.Len(), 1) + } + of := out.floats + // A lone worker runs the wide 8-row panel; a split range runs + // the narrow one (see matMulF64Rows). The whole-batch layout + // runs each kernel alone; the row split mirrors MatMul2D's + // worker test for its rows. + wide := whole || workersFor(nOut) == 1 + drive(func(walk func(visit func(bt, r0, r1 int))) { + walk(func(bt, r0, r1 int) { + aK := af[bt*perA : bt*perA+perA] + bK := bf[bt*perB : bt*perB+perB] + oK := of[bt*tailTotal : bt*tailTotal+tailTotal] + switch { + case bmm: + matMulF64Rows(aK, bK, oK, r0, r1, k, m, wide) + case bmv: + matVecF64Rows(aK, bK, oK, r0, r1, k) + default: + vecMatF64Cols(aK, bK, oK[r0:r1], m, r0, r1) + } + }) + }) + default: + ac := complexPayload(a) + bc := complexPayload(b) + oc := out.complexes + drive(func(walk func(visit func(bt, r0, r1 int))) { + walk(func(bt, r0, r1 int) { + aK := ac[bt*perA : bt*perA+perA] + bK := bc[bt*perB : bt*perB+perB] + oK := oc[bt*tailTotal : bt*tailTotal+tailTotal] + switch { + case bmm: + // The plain i-p-j walk MatMul2D runs, bounded to + // this batch's rows. + for i := r0; i < r1; i++ { + orow := oK[i*m : (i+1)*m] + for p := range k { + av := aK[i*k+p] + for j := range m { + orow[j] += av * bK[p*m+j] + } + } + } + case bmv: + matVecRows(aK, bK, oK, r0, r1, k) + default: + vecMatCols(aK, bK, oK[r0:r1], m, r0, r1) + } + }) + }) + } + return out, nil +} + +// scalarArray wraps a Scalar as a length-1 array of its own dtype, so +// reduced einsum results keep the complex or int payload instead of +// collapsing to the real part. +func scalarArray(s Scalar) (*Array, error) { + switch { + case s.IsComplex(): + return FromComplexes([]complex128{s.Complex()}, 1) + case s.IsFloat(): + return FromFloats([]float64{s.Float()}, 1) + default: + return FromInts([]int64{s.Int()}, 1) + } +} + +func einsumEqLabels(a, b []int) bool { + return slices.Equal(a, b) +} + +// einsumIsPermutation reports whether rhs is a permutation of lhs. A +// label repeated on either side stays a permutation of itself, exactly +// as the membership test it replaces answered. +func einsumIsPermutation(lhs, rhs []int) bool { + if len(lhs) != len(rhs) { + return false + } + for _, r := range rhs { + if !slices.Contains(lhs, r) { + return false + } + } + return true +} + +// einsumTranspose reorders the axes of a according to the mapping +// lhs to rhs. lhs is the current order; rhs is the desired order. +func einsumTranspose(a *Array, lhs, rhs []int) (*Array, error) { + dims := make([]int, len(rhs)) + for i, r := range rhs { + // Find r in lhs. + j := -1 + for k, l := range lhs { + if l == r { + j = k + break + } + } + if j < 0 { + return nil, errf("Einsum: rhs label %d not in lhs", r) + } + dims[i] = j + } + return TransposeAxes(a, dims...) +} + +// einsumReduceAll handles patterns that reduce every axis (no +// "->rhs"). For 1 operand it's a full sum; for 2 operands it's a +// dot/inner product over matching axes: the second operand is first +// permuted into the first's label order, so "ij,ji->" pairs a's +// columns with b's rows instead of multiplying positionally. +func einsumReduceAll(lhs [][]int, operands []*Array) (*Array, error) { + if len(operands) == 1 { + return scalarArray(Sum(operands[0])) + } + if len(operands) == 2 { + b, err := einsumAlign(operands[1], lhs[1], lhs[0]) + if err != nil { + return nil, err + } + prod, err := Mul(operands[0], b) + if err != nil { + return nil, err + } + return scalarArray(Sum(prod)) + } + return nil, errf("Einsum: '->' with %d operands not supported", len(operands)) +} + +// einsumAlign permutes a so its labels read in the target order. Labels +// must appear exactly once in the operand; a already-aligned operand is +// returned unchanged, never a copy. +func einsumAlign(a *Array, labels, target []int) (*Array, error) { + for i, l := range labels { + if slices.Contains(labels[:i], l) { + return nil, errf("Einsum: repeated label in operand") + } + } + dims := make([]int, len(target)) + for i, want := range target { + found := false + for k, l := range labels { + if l == want { + dims[i] = k + found = true + break + } + } + if !found { + return nil, errf("Einsum: label missing from operand, cannot align") + } + } + return TransposeAxes(a, dims...) +} + +// einsumOuter computes the outer product of two 1-D vectors, producing +// a 2-D matrix indexed by the rhs labels. The result dtype follows the +// promotion ladder, so int vectors stay int and complex vectors stay +// complex. +func einsumOuter(a, b *Array, rhs []int, lhs [][]int) (*Array, error) { + if a.NDim() != 1 || b.NDim() != 1 { + return nil, errf("Einsum: outer product needs two 1-D operands, got %d and %d dims", a.NDim(), b.NDim()) + } + // Output shape: in rhs order, each label picks the matching input + // size. + outShape := make([]int, len(rhs)) + for i, label := range rhs { + switch label { + case lhs[0][0]: + outShape[i] = a.Len() + case lhs[1][0]: + outShape[i] = b.Len() + default: + return nil, errf("Einsum: outer label %d not in operands", label) + } + } + out := &Array{shape: outShape, dt: promote(a.dt, b.dt)} + total := 1 + for _, d := range outShape { + total *= d + } + out.alloc(total) + // The first rhs slot is the row (slow) axis. When the rhs order + // swaps the operand labels ("i,j->ji"), the row index selects b's + // elements and the column index a's, so the flat write must follow + // the rhs order, not the operand order. + aIsRow := rhs[0] == lhs[0][0] + for i := range a.Len() { + for j := range b.Len() { + var flat int + if aIsRow { + flat = i*outShape[1] + j + } else { + flat = j*outShape[1] + i + } + switch out.dt { + case Int: + out.ints[flat] = a.ints[i] * b.ints[j] + case Complex: + out.complexes[flat] = a.complexAt(i) * b.complexAt(j) + default: + out.setFromValue(flat, a.floatAt(i)*b.floatAt(j)) + } + } + } + return out, nil +} + +// einsumLabelName spells a label the way the spec writes it: the lower +// case letters first, then the upper case ones. +func einsumLabelName(l int) string { + if l >= 26 { + return string(rune('A' + l - 26)) + } + return string(rune('a' + l)) +} + +// einsumParse splits one operand's label string into explicit labels +// and an ellipsis flag. The ellipsis may appear once and must lead the +// subscripts, because the engine lays it onto the leading axes. Label +// ids run 0-25 for a-z and 26-51 for A-Z. +func einsumParse(s string) ([]int, bool, error) { + var labels []int + ellipsis := false + afterLabel := false + for _, c := range s { + switch { + case c == '.': + // Ellipsis dots arrive as three consecutive runes. + return nil, false, errf("Einsum: stray '.' in %q", s) + case c == '…': + if ellipsis { + return nil, false, errf("Einsum: two ellipses in %q", s) + } + if afterLabel { + // The engine lays the ellipsis axes onto the leading + // dimensions, so subscripts written after a label would + // be read as if they came first. Refuse rather than + // compute a different contraction than the one written. + return nil, false, errf("Einsum: the ellipsis must lead the subscripts, %q has a label before it", s) + } + ellipsis = true + case c >= 'a' && c <= 'z': + labels = append(labels, int(c-'a')) + afterLabel = true + case c >= 'A' && c <= 'Z': + labels = append(labels, int(c-'A')+26) + afterLabel = true + default: + return nil, false, errf("Einsum: bad label %q in %q", string(c), s) + } + } + return labels, ellipsis, nil +} + +// einsumGeneral is the full Einstein-summation engine the pattern +// table above does not reach: ellipsis broadcasting, repeated labels +// (diagonals), reduction-only axes, arbitrary output order, implicit +// output, any operand count. +// +// Evaluation runs over label-position tables built once per call. The +// label union is sorted once; each operand carries, per sorted label +// position, the payload stride its dimensions contribute when that +// label's coordinate advances (repeated labels add, broadcast axes +// contribute zero). Output slots are written in row-major order by an +// odometer over the output labels, and within one slot the summed +// labels run as an inner odometer in ascending sorted order, so every +// slot receives exactly the products, in exactly the summation order, +// a full walk over the sorted label union would produce. +func einsumGeneral(specs []string, rhsSpec string, rhsExplicit bool, operands []*Array) (*Array, error) { + // The three-dot ellipsis becomes its single-rune form for parsing. + for i := range specs { + specs[i] = strings.ReplaceAll(specs[i], "...", "\u2026") + } + rhsSpec = strings.ReplaceAll(rhsSpec, "...", "\u2026") + type opInfo struct { + explicit []int + ellipsis bool + // label of each dimension after ellipsis expansion + dimLabels []int + // stride per dimension, zero on broadcast (size-1) axes + dimStrides []int + } + ops := make([]opInfo, len(operands)) + // Synthetic labels for the ellipsis axes start past the alphabet. + nextSynthetic := 100 + var allLabels []int + labelSeen := make(map[int]bool) + addLabel := func(l int) { + if !labelSeen[l] { + labelSeen[l] = true + allLabels = append(allLabels, l) + } + } + for p, s := range specs { + labels, ell, err := einsumParse(s) + if err != nil { + return nil, err + } + ops[p].explicit = labels + ops[p].ellipsis = ell + nd := operands[p].NDim() + ellDims := 0 + if ell { + ellDims = nd - len(labels) + if ellDims < 0 { + return nil, errf("Einsum: operand %d has %d dims but spec %q wants more", p, nd, s) + } + } else if nd != len(labels) { + return nil, errf("Einsum: operand %d has shape %s but spec %q has %d labels", + p, shapeText(operands[p].shape), s, len(labels)) + } + if ell { + for range ellDims { + addLabel(nextSynthetic) + nextSynthetic++ + } + } + for _, l := range labels { + addLabel(l) + } + } + // Label sizes: explicit labels must agree exactly; an ellipsis axis + // broadcasts the same way, where only a size-1 axis stretches to the + // common size. The slot walk below expresses replication through a + // zero stride, and a zero stride is exactly what a size-1 axis + // contributes: a longer axis that disagrees with the broadcast + // extent has no stride that could read it, so it is refused rather + // than walked past the operand's elements. + sizes := make(map[int]int) + for p := range ops { + nd := operands[p].NDim() + ellDims := nd - len(ops[p].explicit) + dl := make([]int, 0, nd) + if ops[p].ellipsis { + for k := range ellDims { + dl = append(dl, 100+(nextSynthetic-100-ellDims+k)) + } + } + dl = append(dl, ops[p].explicit...) + ops[p].dimLabels = dl + } + // The synthetic ids above were allocated per operand in order; two + // operands' ellipses must line up right-aligned, so renumber by + // distance from the right edge instead. + for p := range ops { + nd := operands[p].NDim() + ellDims := nd - len(ops[p].explicit) + if !ops[p].ellipsis { + continue + } + base := 100 + for k := range ellDims { + // axis offset from the right among the ellipsis dims + ops[p].dimLabels[k] = base + (ellDims - 1 - k) + } + } + // Recompute the label union with the renumbered synthetic ids. + labelSeen = make(map[int]bool) + allLabels = allLabels[:0] + for p := range ops { + for _, l := range ops[p].dimLabels { + if !labelSeen[l] { + labelSeen[l] = true + allLabels = append(allLabels, l) + } + } + } + for l := range labelSeen { + sizes[l] = -1 + } + for p := range ops { + for d, l := range ops[p].dimLabels { + dim := operands[p].shape[d] + cur := sizes[l] + if l >= 100 { + // A synthetic label is an ellipsis axis, and those + // broadcast: the widest operand sets the extent. + if dim > cur { + sizes[l] = dim + } + continue + } + // A written label must have the same length everywhere it + // appears, size 1 included: a lone 1 stretched against a + // longer labelled axis is refused, because a 1 in the wrong + // place is a shape typo far more often than an intended + // broadcast, and the package's own contract is that a + // shared letter means a shared length. Only ellipsis axes + // broadcast. + switch { + case cur == -1: + sizes[l] = dim + case cur != dim: + return nil, errf("Einsum: label %q is %d in one operand and %d in another", + einsumLabelName(l), cur, dim) + } + } + } + for _, s := range sizes { + if s == -1 { + return nil, errf("Einsum: unsized label") + } + } + // Every ellipsis axis must be 1 or the broadcast extent: the widest + // operand sets the size, and a size-1 axis beside it replicates. The + // check runs before any stride or slot walk, so a mismatched axis + // can never reach the readers, not even through a Slice view whose + // payload runs past Len(). + for p := range ops { + for d, l := range ops[p].dimLabels { + if l < 100 { + continue + } + dim := operands[p].shape[d] + if dim == 1 || dim == sizes[l] { + continue + } + return nil, errf("Einsum: operand %d has shape %s, whose ellipsis axis of %d cannot broadcast against the extent %d; only a size-1 axis broadcasts", + p, shapeText(operands[p].shape), dim, sizes[l]) + } + } + // Output labels. + var outLabels []int + if !rhsExplicit { + // Implicit output: labels appearing exactly once, sorted, + // ellipsis axes prepended. + counts := make(map[int]int) + for p := range ops { + for _, l := range ops[p].dimLabels { + counts[l]++ + } + } + var once []int + for l, c := range counts { + if c == 1 && l < 100 { + once = append(once, l) + } + } + slices.Sort(once) + outLabels = append(outLabels, once...) + // Ellipsis dims go first in their broadcast order, leftmost + // operand axis first. The synthetic ids count from the right + // edge, so the operand's left-to-right order is descending id. + var ellAxes []int + for l := range sizes { + if l >= 100 { + ellAxes = append(ellAxes, l) + } + } + slices.Sort(ellAxes) + slices.Reverse(ellAxes) + outLabels = append(ellAxes, outLabels...) + } else { + rhsLabels, rhsEll, err := einsumParse(rhsSpec) + if err != nil { + return nil, err + } + seen := make(map[int]bool) + for _, l := range rhsLabels { + if seen[l] { + return nil, errf("Einsum: output label %q repeats", einsumLabelName(l)) + } + seen[l] = true + if !labelSeen[l] { + return nil, errf("Einsum: output label %q not in any operand", einsumLabelName(l)) + } + outLabels = append(outLabels, l) + } + if rhsEll { + var ellAxes []int + for l := range sizes { + if l >= 100 { + ellAxes = append(ellAxes, l) + } + } + slices.Sort(ellAxes) + slices.Reverse(ellAxes) + outLabels = slices.Concat(ellAxes, outLabels) + } + } + // Strides per operand dimension (0 on broadcast axes) and the + // output strides. + for p := range ops { + a := operands[p] + st := make([]int, a.NDim()) + run := 1 + for d := a.NDim() - 1; d >= 0; d-- { + st[d] = run + run *= a.shape[d] + } + for d, l := range ops[p].dimLabels { + if a.shape[d] == 1 && sizes[l] != 1 { + st[d] = 0 + } + } + ops[p].dimStrides = st + } + outShape := make([]int, len(outLabels)) + for i, l := range outLabels { + outShape[i] = sizes[l] + } + out := &Array{shape: outShape, dt: operands[0].Dtype()} + for _, o := range operands[1:] { + out.dt = promote(out.dt, o.Dtype()) + } + total := 1 + for _, d := range outShape { + total *= d + } + out.alloc(total) + outStrides := make([]int, len(outLabels)) + run := 1 + for i := len(outLabels) - 1; i >= 0; i-- { + outStrides[i] = run + run *= outShape[i] + } + // Label-position tables. sorted lists the label union ascending; + // depthOf maps a label to its position there and outAxisOf a label + // to its output axis, or -1 when the label is summed away. All hot + // lookups below are slice reads. + sorted := append([]int(nil), allLabels...) + slices.Sort(sorted) + nAll := len(sorted) + nOps := len(operands) + nOut := len(outLabels) + maxLabel := -1 + for _, l := range allLabels { + if l > maxLabel { + maxLabel = l + } + } + depthOf := make([]int, maxLabel+1) + outAxisOf := make([]int, maxLabel+1) + for i := range outAxisOf { + outAxisOf[i] = -1 + } + for i, l := range sorted { + depthOf[l] = i + } + for i, l := range outLabels { + outAxisOf[l] = i + } + // opDepth[p*nAll+d] is the payload offset operand p gains when the + // coordinate of sorted[d] advances by one: the sum of its + // dimensions' strides on that label, zero for broadcast axes and + // for labels the operand does not have. Repeated labels sum, which + // reproduces the diagonal extraction of the per-element walk. + opDepth := make([]int, nOps*nAll) + for p := range ops { + for d, l := range ops[p].dimLabels { + opDepth[p*nAll+depthOf[l]] += ops[p].dimStrides[d] + } + } + // The summed labels keep their ascending sorted order: index 0 is + // the outermost inner-loop axis, index nSum-1 the innermost, so the + // per-slot visit order matches the full walk with the output + // coordinates held fixed. + sumDepths := make([]int, 0, nAll) + for d, l := range sorted { + if outAxisOf[l] < 0 { + sumDepths = append(sumDepths, d) + } + } + nSum := len(sumDepths) + sumSize := make([]int, nSum) + sumTotal := 1 + for t, d := range sumDepths { + sumSize[t] = sizes[sorted[d]] + sumTotal *= sumSize[t] + } + innerDelta := make([]int, nOps*nSum) + for p := range nOps { + for t, d := range sumDepths { + innerDelta[p*nSum+t] = opDepth[p*nAll+d] + } + } + outDelta := make([]int, nOps*nOut) + for p := range nOps { + for i, l := range outLabels { + outDelta[p*nOut+i] = opDepth[p*nAll+depthOf[l]] + } + } + // Readers resolve each operand to a dense payload slice indexed by + // the same logical offsets the per-element walk computed, widened + // exactly where its scalar read widened. Every conversion is + // deterministic, so widening once per call reads back the values + // the per-element widening produced. + // Each worker builds its own visitor, so the cursor and the odometer + // scratch below are per-worker: both are rewritten in full at every + // visit, which is what keeps the slots independent. + var newVisit func() func(slot int, base []int) + switch out.dt { + case Int: + // An Int result forces every operand to Int, read raw; a + // strided one is gathered through physIndex first, mirroring + // the complex branch below. + rv := make([][]int64, nOps) + for p, a := range operands { + if a.strides == nil { + rv[p] = a.ints + continue + } + w := make([]int64, a.Len()) + einsumGather(w, func(i int) int64 { return a.ints[a.physIndex(i)] }) + rv[p] = w + } + newVisit = func() func(slot int, base []int) { + off := make([]int, nOps) + coord := make([]int, nSum) + return func(slot int, base []int) { + einsumSlotSum(rv, out.ints, slot, base, off, coord, innerDelta, sumSize, sumTotal) + } + } + case Complex: + rv := make([][]complex128, nOps) + for p, a := range operands { + switch { + case a.dt == Complex && a.strides == nil: + rv[p] = a.complexes + case a.strides != nil: + // A strided view is gathered through the accessor: its + // payload is not the view's own element order. + w := make([]complex128, a.Len()) + einsumGather(w, func(i int) complex128 { return a.complexAt(i) }) + rv[p] = w + default: + // A real operand widens to complex(v, 0) element by + // element, exactly as the scalar read did. + rv[p] = complexPayload(a) + } + } + newVisit = func() func(slot int, base []int) { + off := make([]int, nOps) + coord := make([]int, nSum) + return func(slot int, base []int) { + einsumSlotSum(rv, out.complexes, slot, base, off, coord, innerDelta, sumSize, sumTotal) + } + } + default: + rv := make([][]float64, nOps) + for p, a := range operands { + rv[p] = einsumRealReader(a) + } + if out.dt == Float32 { + // Per-addend narrowing: each product lands through + // float32(float64(out[slot]) + acc), the exact write the + // walk performed per visit. + newVisit = func() func(slot int, base []int) { + off := make([]int, nOps) + coord := make([]int, nSum) + return func(slot int, base []int) { + einsumSlotSumF32(rv, out.floats32, slot, base, off, coord, innerDelta, sumSize, sumTotal) + } + } + } else { + newVisit = func() func(slot int, base []int) { + off := make([]int, nOps) + coord := make([]int, nSum) + return func(slot int, base []int) { + einsumSlotSum(rv, out.floats, slot, base, off, coord, innerDelta, sumSize, sumTotal) + } + } + } + } + einsumDriveSlots(total, sumTotal, nOut, outShape, outDelta, nOps, newVisit) + return out, nil +} + +// einsumRealReader returns the operand's elements as float64, indexed +// by logical flat offset, widened exactly where the engine's scalar +// read widened: float32 and int from their raw payloads, a float64 +// payload as is, a strided view gathered through the accessor. The +// widening splits across workers from widenParallelMin up: every +// element widens on its own into its own slot, so the buffer holds the +// serial walk's values. +func einsumRealReader(a *Array) []float64 { + if a.strides != nil { + n := a.Len() + w := make([]float64, n) + if n < widenParallelMin { + for i := range w { + w[i] = a.floatAt(i) + } + return w + } + engine.Parallel(n, func(ws, we int) { + for i := ws; i < we; i++ { + w[i] = a.floatAt(i) + } + }) + return w + } + switch a.dt { + case Float32: + return einsumWiden(a.floats32) + case Int: + return einsumWiden(a.ints) + default: + return a.floats + } +} + +// einsumWiden widens a raw payload to float64, exactly per element and +// across workers from widenParallelMin up, so the result reads back the +// serial walk's values. +func einsumWiden[T int64 | float32](src []T) []float64 { + w := make([]float64, len(src)) + if len(src) < widenParallelMin { + for i, v := range src { + w[i] = float64(v) + } + return w + } + engine.Parallel(len(src), func(ws, we int) { + for i := ws; i < we; i++ { + w[i] = float64(src[i]) + } + }) + return w +} + +// einsumGather fills dst by mapping every logical element index through +// read, exactly per element and across workers from widenParallelMin +// up, so the buffer reads back the serial walk's values. +func einsumGather[T any](dst []T, read func(int) T) { + if len(dst) < widenParallelMin { + for i := range dst { + dst[i] = read(i) + } + return + } + engine.Parallel(len(dst), func(ws, we int) { + for i := ws; i < we; i++ { + dst[i] = read(i) + } + }) +} + +// einsumSlotFloor is the number of per-slot visits a worker must carry +// before the slot walk is split across goroutines. A visit is a handful +// of multiply-adds and an odometer step, so a chunk below this costs +// more to spawn and synchronise than it runs: measured on a contraction +// of eight visits per slot, a worker carrying a thousand of them ran +// slower split than whole, one carrying two thousand broke even, and +// one carrying four thousand was ahead. +const einsumSlotFloor = 4096 + +// einsumDriveSlots walks the output slots in row-major order, +// maintaining per-operand base offsets through the output labels' +// strides, and hands each slot to the visitor built for the worker that +// owns it. The slots are split across workers as contiguous chunks, +// each walking its own cursor from its own first slot; which slot is +// written when is free to choose, and visit owns the summation order +// within the slot, so the chunking cannot move a bit. +func einsumDriveSlots(outTotal, sumTotal, nOut int, outShape, outDelta []int, nOps int, newVisit func() func(slot int, base []int)) { + if sumTotal == 0 || outTotal == 0 { + // A zero-sized summed label starves every slot, and a + // zero-sized output axis has no slot at all: the output stays + // zeroed, as the walk left it. + return + } + // The floor is per worker, so it converts into slots through the + // per-slot visit count; a chunk below it never splits. + perSlot := sumTotal * nOps + minSlots := 1 + if perSlot < einsumSlotFloor { + minSlots = (einsumSlotFloor + perSlot - 1) / perSlot + } + engine.ParallelMin(outTotal, minSlots, func(ss, se int) { + visit := newVisit() + base := make([]int, nOps) + coord := make([]int, nOut) + // The chunk's own starting cursor: the walk below only advances + // offsets by strides, so it has to begin from the offset the + // first slot's coordinate already carries. + rem := ss + for i := nOut - 1; i >= 0; i-- { + coord[i] = rem % outShape[i] + rem /= outShape[i] + for p := range nOps { + base[p] += coord[i] * outDelta[p*nOut+i] + } + } + for slot := ss; slot < se; slot++ { + visit(slot, base) + for i := nOut - 1; i >= 0; i-- { + coord[i]++ + for p := range nOps { + base[p] += outDelta[p*nOut+i] + } + if coord[i] < outShape[i] { + break + } + coord[i] = 0 + for p := range nOps { + base[p] -= outDelta[p*nOut+i] * outShape[i] + } + } + } + }) +} + +// einsumSlotSum accumulates one output slot. It visits the summed +// labels' combinations in ascending-odometer order (index 0 outermost, +// the last innermost), multiplies the operands in order from the unit +// seed and adds each product into out[slot]. The innermost axis +// advances at every visit, so it is peeled into its own counted run: +// there only the operand cursors move, and the odometer above them +// steps once per run, which takes the digit bookkeeping off the hot +// path without moving a single addend. The run loop is written out per +// operand count up to three: the hot reads and cursor updates then hold +// their indices in registers, while the outer walk stays generic in +// nSum. The arithmetic of every branch is the walk's: seed, operand +// order, coordinate order and wrap bookkeeping included. +func einsumSlotSum[T int64 | float64 | complex128](rv [][]T, out []T, slot int, base, off, coord, delta, sizes []int, total int) { + nSum := len(sizes) + copy(off, base) + clear(coord[:nSum]) + if nSum == 0 { + // Nothing is summed: the slot takes the product of the + // operands at the cursor, and a repeat of the walk would read + // the same elements again. + for range total { + acc := T(1) + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] += acc + } + return + } + // inner is the innermost axis' extent, runs the number of odometer + // combinations above it. + inner := sizes[nSum-1] + runs := 1 + for t := range nSum - 1 { + runs *= sizes[t] + } + switch len(rv) { + case 1: + r0 := rv[0] + o0 := off[0] + d0 := delta[0*nSum:] + a0 := d0[nSum-1] + for range runs { + for range inner { + acc := T(1) + acc *= r0[o0] + out[slot] += acc + o0 += a0 + } + o0 -= a0 * inner + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + o0 += d0[t] + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + o0 -= d0[t] * sizes[t] + } + } + off[0] = o0 + case 2: + r0, r1 := rv[0], rv[1] + o0, o1 := off[0], off[1] + d0 := delta[0*nSum : 1*nSum] + d1 := delta[1*nSum : 2*nSum] + a0, a1 := d0[nSum-1], d1[nSum-1] + for range runs { + for range inner { + acc := T(1) + acc *= r0[o0] + acc *= r1[o1] + out[slot] += acc + o0 += a0 + o1 += a1 + } + o0 -= a0 * inner + o1 -= a1 * inner + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + o0 += d0[t] + o1 += d1[t] + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + o0 -= d0[t] * sizes[t] + o1 -= d1[t] * sizes[t] + } + } + off[0], off[1] = o0, o1 + case 3: + r0, r1, r2 := rv[0], rv[1], rv[2] + o0, o1, o2 := off[0], off[1], off[2] + d0 := delta[0*nSum : 1*nSum] + d1 := delta[1*nSum : 2*nSum] + d2 := delta[2*nSum : 3*nSum] + a0, a1, a2 := d0[nSum-1], d1[nSum-1], d2[nSum-1] + for range runs { + for range inner { + acc := T(1) + acc *= r0[o0] + acc *= r1[o1] + acc *= r2[o2] + out[slot] += acc + o0 += a0 + o1 += a1 + o2 += a2 + } + o0 -= a0 * inner + o1 -= a1 * inner + o2 -= a2 * inner + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + o0 += d0[t] + o1 += d1[t] + o2 += d2[t] + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + o0 -= d0[t] * sizes[t] + o1 -= d1[t] * sizes[t] + o2 -= d2[t] * sizes[t] + } + } + off[0], off[1], off[2] = o0, o1, o2 + default: + for range runs { + for range inner { + acc := T(1) + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] += acc + for p := range rv { + off[p] += delta[p*nSum+nSum-1] + } + } + for p := range rv { + off[p] -= delta[p*nSum+nSum-1] * inner + } + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + for p := range rv { + off[p] += delta[p*nSum+t] + } + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + for p := range rv { + off[p] -= delta[p*nSum+t] * sizes[t] + } + } + } + } +} + +// einsumSlotSumF32 is einsumSlotSum for a Float32 result: the products +// accumulate in float64 and each addend narrows into out[slot] on +// arrival, float32(float64(out[slot]) + acc), one narrowing per visit +// like the scalar walk. The innermost axis is peeled the same way, and +// the run loop is written out per operand count up to three, mirroring +// einsumSlotSum. +func einsumSlotSumF32(rv [][]float64, out []float32, slot int, base, off, coord, delta, sizes []int, total int) { + nSum := len(sizes) + copy(off, base) + clear(coord[:nSum]) + if nSum == 0 { + for range total { + acc := 1.0 + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] = float32(float64(out[slot]) + acc) + } + return + } + inner := sizes[nSum-1] + runs := 1 + for t := range nSum - 1 { + runs *= sizes[t] + } + switch len(rv) { + case 1: + r0 := rv[0] + o0 := off[0] + d0 := delta[0*nSum:] + a0 := d0[nSum-1] + for range runs { + for range inner { + acc := 1.0 + acc *= r0[o0] + out[slot] = float32(float64(out[slot]) + acc) + o0 += a0 + } + o0 -= a0 * inner + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + o0 += d0[t] + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + o0 -= d0[t] * sizes[t] + } + } + off[0] = o0 + case 2: + r0, r1 := rv[0], rv[1] + o0, o1 := off[0], off[1] + d0 := delta[0*nSum : 1*nSum] + d1 := delta[1*nSum : 2*nSum] + a0, a1 := d0[nSum-1], d1[nSum-1] + for range runs { + for range inner { + acc := 1.0 + acc *= r0[o0] + acc *= r1[o1] + out[slot] = float32(float64(out[slot]) + acc) + o0 += a0 + o1 += a1 + } + o0 -= a0 * inner + o1 -= a1 * inner + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + o0 += d0[t] + o1 += d1[t] + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + o0 -= d0[t] * sizes[t] + o1 -= d1[t] * sizes[t] + } + } + off[0], off[1] = o0, o1 + case 3: + r0, r1, r2 := rv[0], rv[1], rv[2] + o0, o1, o2 := off[0], off[1], off[2] + d0 := delta[0*nSum : 1*nSum] + d1 := delta[1*nSum : 2*nSum] + d2 := delta[2*nSum : 3*nSum] + a0, a1, a2 := d0[nSum-1], d1[nSum-1], d2[nSum-1] + for range runs { + for range inner { + acc := 1.0 + acc *= r0[o0] + acc *= r1[o1] + acc *= r2[o2] + out[slot] = float32(float64(out[slot]) + acc) + o0 += a0 + o1 += a1 + o2 += a2 + } + o0 -= a0 * inner + o1 -= a1 * inner + o2 -= a2 * inner + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + o0 += d0[t] + o1 += d1[t] + o2 += d2[t] + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + o0 -= d0[t] * sizes[t] + o1 -= d1[t] * sizes[t] + o2 -= d2[t] * sizes[t] + } + } + off[0], off[1], off[2] = o0, o1, o2 + default: + for range runs { + for range inner { + acc := 1.0 + for p := range rv { + acc *= rv[p][off[p]] + } + out[slot] = float32(float64(out[slot]) + acc) + for p := range rv { + off[p] += delta[p*nSum+nSum-1] + } + } + for p := range rv { + off[p] -= delta[p*nSum+nSum-1] * inner + } + for t := nSum - 2; t >= 0; t-- { + coord[t]++ + for p := range rv { + off[p] += delta[p*nSum+t] + } + if coord[t] < sizes[t] { + break + } + coord[t] = 0 + for p := range rv { + off[p] -= delta[p*nSum+t] * sizes[t] + } + } + } + } +} diff --git a/internal/core/einsum_bench_test.go b/internal/core/einsum_bench_test.go new file mode 100644 index 0000000..ee4ec79 --- /dev/null +++ b/internal/core/einsum_bench_test.go @@ -0,0 +1,98 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Einsum benchmarks guard both the fast paths and the general engine. +// The batched shapes are the training-loop workloads; the general +// engine shapes are the ones the dispatch table declines. + +func benchVals(b *testing.B, seed, n int) []float64 { + b.Helper() + v := make([]float64, n) + for i := range v { + v[i] = float64(i%17)*float64(seed%5) + float64(i%7) - 6 + } + return v +} + +func benchMat(b *testing.B, seed int, shape ...int) *Array { + b.Helper() + n := 1 + for _, d := range shape { + n *= d + } + a, err := FromFloats(benchVals(b, seed, n), shape...) + if err != nil { + b.Fatal(err) + } + return a +} + +// BenchmarkEinsumMatMulPath exercises the dispatch table's MatMul fast +// path as the control. +func BenchmarkEinsumMatMulPath(b *testing.B) { + a := benchMat(b, 1, 64, 64) + c := benchMat(b, 2, 64, 64) + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum("ij,jk->ik", a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkEinsumBatchedMatMul is the batched product the table has no +// fast path for: it falls to the general engine. +func BenchmarkEinsumBatchedMatMul(b *testing.B) { + a := benchMat(b, 3, 16, 64, 64) + c := benchMat(b, 4, 16, 64, 64) + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum("bij,bjk->bik", a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkEinsumBatchedSmall keeps the general engine's fixed costs +// visible at a size where they dominate. +func BenchmarkEinsumBatchedSmall(b *testing.B) { + a := benchMat(b, 5, 4, 32, 32) + c := benchMat(b, 6, 4, 32, 32) + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum("bij,bjk->bik", a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkEinsumEllipsisMatMul is the ellipsis spelling of the batched +// product. +func BenchmarkEinsumEllipsisMatMul(b *testing.B) { + a := benchMat(b, 7, 8, 48, 48) + c := benchMat(b, 8, 48, 48) + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum("...ij,jk->...ik", a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkEinsumBilinear sums a shared label across three operands, +// the pattern attention and tensordot shapes reduce to. +func BenchmarkEinsumBilinear(b *testing.B) { + a := benchMat(b, 9, 16, 32) + c := benchMat(b, 10, 32, 32) + d := benchMat(b, 11, 32, 16) + b.ReportAllocs() + for b.Loop() { + if _, err := Einsum("ik,kj,jl->il", a, c, d); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/einsum_foldorder_test.go b/internal/core/einsum_foldorder_test.go new file mode 100644 index 0000000..0543ac6 --- /dev/null +++ b/internal/core/einsum_foldorder_test.go @@ -0,0 +1,31 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// TestEinsumFoldOrderLargestLabelFirst pins the fold order of the +// single-operand axis reduction: einsumSumAxes closes the dropped axis +// carrying the largest label first, which mirrors the engine's nesting +// (the engine walks the sorted label union smallest outermost, so the +// largest label runs innermost and has to be summed away first). The +// choice is provable only where summation is inexact, so the operand +// carries 1, 1, 1e100 and -1e100 across two summed axes: folding the +// largest label first pairs 1e100 with -1e100 and answers exactly 2, +// while the opposite order would let 1e100 swallow both 1s and answer 0. +func TestEinsumFoldOrderLargestLabelFirst(t *testing.T) { + // a[0][j][k]: the dropped axes j and k carry labels 9 and 10, so + // k folds first and the sum runs sum_j (sum_k a[0][j][k]). + a := mustFloats(t, []float64{1, 1, 1e100, -1e100}, 1, 2, 2) + out, err := Einsum("ijk->i", a) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + if out.NDim() != 1 || out.Len() != 1 { + t.Fatalf("Einsum answered shape %s, want one slot", shapeText(out.Shape())) + } + if got := out.FloatAt(0); got != 2 { + t.Fatalf("fold order answered %v, want exactly 2; the opposite order answers 0", got) + } +} diff --git a/internal/core/einsum_norm_special_pins_test.go b/internal/core/einsum_norm_special_pins_test.go new file mode 100644 index 0000000..0228357 --- /dev/null +++ b/internal/core/einsum_norm_special_pins_test.go @@ -0,0 +1,562 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + "strings" + "testing" + "time" +) + +// Regression pins for the einsum broadcast extents and norms and the +// Bessel evaluation paths: the ellipsis sizing rule, the einsum and +// reduction dtype contracts, the large-argument Bessel cost and +// accuracy, and the elliptic integrals at negative parameters. + +// TestEinsumEllipsisExtentMismatch pins the ellipsis sizing rule: an +// ellipsis axis broadcasts only when it is 1. A wider axis that +// disagrees with the broadcast extent used to be sized by max() and +// walked with its unit stride, reading past the operand (a panic for a +// fresh array, silently the aliased payload's tail for a Slice view). +func TestEinsumEllipsisExtentMismatch(t *testing.T) { + check := func(name string, err error) { + t.Helper() + if err == nil { + t.Fatalf("%s: expected a refusal, got a result", name) + } + if !strings.Contains(err.Error(), "ellipsis") { + t.Fatalf("%s: error %v does not name the ellipsis axis", name, err) + } + if !strings.Contains(err.Error(), "1") { + t.Fatalf("%s: error %v does not name the size-1 rule", name, err) + } + } + // Fresh arrays: 3 ones against 6 twos. + ones, err := FromFloats([]float64{1, 1, 1}, 3) + if err != nil { + t.Fatal(err) + } + twos, err := FromFloats([]float64{2, 2, 2, 2, 2, 2}, 6) + if err != nil { + t.Fatal(err) + } + out, err := Einsum("...,...->...", ones, twos) + if err == nil { + t.Fatalf("Einsum(...,...->...) with extents 3 and 6 = %v, want an error", out) + } + check("1-D extents 3 and 6", err) + + // Only the ellipsis axis disagrees: a (2,3,4) against a (5,3,4). + a, err := FromFloats(seqFloats(24), 2, 3, 4) + if err != nil { + t.Fatal(err) + } + b, err := FromFloats(onesFloats(5*3*4), 5, 3, 4) + if err != nil { + t.Fatal(err) + } + out, err = Einsum("...ij,...ij->...ij", a, b) + if err == nil { + t.Fatalf("Einsum(...ij,...ij->...ij) with ellipsis extents 2 and 5 = %v, want an error", out) + } + check("3-D ellipsis-only mismatch", err) + + // A Slice view whose payload runs past Len(): the same refusal, not + // the aliased tail values. + big, err := FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 9) + if err != nil { + t.Fatal(err) + } + view, err := Slice(big, 0, 0, 3) + if err != nil { + t.Fatal(err) + } + if view.Len() != 3 { + t.Fatalf("Slice view len = %d, want 3", view.Len()) + } + out, err = Einsum("...,...->...", view, twos) + if err == nil { + t.Fatalf("Einsum(...,...->...) on a Slice view = %v, want an error", out) + } + check("Slice view extents 3 and 6", err) + + // The legitimate broadcasts keep working: a size-1 axis stretches, + // and equal extents pass through. + bcast, err := FromFloats([]float64{1, 2, 3, 4}, 1, 4) + if err != nil { + t.Fatal(err) + } + rows, err := FromFloats(seqFloats(12), 3, 4) + if err != nil { + t.Fatal(err) + } + got, err := Einsum("...i,...i->...i", bcast, rows) + if err != nil { + t.Fatalf("size-1 ellipsis broadcast: %v", err) + } + want := []float64{1, 4, 9, 16, 5, 12, 21, 32, 9, 20, 33, 48} + if len(got.RawFloats()) != len(want) { + t.Fatalf("size-1 broadcast len = %d, want %d", got.Len(), len(want)) + } + for i, w := range want { + if got.FloatAt(i) != w { + t.Fatalf("size-1 broadcast [%d] = %v, want %v", i, got.FloatAt(i), w) + } + } + if _, err := Einsum("...,...->...", rows, rows); err != nil { + t.Fatalf("equal ellipsis extents: %v", err) + } +} + +// seqFloats returns 1..n as float64, the element pattern the einsum +// tests read back. +func seqFloats(n int) []float64 { + out := make([]float64, n) + for i := range out { + out[i] = float64(i + 1) + } + return out +} + +// onesFloats returns n copies of 1. +func onesFloats(n int) []float64 { + out := make([]float64, n) + for i := range out { + out[i] = 1 + } + return out +} + +// TestEinsumRepeatedOutputLabel pins the output-label uniqueness rule. +// "ii->ii" is invalid einsum (a typo for the diagonal "ii->i"), but the +// identity fast path used to accept it and return a copy of the whole +// matrix. +func TestEinsumRepeatedOutputLabel(t *testing.T) { + m, err := FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatal(err) + } + out, err := Einsum("ii->ii", m) + if err == nil { + t.Fatalf("Einsum(ii->ii) = %v, want an error", out) + } + if !strings.Contains(err.Error(), "repeats") { + t.Fatalf("Einsum(ii->ii) error = %v, want a repeated-output-label refusal", err) + } + // The intended diagonal spelling still works. + diag, err := Einsum("ii->i", m) + if err != nil { + t.Fatalf("Einsum(ii->i): %v", err) + } + if diag.Len() != 2 || diag.FloatAt(0) != 1 || diag.FloatAt(1) != 4 { + t.Fatalf("Einsum(ii->i) = %v, want [1 4]", diag.RawFloats()[:2]) + } + // Distinct output labels, including a permutation, are untouched. + if _, err := Einsum("ij->ij", m); err != nil { + t.Fatalf("Einsum(ij->ij): %v", err) + } + if _, err := Einsum("ij->ji", m); err != nil { + t.Fatalf("Einsum(ij->ji): %v", err) + } + // A repeated lhs label with a valid output is still the diagonal. + trace, err := Einsum("ii->", m) + if err != nil { + t.Fatalf("Einsum(ii->): %v", err) + } + if got := trace.FloatAt(0); got != 5 { + t.Fatalf("Einsum(ii->) = %v, want 5", got) + } +} + +// TestNormRejectsNaNP pins the Norm parameter check: NaN used to slip +// through the `p <= 0` test and produce a silent all-NaN result, while +// every other invalid p was refused loudly. +func TestNormRejectsNaNP(t *testing.T) { + a, err := FromFloats([]float64{3, 4, 5, 6, 7}, 5) + if err != nil { + t.Fatal(err) + } + out, err := Norm(a, math.NaN(), 0, false) + if err == nil { + t.Fatalf("Norm(p=NaN) = %v, want an error", out.RawFloats()) + } + if !strings.Contains(err.Error(), "must be positive") { + t.Fatalf("Norm(p=NaN) error = %v, want the positivity refusal", err) + } + // The valid parameters keep their values: the 2-norm and the Inf + // norm (max-abs) are both pinned. + got, err := Norm(a, 2, 0, false) + if err != nil { + t.Fatalf("Norm(p=2): %v", err) + } + if want := math.Sqrt(9 + 16 + 25 + 36 + 49); got.FloatAt(0) != want { + t.Fatalf("Norm(p=2) = %v, want %v", got.FloatAt(0), want) + } + got, err = Norm(a, math.Inf(1), 0, false) + if err != nil { + t.Fatalf("Norm(p=+Inf): %v", err) + } + if got.FloatAt(0) != 7 { + t.Fatalf("Norm(p=+Inf) = %v, want 7", got.FloatAt(0)) + } +} + +// TestAiryPositiveBranchAccuracy pins Ai over the upper half of the +// documented window. The float64 reconstruction c1·f − c2·g cancels +// e^{2·zeta}, with zeta = 2x^{3/2}/3, so the old single-precision +// branch was off by 5.7e-8 at x = 6 and 3.6e-3 at x = 8; the values +// below are an independent 400-bit evaluation of the same Maclaurin +// pair, cross-checked against the Bessel connection +// Ai(x) = (1/pi)·sqrt(x/3)·K_{1/3}(2x^{3/2}/3) to better than 1e-85. +func TestAiryPositiveBranchAccuracy(t *testing.T) { + cases := []struct { + x float64 + want float64 + }{ + {1, 0.135292416312881415524147423515}, + {2, 0.0349241304232743791353220807918}, + {3, 0.00659113935746071897805587355903}, + {3.5, 0.0025840987869896349632771447833}, + {4, 0.0009515638512048018736214999689}, + {5, 0.000108344428136074417349865025033}, + {6, 9.94769436025288957023884766883e-06}, + {7, 7.4921288639971670807710402721e-07}, + {8, 4.69220761609923162564908170349e-08}, + } + for _, c := range cases { + ai, _, err := Airy(mustFloats(t, []float64{c.x}, 1)) + if err != nil { + t.Fatalf("Airy(%v): %v", c.x, err) + } + got := ai.FloatAt(0) + if rel := math.Abs(got-c.want) / c.want; rel > 1e-13 { + t.Errorf("Ai(%v) = %.17g, want %.17g, relative error %.3g", c.x, got, c.want, rel) + } + } + // The extended-precision constants must round to the float64 + // constants the series branch uses, so neither can drift alone. + for _, c := range []struct { + digits string + asF64 float64 + }{ + {airyC1Digits, 0.35502805388781723926}, + {airyC2Digits, 0.25881940379280679840}, + } { + f, _, err := big.ParseFloat(c.digits, 10, 256, big.ToNearestEven) + if err != nil { + t.Fatalf("ParseFloat(%q): %v", c.digits, err) + } + got, _ := f.Float64() + if got != c.asF64 { + t.Errorf("%s rounds to %.20g, want %.20g", c.digits, got, c.asF64) + } + } +} + +// TestEllipticKExtremeNegativeParameter pins K(-s) for s past 2^53, +// where the float64 sum 1+s rounds to s. The transformed parameter +// s/(1+s) then rounds to exactly 1 and the AGM returned its round-off +// floor: 1.77e7 at s = 1e16 and 1.3e-139 at s = MaxFloat64, against +// the true K, which stays finite and decays like +// ln(4 sqrt(1+s))/sqrt(1+s). The wanted values are an independent +// 300-bit AGM, K(1-eps) = pi/(2 AGM(1, sqrt(eps))), fed with the exact +// complement eps = 1/(1+s). +func TestEllipticKExtremeNegativeParameter(t *testing.T) { + cases := []struct { + s float64 + want float64 + }{ + {1.5, 1.23301490843926637950440475303}, + {1e8, 0.00105966347091044866547751716942}, + {1e15, 5.89944481901753189416588265882e-07}, + {1e16, 1.98069751050722556208040182536e-07}, + {1e18, 2.21095601980663017697189972856e-08}, + {1e100, 1.16515549010822173901218438276e-48}, + {1e292, 3.37563717938250558254616096941e-144}, + {1e300, 3.46774058310226734144141165422e-148}, + {math.MaxFloat64, 2.65724011463622780028451982264e-152}, + } + for _, c := range cases { + k, err := EllipticK(mustFloats(t, []float64{-c.s}, 1)) + if err != nil { + t.Fatalf("EllipticK(-%v): %v", c.s, err) + } + got := k.FloatAt(0) + if math.IsInf(got, 0) || math.IsNaN(got) { + t.Errorf("K(-%v) = %v, want %v", c.s, got, c.want) + continue + } + if rel := math.Abs(got-c.want) / c.want; rel > 1e-14 { + t.Errorf("K(-%v) = %.17g, want %.17g, relative error %.3g", c.s, got, c.want, rel) + } + } +} + +// TestBesselLargeArgumentCostAndAccuracy pins the large-argument +// evaluation paths. The Miller recurrence starts at order n + |x| + 40, +// so one evaluation at x = 1e9 used to cost about 1e9 steps (minutes +// per element on a quiet machine) however small the order was, and +// the plain float64 phase x − nπ/2 − π/4 had already rounded the +// fraction of a large x away (1e-8 relative at x = 1e9). Orders far +// below the argument now climb from the asymptotic seeds at O(n) cost, +// the phase is reduced modulo 2π in two parts, and the downward walk +// keeps the orders where the climb would amplify the parasitic K +// component it inherits from the seeds. +// +// The wanted values come from an independent 420-bit reference: the +// ascending series at the moderate arguments, cross-checked against the +// large-argument expansion with an exactly reduced phase to better than +// 1e-30 there and against the closed forms of the spherical functions; +// at 1e6 and above the expansion alone, whose truncation error is below +// 1e-100. +func TestBesselLargeArgumentCostAndAccuracy(t *testing.T) { + // The cost is the point, so it is pinned first and tightly enough + // that the old walk cannot pass it. The climb does this work in + // microseconds, which leaves five orders of magnitude of margin. + slow := func(name string, f func()) { + t.Helper() + began := time.Now() + f() + if elapsed := time.Since(began); elapsed > 250*time.Millisecond { + t.Fatalf("%s at x = 1e9 took %v: the O(|x|) walk is back", name, elapsed) + } + } + slow("BesselJ", func() { BesselJ(3, 1e9) }) + slow("BesselIn", func() { + if _, err := BesselIn(3, mustFloats(t, []float64{1e9}, 1)); err != nil { + t.Fatal(err) + } + }) + slow("SphericalBesselJ", func() { + if _, err := SphericalBesselJ(2, mustFloats(t, []float64{1e9}, 1)); err != nil { + t.Fatal(err) + } + }) + close := func(got, want float64, tol float64) bool { + if math.IsNaN(got) || math.IsInf(got, 0) { + return false + } + return math.Abs(got-want) <= tol*math.Abs(want) + } + for _, c := range []struct { + n int + x float64 + want float64 + }{ + // The phase reduction: order zero rides it directly, and 1e12 is + // the argument whose Miller walk used to hang. + {0, 1e6, 0.00033104301373987374098796304}, + {0, 1e9, 2.4687471886269195114428159e-05}, + {0, 1e12, 1.016712505004068170196092e-07}, + // The climb from the asymptotic seeds. + {3, 1e6, 0.00072596703263590033550304933}, + {3, 1e9, 5.2104225428039902129072166e-06}, + {3, 1e12, 7.9138026838463740158213673e-07}, + {7, 1e9, 5.2104220490545510286050221e-06}, + // The parity law survives the climb at a large argument. + {3, -1e9, -5.2104225428039902129072166e-06}, + // The crossover and the split at n = |x|: both sides of the + // boundary must agree with the reference. + {0, 15.5, -0.10923065090005016848282335}, + {3, 15.5, -0.13345665257394448939628244}, + {7, 15.5, 0.12160445971276508316409274}, + {7, 500, -0.008824247657353861568197967}, + {19, 20, 0.21886190352168099119079698}, + {20, 20, 0.16474777377532653234118514}, + {21, 20, 0.1106336440289720734915733}, + {50, 50, 0.121409021897615063820108}, + {51, 50, 0.0916229012737578895575967}, + {99, 100, 0.11524392532303779883241611}, + {100, 100, 0.096366673295861559674314025}, + {101, 100, 0.077489421268685320516211937}, + } { + if got := BesselJ(c.n, c.x); !close(got, c.want, 1e-12) { + t.Errorf("BesselJ(%d, %g) = %.17g, want %.17g", c.n, c.x, got, c.want) + } + } + for _, c := range []struct { + l int + x float64 + want float64 + }{ + // The spherical climb, including orders a quarter and a half of + // the argument where the recurrence sits near the turning point. + {0, 20, 0.045647262536381382718804999}, + {1, 20, -0.018121739963850530167173143}, + {2, 20, -0.048365523530958962243880971}, + {3, 20, 0.0060303590811107896062029}, + {5, 20, 0.0166839080630956927665205}, + {10, 20, 0.0396866986446263713096385}, + {2, 100, 0.0048034416524879534799540249}, + {2, 1e6, 3.4999069191386037217677779e-07}, + {0, 1e9, 5.4584344944869956424438727e-10}, + {1, 1e9, -8.3788718081805888494107608e-10}, + {2, 1e9, -5.4584345196236110669856393e-10}, + {3, 1e9, 8.37887178088841625129271e-10}, + {2, 1e12, 6.1123870237451515928646e-13}, + // The parity law survives the climb at a large argument. + {3, -1e9, -8.37887178088841625129271e-10}, + } { + out, err := SphericalBesselJ(c.l, mustFloats(t, []float64{c.x}, 1)) + if err != nil { + t.Fatalf("SphericalBesselJ(%d, %g): %v", c.l, c.x, err) + } + if got := out.FloatAt(0); !close(got, c.want, 1e-12) { + t.Errorf("j(%d, %g) = %.17g, want %.17g", c.l, c.x, got, c.want) + } + } + for _, c := range []struct { + n int + x float64 + want float64 + }{ + // The upward climb, at the split boundary n² = 4x as well as far + // below it. Without the split the near-boundary orders come back + // with only a handful of correct digits. + {2, 500, 2.49480026292137369656865e+215}, + {3, 700, 1.5197848119270448202752878e+302}, + {30, 700, 8.0395148044586219063553629e+301}, + {30, 100, 1.20615487044984340057805e+40}, + {20, 20, 3188.75032885361480155313}, + // The downward walk above the split: the quotient must be formed + // before the multiplication, or I₀(x)·jₙ overflows to +Inf even + // though the answer is an ordinary number. + {50, 500, 2.05521801630540860858473e+214}, + {100, 500, 1.16377328686043708878981e+211}, + {150, 500, 4.8882964482285023996976e+205}, + // The odd order keeps its parity through the climb as well. + {3, -700, -1.5197848119270448202752878e+302}, + } { + out, err := BesselIn(c.n, mustFloats(t, []float64{c.x}, 1)) + if err != nil { + t.Fatalf("BesselIn(%d, %g): %v", c.n, c.x, err) + } + if got := out.FloatAt(0); !close(got, c.want, 1e-12) { + t.Errorf("I(%d, %g) = %.17g, want %.17g", c.n, c.x, got, c.want) + } + } + // The same phase and asymptotic seeds feed the second kind, whose + // table is the Wronskian J₁·Y₀ − J₀·Y₁ = 2/(πx) against the J values + // above (it holds to 3e-19 at x = 20, better further out). + for _, c := range []struct { + n int + x float64 + want float64 + }{ + {0, 1e6, -0.000725968522335179165682722}, + {1, 1e6, -0.000331043376724176288863517}, + {0, 1e9, -5.21042265389761374215067e-06}, + {1, 1e9, -2.46874718888744064444629e-05}, + {0, 1e12, -7.91380268385094922209395e-07}, + {1, 1e12, -1.01671250500802507153802e-07}, + } { + got, err := BesselY(c.n, c.x) + if err != nil { + t.Fatalf("BesselY(%d, %g): %v", c.n, c.x, err) + } + if !close(got, c.want, 1e-12) { + t.Errorf("Y(%d, %g) = %.17g, want %.17g", c.n, c.x, got, c.want) + } + } + // The Wronskian ties the two tables: it must hold to machine + // precision at a moderate argument, where both branches are exact. + for _, x := range []float64{20, 100} { + j0 := BesselJ(0, x) + j1 := BesselJ(1, x) + y0, err := BesselY(0, x) + if err != nil { + t.Fatal(err) + } + y1, err := BesselY(1, x) + if err != nil { + t.Fatal(err) + } + want := 2 / (math.Pi * x) + if got := j1*y0 - j0*y1; math.Abs(got-want) > 1e-13*want { + t.Errorf("Wronskian at %g = %.17g, want %.17g", x, got, want) + } + } + // I₃(1e9) genuinely overflows float64, and must say so rather than + // fold the overflow into a NaN. + big, err := BesselIn(3, mustFloats(t, []float64{1e9}, 1)) + if err != nil { + t.Fatalf("BesselIn(3, 1e9): %v", err) + } + if !math.IsInf(big.FloatAt(0), 1) { + t.Errorf("I₃(1e9) = %g, want +Inf", big.FloatAt(0)) + } + // The phase itself, against the exact reduction of x − nπ/2 − π/4 + // modulo 2π into [−π, π] computed in 420-bit arithmetic. + for _, c := range []struct { + x float64 + want float64 + }{ + {15.5, 2.1482312222433787365337656}, + {1e3, 0.18813799504830185926374327}, + {1e6, -1.1429623304831833536309925}, + {1e9, -0.20800273989606314020601372}, + {1e12, -1.4430229225342347770949126}, + } { + if got := besselPhase(0, c.x); math.Abs(got-c.want) > 1e-15 { + t.Errorf("besselPhase(0, %g) = %.17g, want %.17g", c.x, got, c.want) + } + } +} + +// TestNewSparseCOONilOperands pins the constructor's contract: a nil +// array, which is what a caller holds after a failed constructor, is a +// refusal, not a nil-pointer panic. NewSparseCOO used to reach +// values.Len() unchecked, and indices.dt before it. +func TestNewSparseCOONilOperands(t *testing.T) { + idx, err := FromInts([]int64{0, 1, 2}, 3, 1) + if err != nil { + t.Fatal(err) + } + vals, err := FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + t.Fatal(err) + } + // The nil a failed constructor hands back. + failed, err := FromFloats([]float64{1, 2}, 3) + if err == nil || failed != nil { + t.Fatalf("FromFloats with the wrong length = %v, %v; want a nil array and an error", failed, err) + } + failedIdx, err := FromInts([]int64{0}, 9, 1) + if err == nil || failedIdx != nil { + t.Fatalf("FromInts with the wrong length = %v, %v; want a nil array and an error", failedIdx, err) + } + for _, c := range []struct { + name string + idx *Array + vals *Array + wantMsg string + }{ + {"nil values", idx, failed, "values"}, + {"nil indices", failedIdx, vals, "indices"}, + {"both nil", nil, nil, "indices"}, + } { + got, err := NewSparseCOO(c.idx, c.vals, []int{3}) + if err == nil { + t.Fatalf("NewSparseCOO(%s) = %+v, want an error", c.name, got) + } + if !strings.Contains(err.Error(), c.wantMsg) { + t.Errorf("NewSparseCOO(%s) error = %v, want it to name %s", c.name, err, c.wantMsg) + } + } + // The well-formed pair still builds and materialises. + s, err := NewSparseCOO(idx, vals, []int{3}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + dense, err := s.Dense() + if err != nil { + t.Fatalf("Dense: %v", err) + } + if want := []float64{1, 2, 3}; dense.Len() != 3 || + dense.FloatAt(0) != want[0] || dense.FloatAt(1) != want[1] || dense.FloatAt(2) != want[2] { + t.Fatalf("Dense = %v, want %v", dense.RawFloats()[:dense.Len()], want) + } +} diff --git a/internal/core/einsum_special_pins_test.go b/internal/core/einsum_special_pins_test.go new file mode 100644 index 0000000..e418338 --- /dev/null +++ b/internal/core/einsum_special_pins_test.go @@ -0,0 +1,279 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// Regression pins for the einsum special forms and the special-function +// extremes: the outer product under swapped output labels, the complex +// dot, spherical Bessel and Fresnel at tiny and large arguments, Ei and +// beta at large arguments, high-order Bessel, NaN ordering in the +// extrema and the complex Jacobian probe. + +// TestEinsumOuterSwappedRHS pins the flat-index mapping of the outer +// product when the rhs labels swap the operand order: "i,j->ji" must +// return the transpose of "i,j->ij". +func TestEinsumOuterSwappedRHS(t *testing.T) { + a, _ := FromFloats([]float64{1, 2, 3}, 3) + b, _ := FromFloats([]float64{4, 5}, 2) + ij, err := Einsum("i,j->ij", a, b) + if err != nil { + t.Fatalf("Einsum ij: %v", err) + } + ji, err := Einsum("i,j->ji", a, b) + if err != nil { + t.Fatalf("Einsum ji: %v", err) + } + if ji.Shape()[0] != 2 || ji.Shape()[1] != 3 { + t.Fatalf("ji shape %v, want [2 3]", ji.Shape()) + } + for i := range 2 { + for j := range 3 { + if ji.FloatAt(i*3+j) != ij.FloatAt(j*2+i) { + t.Fatalf("ji[%d][%d] = %g, want %g", i, j, ji.FloatAt(i*3+j), ij.FloatAt(j*2+i)) + } + } + } +} + +// TestEinsumComplexDot pins the complex contract of the reduced +// patterns: the imaginary part must survive "i,i->" and "ij,ji->". +func TestEinsumComplexDot(t *testing.T) { + a, _ := FromComplexes([]complex128{1 + 2i, 3 - 1i}, 2) + got, err := Einsum("i,i->", a, a) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + if got.Dtype() != Complex { + t.Fatalf("dtype %s, want complex", got.Dtype()) + } + want := (1+2i)*(1+2i) + (3-1i)*(3-1i) + z := got.ComplexAt(0) + if real(z) != real(want) || imag(z) != imag(want) { + t.Fatalf("dot = %v, want %v", z, want) + } + m, _ := FromComplexes([]complex128{1 + 1i, 2 - 1i, 0, 1}, 2, 2) + got2, err := Einsum("ij,ji->", m, m) + if err != nil { + t.Fatalf("Einsum ij,ji: %v", err) + } + if got2.Dtype() != Complex { + t.Fatalf("dtype %s, want complex", got2.Dtype()) + } + // sum over i,j of m[i][j]·m[j][i] + var want2 complex128 + for i := range 2 { + for j := range 2 { + want2 += m.ComplexAt(i*2+j) * m.ComplexAt(j*2+i) + } + } + z2 := got2.ComplexAt(0) + if real(z2) != real(want2) || imag(z2) != imag(want2) { + t.Fatalf("inner = %v, want %v", z2, want2) + } +} + +// TestSphericalBesselTinyArgument pins the small-x series branch: the +// Miller walk used to overflow its unscaled seed and return NaN for +// |x| below ~1e-6. +func TestSphericalBesselTinyArgument(t *testing.T) { + for _, l := range []int{0, 1, 3, 5} { + x, _ := FromFloats([]float64{1e-7}, 1) + out, err := SphericalBesselJ(l, x) + if err != nil { + t.Fatalf("SphericalBesselJ(%d): %v", l, err) + } + got := out.FloatAt(0) + if math.IsNaN(got) || math.IsInf(got, 0) { + t.Fatalf("j_%d(1e-7) = %g, want finite", l, got) + } + // j_l(x) ≈ x^l/(2l+1)!!·(1 − x²/(2(2l+3))) for tiny x; the + // tolerance covers the next series order. + dbl := 1.0 + for k := 3; k <= 2*l+1; k += 2 { + dbl *= float64(k) + } + leading := math.Pow(1e-7, float64(l)) / dbl + tol := leading*(1e-7*1e-7/(2*float64(2*l+3)))*1.1 + 4*math.Abs(got)*1e-16 + 1e-300 + if math.Abs(got-leading) > tol { + t.Fatalf("j_%d(1e-7) = %g, want %g ± %g", l, got, leading, tol) + } + } +} + +// TestSphericalBesselHighDegree pins the Miller rescaling: a +// degree-argument combination whose unscaled walk overflows must still +// return the representable true value. +func TestSphericalBesselHighDegree(t *testing.T) { + x, _ := FromFloats([]float64{5}, 1) + out, err := SphericalBesselJ(200, x) + if err != nil { + t.Fatalf("SphericalBesselJ: %v", err) + } + got := out.FloatAt(0) + if math.IsNaN(got) || math.IsInf(got, 0) { + t.Fatalf("j_200(5) = %g, want finite", got) + } + // Cross-check against the power series form: j_l(x) ≈ x^l/(2l+1)!! + // is the leading term; x = 5 is not tiny, so verify against the + // known recurrence instead: j_{l-1} and j_{l+1} bracket j_l·(2l+1)/x. + jm, _ := SphericalBesselJ(199, x) + jp, _ := SphericalBesselJ(201, x) + // The three-term recurrence (2l+1)/x·j_l = j_{l-1} + j_{l+1} holds + // to rounding for exact Bessel values. + recon := (jm.FloatAt(0) + jp.FloatAt(0)) * 5 / 401 + if math.Abs(recon-got) > 1e-12*math.Max(math.Abs(got), 1e-300) { + t.Fatalf("recurrence check: j_200(5) = %g, reconstruction %g", got, recon) + } +} + +// TestFresnelLargeArgument pins the big-float precision budget: at +// x = 25 the series must still land on the asymptotic limit. +func TestFresnelLargeArgument(t *testing.T) { + x, _ := FromFloats([]float64{25}, 1) + c, err := FresnelC(x) + if err != nil { + t.Fatalf("FresnelC: %v", err) + } + s, err := FresnelS(x) + if err != nil { + t.Fatalf("FresnelS: %v", err) + } + // The asymptotic limits are C(∞) = S(∞) = 1/2, approached with an + // oscillation of amplitude ~1/(πx) ≈ 0.013 at x = 25; the previous + // precision deficit produced values around 1e40 here. + if math.Abs(c.FloatAt(0)-0.5) > 0.02 || math.Abs(s.FloatAt(0)-0.5) > 0.02 { + t.Fatalf("C(25) = %g, S(25) = %g, both within 0.02 of 0.5", c.FloatAt(0), s.FloatAt(0)) + } +} + +// TestExpIntegralEiLarge pins the asymptotic branch: the convergent +// series cap used to truncate silently for x ≳ 250. +func TestExpIntegralEiLarge(t *testing.T) { + x, _ := FromFloats([]float64{100, 250}, 2) + out, err := ExpIntegralEi(x) + if err != nil { + t.Fatalf("ExpIntegralEi: %v", err) + } + for i, xv := range []float64{100, 250} { + got := out.FloatAt(i) + // Ei(x)·x/e^x = 1 + 1/x + 2/x² + ... > 1 always. + scaled := got * xv / math.Exp(xv) + if scaled < 1 || scaled > 1.06 { + t.Fatalf("Ei(%g)·x/e^x = %g, want just above 1", xv, scaled) + } + } +} + +// TestBetaLargeArguments pins the log-space evaluation: Γ(100) overflows +// but B(100, 80) ≈ 3.4e-67 is representable. +func TestBetaLargeArguments(t *testing.T) { + x, _ := FromFloats([]float64{100}, 1) + y, _ := FromFloats([]float64{80}, 1) + out, err := Beta(x, y) + if err != nil { + t.Fatalf("Beta: %v", err) + } + want := math.Exp(lgammaF(100) + lgammaF(80) - lgammaF(180)) + if math.Abs(out.FloatAt(0)-want) > 1e-10*want { + t.Fatalf("B(100,80) = %g, want %g", out.FloatAt(0), want) + } +} + +func lgammaF(v float64) float64 { + l, _ := math.Lgamma(v) + return l +} + +// TestBesselHighOrder pins the Miller rescaling for BesselJ and BesselIn: +// orders whose unscaled descent overflows must return the representable +// true value, not NaN. +func TestBesselHighOrder(t *testing.T) { + jv := BesselJ(240, 20) + if !isFinite(jv) || jv == 0 { + t.Fatalf("J_240(20) = %g, want finite non-zero", jv) + } + xs, _ := FromFloats([]float64{1}, 1) + in, err := BesselIn(120, xs) + if err != nil { + t.Fatalf("BesselIn: %v", err) + } + // I_120(1) ≈ (1/2)^120/120! · correction, a tiny positive number. + if !isFinite(in.FloatAt(0)) || in.FloatAt(0) <= 0 { + t.Fatalf("I_120(1) = %g, want finite positive", in.FloatAt(0)) + } + xt, _ := FromFloats([]float64{1e-7}, 1) + in2, err := BesselIn(2, xt) + if err != nil { + t.Fatalf("BesselIn tiny: %v", err) + } + want := math.Pow(1e-7/2, 2) / 2 // I_2(x) ≈ x²/8 for tiny x + if math.Abs(in2.FloatAt(0)-want) > 1e-24 { + t.Fatalf("I_2(1e-7) = %g, want %g", in2.FloatAt(0), want) + } +} + +func isFinite(v float64) bool { return !math.IsNaN(v) && !math.IsInf(v, 0) } + +// TestSphericalHarmonicRealM0 pins the dtype contract: m = 0 must +// return a Float array like every other order. +func TestSphericalHarmonicRealM0(t *testing.T) { + theta, _ := FromFloats([]float64{0.3, 1.1}, 2) + phi, _ := FromFloats([]float64{0.5, 2.0}, 2) + out, err := SphericalHarmonicReal(2, 0, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonicReal: %v", err) + } + if out.Dtype() != Float { + t.Fatalf("dtype %s, want float", out.Dtype()) + } + for i := range 2 { + _ = out.FloatAt(i) // must not panic + } +} + +// TestMinMaxNaNOrder pins the seeding rule: a NaN anywhere must never +// win and the result must not depend on element order. +func TestMinMaxNaNOrder(t *testing.T) { + for _, vals := range [][]float64{ + {math.NaN(), 1, 2}, + {1, math.NaN(), 2}, + {1, 2, math.NaN()}, + } { + a, _ := FromFloats(vals, 3) + mn, err := Min(a) + if err != nil { + t.Fatalf("Min: %v", err) + } + mx, err := Max(a) + if err != nil { + t.Fatalf("Max: %v", err) + } + if mn.Float() != 1 || mx.Float() != 2 { + t.Fatalf("Min/Max of %v = (%g, %g), want (1, 2)", vals, mn.Float(), mx.Float()) + } + } + allNaN, _ := FromFloats([]float64{math.NaN(), math.NaN()}, 2) + mn, err := Min(allNaN) + if err != nil { + t.Fatalf("Min: %v", err) + } + if !math.IsNaN(mn.Float()) { + t.Fatalf("Min of all-NaN = %g, want NaN", mn.Float()) + } +} + +// TestJacobianComplexProbe pins the dtype guard on the probe outputs. +func TestJacobianComplexProbe(t *testing.T) { + f := func(x *Array) (*Array, error) { + return FromComplexes([]complex128{1i, 2}, 2) + } + x, _ := FromFloats([]float64{1, 2}, 2) + if _, err := Jacobian(f, x, JacobianOptions{}); err == nil { + t.Fatal("Jacobian accepted a complex probe output") + } +} diff --git a/internal/core/einsum_test.go b/internal/core/einsum_test.go new file mode 100644 index 0000000..2803f60 --- /dev/null +++ b/internal/core/einsum_test.go @@ -0,0 +1,256 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// TestEinsumGeneralReductions pins patterns only the general engine +// reaches: reduction-only axes, implicit output, arbitrary order. +func TestEinsumGeneralReductions(t *testing.T) { + a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + // "ij->i": row sums. + rs, err := Einsum("ij->i", a) + if err != nil { + t.Fatalf("ij->i: %v", err) + } + if rs.FloatAt(0) != 6 || rs.FloatAt(1) != 15 { + t.Fatalf("row sums = (%g, %g), want (6, 15)", rs.FloatAt(0), rs.FloatAt(1)) + } + // "ij->j": column sums. + cs, err := Einsum("ij->j", a) + if err != nil { + t.Fatalf("ij->j: %v", err) + } + for j, want := range []float64{5, 7, 9} { + if cs.FloatAt(j) != want { + t.Fatalf("col sum %d = %g, want %g", j, cs.FloatAt(j), want) + } + } + // Implicit output: "ij,jk" == "ij,jk->ik". + b, _ := FromFloats([]float64{1, 0, 0, 1, 2, -1}, 3, 2) + m1, err := Einsum("ij,jk", a, b) + if err != nil { + t.Fatalf("implicit: %v", err) + } + m2, _ := Einsum("ij,jk->ik", a, b) + if !sameValues(m1, m2) { + t.Fatal("implicit output differs from explicit") + } + // Output order "ij->ji" through the general engine as well. + tr, err := Einsum("ij->ji", a) + if err != nil { + t.Fatalf("ij->ji: %v", err) + } + if tr.FloatAt(0) != 1 || tr.FloatAt(1) != 4 || tr.FloatAt(2) != 2 { + t.Fatalf("transpose wrong: %v %v %v", tr.FloatAt(0), tr.FloatAt(1), tr.FloatAt(2)) + } +} + +func sameValues(a, b *Array) bool { + if a.Len() != b.Len() { + return false + } + for i := range a.Len() { + if a.Dtype() == Complex { + if a.ComplexAt(i) != b.ComplexAt(i) { + return false + } + } else if a.FloatAt(i) != b.FloatAt(i) { + return false + } + } + return true +} + +// TestEinsumEllipsis pins batched and broadcast patterns. +func TestEinsumEllipsis(t *testing.T) { + // Batched matmul "...ij,...jk->...ik" against the manual loop. + av := make([]float64, 2*3*4) + bv := make([]float64, 2*4*5) + for i := range av { + av[i] = float64(i%7) - 3 + } + for i := range bv { + bv[i] = float64(i%5) - 2 + } + a, _ := FromFloats(av, 2, 3, 4) + b, _ := FromFloats(bv, 2, 4, 5) + got, err := Einsum("...ij,...jk->...ik", a, b) + if err != nil { + t.Fatalf("ellipsis batched: %v", err) + } + if got.NDim() != 3 || got.Shape()[0] != 2 || got.Shape()[1] != 3 || got.Shape()[2] != 5 { + t.Fatalf("shape %v, want (2, 3, 5)", got.Shape()) + } + for bt := range 2 { + for i := range 3 { + for k := range 5 { + want := 0.0 + for j := range 4 { + want += a.FloatAt((bt*3+i)*4+j) * b.FloatAt((bt*4+j)*5+k) + } + if math.Abs(got.FloatAt((bt*3+i)*5+k)-want) > 1e-12 { + t.Fatalf("batched [%d,%d,%d] = %g, want %g", bt, i, k, got.FloatAt((bt*3+i)*5+k), want) + } + } + } + } + // Broadcast: (1, 4) x (2, 4) -> (2, ...) inner product per batch. + u, _ := FromFloats([]float64{1, 2, 3, 4}, 4) + v, _ := FromFloats([]float64{1, 0, 0, 0, 0, 1, 0, 0}, 2, 4) + dots, err := Einsum("i,...i->...", u, v) + if err != nil { + t.Fatalf("broadcast dots: %v", err) + } + if dots.Shape()[0] != 2 || dots.FloatAt(0) != 1 || dots.FloatAt(1) != 2 { + t.Fatalf("broadcast dots = %v %v", dots.FloatAt(0), dots.FloatAt(1)) + } +} + +// TestEinsumDiagonalGeneral pins repeated labels through the engine. +func TestEinsumDiagonalGeneral(t *testing.T) { + a, _ := FromFloats([]float64{1, 2, 3, 4}, 2, 2) + // "ii->i" still hits the fast path; "iij->ij" only the engine can. + b, _ := FromFloats([]float64{ + 1, 2, 3, 4, 5, 6, 7, 8, + }, 2, 2, 2) + out, err := Einsum("iij->ij", b) + if err != nil { + t.Fatalf("iij->ij: %v", err) + } + if out.Shape()[0] != 2 || out.Shape()[1] != 2 { + t.Fatalf("shape %v, want (2, 2)", out.Shape()) + } + // element [i][j] = b[i][i][j]. + for i := range 2 { + for j := range 2 { + want := b.FloatAt((i*2+i)*2 + j) + if out.FloatAt(i*2+j) != want { + t.Fatalf("out[%d][%d] = %g, want %g", i, j, out.FloatAt(i*2+j), want) + } + } + } + // Bad specs stay errors. + if _, err := Einsum("ij->jj", a); err == nil { + t.Fatal("repeated output label accepted") + } + if _, err := Einsum("ij->k", a); err == nil { + t.Fatal("unknown output label accepted") + } +} + +// TestEinsumSlotSplitWideWorkload pins the parallel slot walk's chunk +// cursor. The general engine hands every worker a contiguous chunk of +// output slots and each worker rebuilds the cursor of its own first slot +// from the slot index, so the walk must return the same bits however the +// slots are split. Both workloads hold thousands of slots, and the split +// is asserted before it is compared so a shrunken workload cannot +// quietly fall back to the single-chunk walk. +func TestEinsumSlotSplitWideWorkload(t *testing.T) { + const workers = 4 + cases := []struct { + spec string + shapes [][]int + sumTotal int // product of the summed labels' sizes + }{ + // 64·8·8·4 = 16384 slots and nothing summed, so the cursor is + // the only source of every operand offset. + {"ij,kl->ijkl", [][]int{{64, 8}, {8, 4}}, 1}, + // 64·16 = 1024 slots of 8·8·3 = 192 visits: the sum runs inside + // the worker that owns the slot. + {"ik,kj,jl->il", [][]int{{64, 8}, {8, 8}, {8, 16}}, 64}, + } + prev := NumWorkers() + defer SetNumCPU(prev) + for _, tc := range cases { + for _, dt := range []Dtype{Int, Float} { + operands := make([]*Array, len(tc.shapes)) + for i, shape := range tc.shapes { + operands[i] = einsumOperand(t, dt, shape...) + } + SetNumCPU(1) + want, err := Einsum(tc.spec, operands...) + if err != nil { + t.Fatalf("%s/%s: %v", tc.spec, dt, err) + } + // The dispatch splits only while a worker's chunk reaches + // the visit floor, so assert the workload still does. + perSlot := tc.sumTotal * len(operands) + minSlots := 1 + if perSlot < einsumSlotFloor { + minSlots = (einsumSlotFloor + perSlot - 1) / perSlot + } + slots := want.Len() + if chunk := (slots + workers - 1) / workers; chunk < minSlots { + t.Fatalf("%s/%s: %d slots no longer split at %d workers: chunk %d below the %d-slot floor", + tc.spec, dt, slots, workers, chunk, minSlots) + } + // A pure outer product is the product of one element of each + // operand, so its corners are checked against that definition: + // the comparison below cannot pass on a walk that is wrong in + // every chunk. + if tc.sumTotal == 1 { + for _, c := range [][4]int{{0, 0, 0, 0}, {7, 3, 5, 1}, {63, 7, 7, 3}} { + var x, y, g float64 + if dt == Int { + xi, _ := IntAt(operands[0], c[0], c[1]) + yi, _ := IntAt(operands[1], c[2], c[3]) + gi, _ := IntAt(want, c[0], c[1], c[2], c[3]) + x, y, g = float64(xi), float64(yi), float64(gi) + } else { + x, _ = FloatAt(operands[0], c[0], c[1]) + y, _ = FloatAt(operands[1], c[2], c[3]) + g, _ = FloatAt(want, c[0], c[1], c[2], c[3]) + } + if g != x*y { + t.Fatalf("%s/%s: slot %v = %v, want %v·%v = %v", + tc.spec, dt, c, g, x, y, x*y) + } + } + } + SetNumCPU(workers) + got, err := Einsum(tc.spec, operands...) + if err != nil { + t.Fatalf("%s/%s: split walk: %v", tc.spec, dt, err) + } + if !einsumBitsEqual(got, want) { + t.Fatalf("%s/%s: the split walk disagrees with the serial one at element %d", + tc.spec, dt, firstDifferingSlot(got, want)) + } + } + } +} + +// firstDifferingSlot returns the flat index of the first element on +// which two results disagree, or -1 when none does: a failure names one +// element instead of printing two whole payloads. +func firstDifferingSlot(a, b *Array) int { + if a.Len() != b.Len() || a.Dtype() != b.Dtype() { + return -1 + } + for i := range a.Len() { + switch a.dt { + case Int: + if a.ints[i] != b.ints[i] { + return i + } + case Float32: + if a.floats32[i] != b.floats32[i] { + return i + } + case Float: + if a.floats[i] != b.floats[i] { + return i + } + case Complex: + if a.complexes[i] != b.complexes[i] { + return i + } + } + } + return -1 +} diff --git a/internal/core/elementwise_view_test.go b/internal/core/elementwise_view_test.go new file mode 100644 index 0000000..a158bc9 --- /dev/null +++ b/internal/core/elementwise_view_test.go @@ -0,0 +1,219 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// TestElementwiseOpTakesTheStridedFallback pins the contiguity guard of +// elementwiseOp: operands carrying a stride table keep out of the raw +// payload kernels and take the accessor walk instead. No public +// constructor sets strides, so the guard is driven here directly: the +// view arrays carry their own stride tables and the unexported entry +// point is called in package. The result must read the strided elements +// in logical order and come back dense, and neither operand's storage +// may move. +func TestElementwiseOpTakesTheStridedFallback(t *testing.T) { + t.Run("same dtype, non-canonical strides", func(t *testing.T) { + // The strides gather rows 0 and 2 of a (3, 2) storage, so a + // payload-order read answers 1, 2, 3, 4 while the logical one + // answers 1, 2, 4, 5: the same-width raw-payload kernel must + // stay out of reach behind the dense gate, not just behind the + // dispatcher's stride disjunct. + a := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{1, 2, 3, 4, 5, 6}, strides: []int{3, 1}} + b := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{10, 20, 30, 40}} + out, err := elementwiseOp(a, b, "Add", pairAdd) + if err != nil { + t.Fatalf("elementwiseOp: %v", err) + } + want := []float64{11, 22, 34, 45} + for i, w := range want { + if got := out.FloatAt(i); got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + if out.Strided() { + t.Fatal("the result came back strided, want a dense payload") + } + }) + t.Run("mixed dtypes, rebased window", func(t *testing.T) { + // The int view gathers every second row of a (3, 2) storage, so + // its logical elements are 1, 2, 4, 5. Being the lower operand of + // the mixed pair, it is read through the accessor walk, which + // resolves its stride table. + ai := []int64{1, 2, 3, 4, 5, 6} + a := &Array{shape: []int{2, 2}, dt: Int, ints: ai, strides: []int{3, 1}} + b := mustFloats(t, []float64{10, 20, 30, 40}, 2, 2) + out, err := elementwiseOp(a, b, "Add", pairAdd) + if err != nil { + t.Fatalf("elementwiseOp: %v", err) + } + if out.Dtype() != Float { + t.Fatalf("elementwiseOp answered dtype %s, want float", out.Dtype()) + } + want := []float64{11, 22, 34, 45} + for i, w := range want { + if got := out.FloatAt(i); got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + if out.Strided() { + t.Fatal("the result came back strided, want a dense payload") + } + for i, v := range ai { + if v != int64(i+1) { + t.Fatalf("the operand's storage moved at %d", i) + } + } + }) +} + +// TestElementwiseStridedOperandsReadLogically pins the gate that keeps +// the accessor walk honest inside elementwise and elementwiseDiv: every +// raw-payload branch is reserved to dense operands, so a strided view +// reads in its logical order on every entry point that reaches the +// fallback, Div and the extrema family included. The view gathers rows +// 0 and 2 of a (3, 2) storage, so its logical elements are 1, 2, 4, 5 +// while a payload-order read answers 1, 2, 3, 4. +func TestElementwiseStridedOperandsReadLogically(t *testing.T) { + storage := []float64{1, 2, 3, 4, 5, 6} + view := &Array{shape: []int{2, 2}, dt: Float, floats: storage, strides: []int{3, 1}} + denseB := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{10, 20, 30, 40}} + want := []float64{11, 22, 34, 45} + + t.Run("add float64", func(t *testing.T) { + out, err := Add(view, denseB) + if err != nil { + t.Fatalf("Add: %v", err) + } + for i, w := range want { + if got := out.floats[i]; got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + }) + t.Run("minimum float64", func(t *testing.T) { + out, err := Minimum(view, denseB) + if err != nil { + t.Fatalf("Minimum: %v", err) + } + for i, w := range []float64{1, 2, 4, 5} { + if got := out.floats[i]; got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + }) + t.Run("div float32", func(t *testing.T) { + a := &Array{shape: []int{2, 2}, dt: Float32, floats32: []float32{1, 2, 3, 4, 5, 6}, strides: []int{3, 1}} + b := &Array{shape: []int{2, 2}, dt: Float32, floats32: []float32{4, 8, 16, 40}} + out, err := Div(a, b) + if err != nil { + t.Fatalf("Div: %v", err) + } + for i, w := range []float32{0.25, 0.25, 0.25, 0.125} { + if got := out.floats32[i]; got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + }) + t.Run("add float16", func(t *testing.T) { + halves := make([]uint16, len(storage)) + for i, v := range storage { + halves[i] = HalfFromFloat64(v) + } + a := &Array{shape: []int{2, 2}, dt: Float16, halves: halves, strides: []int{3, 1}} + bh := []uint16{HalfFromFloat64(10), HalfFromFloat64(20), HalfFromFloat64(30), HalfFromFloat64(40)} + b := &Array{shape: []int{2, 2}, dt: Float16, halves: bh} + out, err := Add(a, b) + if err != nil { + t.Fatalf("Add: %v", err) + } + for i, w := range want { + if got := HalfToFloat64(out.halves[i]); got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + }) + t.Run("add int64", func(t *testing.T) { + ints := []int64{1, 2, 3, 4, 5, 6} + a := &Array{shape: []int{2, 2}, dt: Int, ints: ints, strides: []int{3, 1}} + b := &Array{shape: []int{2, 2}, dt: Int, ints: []int64{10, 20, 30, 40}} + out, err := Add(a, b) + if err != nil { + t.Fatalf("Add: %v", err) + } + for i, w := range []int64{11, 22, 34, 45} { + if got := out.ints[i]; got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + for i, v := range ints { + if v != int64(i+1) { + t.Fatalf("the operand's storage moved at %d", i) + } + } + }) + t.Run("div float64", func(t *testing.T) { + a := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{1, 2, 3, 4, 5, 6}, strides: []int{3, 1}} + b := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{4, 8, 16, 40}} + out, err := Div(a, b) + if err != nil { + t.Fatalf("Div: %v", err) + } + for i, w := range []float64{0.25, 0.25, 0.25, 0.125} { + if got := out.floats[i]; got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + }) + t.Run("div float16", func(t *testing.T) { + halves := make([]uint16, 6) + for i, v := range []float64{1, 2, 3, 4, 5, 6} { + halves[i] = HalfFromFloat64(v) + } + a := &Array{shape: []int{2, 2}, dt: Float16, halves: halves, strides: []int{3, 1}} + bh := []uint16{HalfFromFloat64(4), HalfFromFloat64(8), HalfFromFloat64(16), HalfFromFloat64(40)} + b := &Array{shape: []int{2, 2}, dt: Float16, halves: bh} + out, err := Div(a, b) + if err != nil { + t.Fatalf("Div: %v", err) + } + for i, w := range []float64{0.25, 0.25, 0.25, 0.125} { + if got := HalfToFloat64(out.halves[i]); got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + }) + t.Run("add complex128", func(t *testing.T) { + complexes := []complex128{1 + 1i, 2 + 2i, 3 + 3i, 4 + 4i, 5 + 5i, 6 + 6i} + a := &Array{shape: []int{2, 2}, dt: Complex, complexes: complexes, strides: []int{3, 1}} + b := &Array{shape: []int{2, 2}, dt: Complex, complexes: []complex128{10, 20, 30, 40}} + out, err := Add(a, b) + if err != nil { + t.Fatalf("Add: %v", err) + } + for i, w := range []complex128{11 + 1i, 22 + 2i, 34 + 4i, 45 + 5i} { + if got := out.complexes[i]; got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + for i, v := range complexes { + if v != complex(float64(i+1), float64(i+1)) { + t.Fatalf("the operand's storage moved at %d", i) + } + } + }) + t.Run("div complex128", func(t *testing.T) { + a := &Array{shape: []int{2, 2}, dt: Complex, complexes: []complex128{1 + 1i, 2 + 2i, 3 + 3i, 4 + 4i, 5 + 5i, 6 + 6i}, strides: []int{3, 1}} + b := &Array{shape: []int{2, 2}, dt: Complex, complexes: []complex128{4, 8, 16, 40}} + out, err := Div(a, b) + if err != nil { + t.Fatalf("Div: %v", err) + } + for i, w := range []complex128{0.25 + 0.25i, 0.25 + 0.25i, 0.25 + 0.25i, 0.125 + 0.125i} { + if got := out.complexes[i]; got != w { + t.Fatalf("element %d = %v, want %v", i, got, w) + } + } + }) +} diff --git a/internal/core/elliptic.go b/internal/core/elliptic.go new file mode 100644 index 0000000..1edfa66 --- /dev/null +++ b/internal/core/elliptic.go @@ -0,0 +1,510 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Elliptic integrals, Jacobi elliptic functions and the Gauss +// hypergeometric function. The complete integrals use the +// arithmetic-geometric mean where it is exact (K directly, E through +// the companion series) and a high-order Gauss-Legendre product where +// it is not (Pi); the Jacobi functions invert the incomplete integral +// F by Newton iteration, which makes sn, cn and dn correct by +// construction: they are the sin, cos and dn of the amplitude whose +// integral is the argument. + +// ellipticF evaluates the incomplete integral F(φ, m) = ∫₀^φ dθ / √(1 +// − m·sin²θ) through Carlson's symmetric form +// +// F(φ, m) = sin φ · R_F(cos²φ, 1 − m·sin²φ, 1), +// +// which is exact for every m ∈ [0, 1], φ ∈ [−π/2, π/2]: the integrand's +// endpoint singularity at m tending to 1, φ = π/2 is a zero argument +// of R_F, a regular point of the duplication algorithm. The +// Gauss-Legendre rule this replaces lost accuracy there (2e-16 at +// m = 0.99, 3e-3 at m = 1−1e-6), because a product rule cannot follow a +// square-root singularity. +func ellipticF(phi, m float64) float64 { + if m == 0 { + return phi + } + // The integrand has period π and a full period integrates to 2K, + // so any φ reduces to the principal range: F(φ, m) = F(φ̂, m) + + // 2n·K(m) with φ̂ = φ − nπ the nearest point of [−π/2, π/2]. + n := math.Round(phi / math.Pi) + phi -= n * math.Pi + s, c := math.Sincos(phi) + f := s * carlsonRF(c*c, 1-m*s*s, 1) + if n != 0 { + f += 2 * n * EllipticKScalar(m) + } + return f +} + +// carlsonRF returns Carlson's symmetric elliptic integral of the first +// kind, R_F(x, y, z) = ½∫₀^∞ dt/√((t+x)(t+y)(t+z)), by the +// duplication algorithm (Numerical Recipes, 6.11): each pass averages +// the three arguments, and the series in the final deviations +// converges to double precision in a handful of passes. The arguments +// must be non-negative with at most one zero, which is the case for +// the calls made here. +func carlsonRF(x, y, z float64) float64 { + // The textbook cut-off is 0.0025, which leaves the third-order + // series in the deviations good to about 1e-10 relative; halving the + // deviations quadratically costs two extra passes and buys six + // digits, so the cut is pushed to 1e-9 here. + const ( + errtol = 1e-9 + c1 = 1.0 / 24 + c2 = 0.1 + c3 = 3.0 / 44 + c4 = 1.0 / 14 + ) + for range 100 { + sx, sy, sz := math.Sqrt(x), math.Sqrt(y), math.Sqrt(z) + lambda := sx*(sy+sz) + sy*sz + x = 0.25 * (x + lambda) + y = 0.25 * (y + lambda) + z = 0.25 * (z + lambda) + ave := (x + y + z) / 3 + delx := (ave - x) / ave + dely := (ave - y) / ave + delz := (ave - z) / ave + if math.Abs(delx) <= errtol && math.Abs(dely) <= errtol && math.Abs(delz) <= errtol { + e2 := delx*dely - delz*delz + e3 := delx * dely * delz + return (1 + (c1*e2-c2-c3*e3)*e2 + c4*e3) / math.Sqrt(ave) + } + } + // The loop above converges in under ten passes for every admissible + // input; the fallback keeps the contract of never looping for ever. + return math.NaN() +} + +// EllipticK returns the complete elliptic integral of the first kind +// K(m) = F(π/2, m), evaluated by the arithmetic-geometric mean: +// K(m) = π / (2·AGM(1, √(1−m))). The parameter convention is m = k². +// m must be below 1; as m tends to 1 it diverges +// and K returns +Inf there, m = 1 included as the limit only through +// the caller's rounding. +func EllipticK(m *Array) (*Array, error) { + return m.realFunc("EllipticK", func(v float64) float64 { + if math.IsNaN(v) { + return math.NaN() + } + if v >= 1 { + if v == 1 { + return math.Inf(1) + } + return math.NaN() + } + if v < 0 { + // Negative parameter: transform to a positive one, + // K(−s) = K(s/(1+s))/√(1+s). The complement 1 − s/(1+s) = + // 1/(1+s) is handed to the AGM directly: for s above 2^53 + // the float64 sum 1+s rounds to s, so the transformed + // parameter would round to exactly 1 and the AGM would + // return its round-off floor (~1.8e15) instead of the true + // K, which stays finite, and decays to 0, for every finite + // s. + s := -v + return agmKFromComplement(1/(1+s)) / math.Sqrt(1+s) + } + return EllipticKScalar(v) + }) +} + +// EllipticKScalar is K(m) at one point (0 ≤ m < 1) by the AGM. +func EllipticKScalar(m float64) float64 { + return agmKFromComplement(1 - m) +} + +// agmKFromComplement is K(1−ε) = π/(2·AGM(1, √ε)) evaluated from the +// complementary parameter ε = 1−m, the form the negative-parameter +// transform needs: there ε = 1/(1+s) stays exact where 1−ε would round +// to 1. +func agmKFromComplement(eps float64) float64 { + a, b := 1.0, math.Sqrt(eps) + // The stop test sits a few ulps above machine zero: rounding can + // pin a and b one ulp apart forever, and any threshold below that + // is an infinite loop, not extra accuracy. + for math.Abs(a-b) > 4*epsF*math.Max(1, a) { + a, b = 0.5*(a+b), math.Sqrt(a*b) + } + return math.Pi / (2 * a) +} + +// EllipticE returns the complete elliptic integral of the second kind +// E(m) = ∫₀^{π/2} √(1 − m·sin²θ) dθ through the AGM companion series +// E = K·(1 − Σ 2^{n−1} c_n²), with c_n² = a_n² − b_n² the AGM +// remainders. m = 1 gives 1; m > 1 is NaN; negative m transforms like +// K's does. +func EllipticE(m *Array) (*Array, error) { + return m.realFunc("EllipticE", ellipticEScalar) +} + +// ellipticEScalar is E(m) at one point. +func ellipticEScalar(m float64) float64 { + if math.IsNaN(m) { + return math.NaN() + } + if m == 1 { + return 1 + } + if m > 1 { + return math.NaN() + } + if m < 0 { + // E(−s) = √(1+s)·E(s/(1+s)). + s := -m + return math.Sqrt(1+s) * ellipticEScalar(s/(1+s)) + } + a, b := 1.0, math.Sqrt(1-m) + k := EllipticKScalar(m) + sum := 0.0 + pow2 := 0.5 // 2^{n-1} starting at n = 1 + for math.Abs(a-b) > 4*epsF*math.Max(1, a) { + c2 := a*a - b*b + sum += pow2 * c2 + pow2 *= 2 + a, b = 0.5*(a+b), math.Sqrt(a*b) + } + return k * (1 - sum) +} + +// EllipticPi returns the complete elliptic integral of the third kind +// Π(n, m) = ∫₀^{π/2} dθ / ((1 − n·sin²θ)√(1 − m·sin²θ)) element-wise +// over paired arrays, through Carlson's symmetric forms +// +// Π(n, m) = R_F(0, 1−m, 1) + (n/3)·R_J(0, 1−m, 1, 1−n), +// +// whose duplication algorithms are exact at the singular ends of the +// parameter square: full double precision for every n < 1 and m < 1, +// where the 64-point product rule this replaces lost seven digits at +// m = 1−1e-6 and more as n approached 1 (measured against mpmath: the +// worst relative error over the pinned table is 6e-16). At n = 1 and +// at m = 1 the integral diverges: the value there is +Inf, and beyond +// it NaN. +func EllipticPi(n, m *Array) (*Array, error) { + if !sameShape(n.shape, m.shape) { + return nil, errf("EllipticPi: shape mismatch %s vs %s", shapeText(n.shape), shapeText(m.shape)) + } + if n.dt == Complex || m.dt == Complex { + return nil, errf("EllipticPi: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, n.shape...), dt: Float} + out.alloc(n.Len()) + engine.Parallel(n.Len(), func(s, e int) { + for i := s; i < e; i++ { + out.floats[i] = ellipticPiScalar(n.floatAt(i), m.floatAt(i)) + } + }) + return out, nil +} + +// ellipticPiScalar is Π(n, m) at one point, through Carlson's +// symmetric forms: +// +// Π(n, m) = R_F(0, 1−m, 1) + (n/3)·R_J(0, 1−m, 1, 1−n). +// +// Both forms are exact at the singular ends of the parameter square, +// which is where the product rule this replaces lost its digits. +func ellipticPiScalar(n, m float64) float64 { + if math.IsNaN(n) || math.IsNaN(m) || m >= 1 || n >= 1 { + if m == 1 || n == 1 { + return math.Inf(1) + } + return math.NaN() + } + rf := carlsonRF(0, 1-m, 1) + if n == 0 { + return rf + } + return rf + n/3*carlsonRJ(0, 1-m, 1, 1-n) +} + +// carlsonRC returns Carlson's degenerate symmetric integral +// R_C(x, y) = R_F(x, y, y) for x ≥ 0 and y ≠ 0, by the duplication +// algorithm with the single-deviation series (Carlson 1995, +// "Numerical computation of real or complex elliptic integrals", +// (16)-(20)). A negative y is its Cauchy principal value, through the +// paper's (21); that case never arises for the calls made here, and is +// implemented so the function stands on its own. +func carlsonRC(x, y float64) float64 { + // r is the target relative truncation error; 1e-16 asks for double + // precision, and the seven-term series is good to that. + const r = 1e-16 + if y < 0 { + // (21) with the positive magnitude: R_C(x, −Y) = + // sqrt(x/(x+Y))·R_C(x+Y, Y). + return math.Sqrt(x/(x-y)) * carlsonRC(x-y, -y) + } + a0 := (x + 2*y) / 3 + q := math.Pow(3*r, -0.125) * math.Abs(a0-x) + a, xc, yc := a0, x, y + pow4 := 1.0 // 4^{−m} + for range 200 { + sx, sy := math.Sqrt(xc), math.Sqrt(yc) + lambda := 2*sx*sy + yc + an := (a + lambda) / 4 + pow4n := pow4 / 4 + if pow4n*q < math.Abs(an) { + s := (y - a0) / (an / pow4n) + s2 := s * s + series := 1 + (3.0/10)*s2 + (1.0/7)*s2*s + (3.0/8)*s2*s2 + + (9.0/22)*s2*s2*s + (159.0/208)*s2*s2*s2 + (9.0/8)*s2*s2*s2*s + return series / math.Sqrt(an) + } + xc, yc = (xc+lambda)/4, (yc+lambda)/4 + a, pow4 = an, pow4n + } + return math.NaN() +} + +// carlsonRJ returns Carlson's symmetric integral +// R_J(x, y, z, p) = (3/2)∫₀^∞ dt/((t+p)√((t+x)(t+y)(t+z))) for x, y, +// z ≥ 0 with at most one zero and p > 0, by the duplication theorem +// with the five-variable series (Carlson 1995, (24)-(32)): each pass +// averages the variables through lambda and accumulates the +// R_C(1, 1+e) correction the theorem contributes, then a sixth-order +// series in the deviations from the mean finishes the job. +func carlsonRJ(x, y, z, p float64) float64 { + const r = 1e-16 + if p <= 0 { + // The principal value for negative p is the paper's (33), + // which needs a permutation of x, y, z; none of the callers + // reaches it, so it is refused rather than approximated. + return math.NaN() + } + a0 := (x + y + z + 2*p) / 5 + delta := (p - x) * (p - y) * (p - z) + q := math.Pow(r/4, -1.0/6) * math.Max(math.Max(math.Abs(a0-x), math.Abs(a0-y)), + math.Max(math.Abs(a0-z), math.Abs(a0-p))) + a, xc, yc, zc, pc := a0, x, y, z, p + pow4 := 1.0 + total := 0.0 + for range 200 { + sx, sy, sz, sp := math.Sqrt(xc), math.Sqrt(yc), math.Sqrt(zc), math.Sqrt(pc) + lambda := sx*sy + sx*sz + sy*sz + an := (a + lambda) / 4 + pow4n := pow4 / 4 + // The m-th correction carries 4^{−3m}; pow4 is 4^{−m} already. + d := (sp + sx) * (sp + sy) * (sp + sz) + e := pow4 * pow4 * pow4 * delta / (d * d) + total += 6 * pow4 / d * carlsonRC(1, 1+e) + if pow4n*q < math.Abs(an) { + scaled := an / pow4n // 4^n·A_n + xx := (a0 - x) / scaled + yy := (a0 - y) / scaled + zz := (a0 - z) / scaled + pp := (-xx - yy - zz) / 2 + e2 := xx*yy + xx*zz + yy*zz - 3*pp*pp + e3 := xx*yy*zz + 2*e2*pp + 4*pp*pp*pp + e4 := (2*xx*yy*zz + e2*pp + 3*pp*pp*pp) * pp + e5 := xx * yy * zz * pp * pp + series := 1 - (3.0/14)*e2 + (1.0/6)*e3 + (9.0/88)*e2*e2 - + (3.0/22)*e4 - (9.0/52)*e2*e3 + (3.0/26)*e5 + return pow4n*series/(an*math.Sqrt(an)) + total + } + xc, yc, zc, pc = (xc+lambda)/4, (yc+lambda)/4, (zc+lambda)/4, (pc+lambda)/4 + a, pow4 = an, pow4n + } + return math.NaN() +} + +// EllipticFScalar is the incomplete elliptic integral of the first +// kind F(φ, m) = ∫₀^φ dθ/√(1−m·sin²θ) at one point, through Carlson's +// symmetric form; valid for every m ∈ [0, 1) and any real φ. The +// inverse view is the amplitude: u = F(φ, m) means φ = am(u, m), the +// reading behind the Jacobi functions. +func EllipticFScalar(phi, m float64) float64 { + return ellipticF(phi, m) +} + +// JacobiSN returns sn(u, m), the Jacobi elliptic sine, element-wise +// over u with the parameter m shared: sn inverts the incomplete +// integral u = F(φ, m) through φ, giving sn = sin φ. The same +// inversion serves cn and dn, which are the cos and the sqrt factor +// of the same amplitude. Valid for m ∈ [0, 1) and any real u, where +// m = 0 degenerates to the circular functions; m = 1 is refused +// because the amplitude would have to travel through the singular +// complete integral. +func JacobiSN(u *Array, m float64) (*Array, error) { + return jacobi(u, m, true, func(phi, mm float64) float64 { return math.Sin(phi) }) +} + +// JacobiCN returns cn(u, m) element-wise. +func JacobiCN(u *Array, m float64) (*Array, error) { + return jacobi(u, m, false, func(phi, mm float64) float64 { return math.Cos(phi) }) +} + +// JacobiDN returns dn(u, m) element-wise. +func JacobiDN(u *Array, m float64) (*Array, error) { + return jacobi(u, m, false, func(phi, mm float64) float64 { + return math.Sqrt(1 - mm*math.Sin(phi)*math.Sin(phi)) + }) +} + +// jacobi inverts u = F(φ, m) by a bracketed Newton iteration and maps +// the amplitude through f. +func jacobi(u *Array, m float64, odd bool, f func(phi, mm float64) float64) (*Array, error) { + if u.dt == Complex { + return nil, errf("Jacobi: complex arrays are not supported") + } + if math.IsNaN(m) || m < 0 || m >= 1 { + return nil, errf("Jacobi: the parameter m must lie in [0, 1), got %g", m) + } + // K scales the period; the amplitude search brackets φ in [0, hi] + // by doubling until F covers u. + k := EllipticKScalar(m) + out := &Array{shape: append([]int{}, u.shape...), dt: Float} + out.alloc(u.Len()) + engine.Parallel(u.Len(), func(s, e int) { + for i := s; i < e; i++ { + uu := u.floatAt(i) + // Sign symmetry first: work with |u|. + sign := 1.0 + if uu < 0 { + sign = -1 + uu = -uu + } + // Bracket: F(φ) ≥ φ·(1−m)^... the integrand is at least + // 1 on [0, π/2] and periodic beyond; hi = u + K covers + // every u ≥ 0 because F(u + K-margin)... doubling is the + // safe route. + hi := math.Max(uu, k) + for ellipticF(hi, m) < uu { + hi *= 2 + } + lo := 0.0 + phi := 0.5 * (lo + hi) + for range 200 { + // Newton from the midpoint of the current bracket, + // re-bracketed every step: F is strictly increasing. + val := ellipticF(phi, m) + if val < uu { + lo = phi + } else { + hi = phi + } + den := math.Sqrt(1 - m*math.Sin(phi)*math.Sin(phi)) + // F'(φ) = 1/den, so the Newton step multiplies by den. + next := phi + (uu-val)*den + if next <= lo || next >= hi { + next = 0.5 * (lo + hi) + } + if math.Abs(next-phi) < 1e-16*math.Max(1, math.Abs(phi)) { + phi = next + break + } + phi = next + } + if odd { + out.floats[i] = sign * f(phi, m) + } else { + out.floats[i] = f(phi, m) // dn is an even function + } + } + }) + return out, nil +} + +// JacobiCDScalar is cd(u, m) = cn(u, m)/dn(u, m) at one point, by the +// Gauss AGM: the descending Landen sequence converges quadratically +// and the amplitude folds back down the chain, so a handful of passes +// covers any u. The parameter convention is m = k². Valid for m ∈ [0, +// 1) and any real u, where m = 0 degenerates cd to cos. This is the +// function the elliptic rational function needs for the zeroes of a +// Cauer filter design. +func JacobiCDScalar(u, m float64) float64 { + if math.IsNaN(m) || m < 0 || m >= 1 || math.IsNaN(u) { + return math.NaN() + } + if m == 0 { + return math.Cos(u) + } + // The AGM chain: a ascends to the mean, b descends, c is half + // their gap and vanishes quadratically. + k := EllipticKScalar(m) + // cd has period 4K and antisymmetry about 2K: reduce to [0, 2K) + // and carry the sign. + sign := 1.0 + u = math.Mod(u, 4*k) + if u < 0 { + u += 4 * k + } + if u >= 2*k { + u -= 2 * k + sign = -1 + } + var a, b, c [64]float64 + a[0], b[0], c[0] = 1, math.Sqrt(1-m), math.Sqrt(m) + n := 0 + for n < 63 { + n++ + a[n] = 0.5 * (a[n-1] + b[n-1]) + b[n] = math.Sqrt(a[n-1] * b[n-1]) + c[n] = 0.5 * (a[n-1] - b[n-1]) + // The gap between a and b stops shrinking a few ulps above + // machine zero, so the stop test carries the same floor the + // complete integral's AGM uses; below it more passes would be + // an infinite loop, not extra accuracy. + if math.Abs(c[n]) <= 4*epsF*math.Abs(a[n]) { + break + } + } + // The amplitude climbs the chain, then folds back down it. + phi := float64(int(1)< 0; i-- { + phi = 0.5 * (phi + math.Asin(c[i]/a[i]*math.Sin(phi))) + } + return sign * math.Cos(phi) / math.Sqrt(1-m*math.Sin(phi)*math.Sin(phi)) +} + +// Hypergeometric2F1 returns the Gauss hypergeometric function +// ₂F₁(a, b; c; x) = Σ (a)_k (b)_k / (c)_k · x^k / k! element-wise +// over x, with a, b, c scalar parameters. The series converges for +// |x| < 1; a or b a non-positive integer terminates it as a +// polynomial. A non-positive integer c is an error. Outside the disc +// of convergence the result is the Gauss value at x = 1 when c > a+b +// (finite), the divergence limit +Inf when x = 1 and c ≤ a+b, and NaN +// for every other |x| ≥ 1. +func Hypergeometric2F1(a, b, c float64, x *Array) (*Array, error) { + if c == math.Trunc(c) && c <= 0 { + return nil, errf("Hypergeometric2F1: c must not be a non-positive integer, got %g", c) + } + terminated := (a == math.Trunc(a) && a <= 0) || (b == math.Trunc(b) && b <= 0) + return x.realFunc("Hypergeometric2F1", func(v float64) float64 { + if math.IsNaN(v) { + return math.NaN() + } + if math.Abs(v) >= 1 && !terminated { + if v == 1 { + // The Gauss value at x = 1 exists when c > a+b. + if c > a+b { + return math.Gamma(c) * math.Gamma(c-a-b) / (math.Gamma(c-a) * math.Gamma(c-b)) + } + return math.Inf(1) + } + return math.NaN() + } + term := 1.0 + sum := 1.0 + for k := 1; k <= 100000; k++ { + term *= (a + float64(k-1)) * (b + float64(k-1)) / (c + float64(k-1)) * v / float64(k) + sum += term + if math.Abs(term) < 1e-18*math.Abs(sum) { + break + } + if math.IsInf(sum, 0) { + return sum + } + } + return sum + }) +} diff --git a/internal/core/elliptic2_test.go b/internal/core/elliptic2_test.go new file mode 100644 index 0000000..6252ce7 --- /dev/null +++ b/internal/core/elliptic2_test.go @@ -0,0 +1,134 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// The third-kind integral and the Carlson forms behind it. Reference +// values are mpmath at 30 to 35 digits (elliprf, elliprj, elliprc, and +// Pi through the identity below, cross-checked against mpmath's +// quadrature); the library must reach double precision at the singular +// ends of the parameter square, where the product rule it replaced +// lost seven digits and worse. + +// TestEllipticPiAccuracy pins Π against mpmath. The covered corners +// are n and m approaching 1 together (where Π reaches 1e6), n +// approaching 1, m approaching 1, negative parameters, and the +// degenerate m = 0. +func TestEllipticPiAccuracy(t *testing.T) { + cases := []struct { + n, m, want float64 + }{ + {0, 0.5, 1.8540746773013719184}, + {0.5, 0.8, 3.4166601403243870137}, + {-1, 0.3, 1.1936018953043909136}, + {0.9, 0.9, 11.047747327040735532}, + {0.99, 0.999, 193.39638212991131552}, + {0.999, 0.5, 69.434652042115458236}, + {0.5, 0, 2.2214414690791831235}, + {-0.5, 0.5, 1.4878469926687983853}, + {0.9, -0.5, 4.2505802986876906623}, + {-2, -3, 0.68305896638359993325}, + {0.999999, 0.999999, 1000003.8969974163894}, + {0, 0.999999999, 11.747927296421043878}, + {0.3, 0.999999999, 16.301444430032469214}, + {0.99, 0.999999, 531.60692473385781712}, + } + for _, tc := range cases { + got := ellipticPiScalar(tc.n, tc.m) + if rel := math.Abs(got-tc.want) / math.Abs(tc.want); rel > 1e-14 { + t.Errorf("Π(%g, %g) = %.17g, mpmath says %.17g (relative %.2g)", + tc.n, tc.m, got, tc.want, rel) + } + } + // The identities: Π(0, m) = K(m), and the divergence contract. + n := mustFloats(t, []float64{0}, 1) + m := mustFloats(t, []float64{0.999999}, 1) + pi, err := EllipticPi(n, m) + if err != nil { + t.Fatalf("EllipticPi: %v", err) + } + if rel := math.Abs(pi.FloatAt(0)-EllipticKScalar(0.999999)) / EllipticKScalar(0.999999); rel > 1e-15 { + t.Errorf("Π(0, 0.999999) = %.17g, K = %.17g", pi.FloatAt(0), EllipticKScalar(0.999999)) + } + edgeN := mustFloats(t, []float64{1, 1.5}, 2) + edgeM := mustFloats(t, []float64{0.5, 0.5}, 2) + edge, err := EllipticPi(edgeN, edgeM) + if err != nil { + t.Fatalf("EllipticPi: %v", err) + } + if !math.IsInf(edge.FloatAt(0), 1) { + t.Errorf("Π(1, 0.5) = %v, want +Inf", edge.FloatAt(0)) + } + if !math.IsNaN(edge.FloatAt(1)) { + t.Errorf("Π(1.5, 0.5) = %v, want NaN", edge.FloatAt(1)) + } +} + +// TestCarlsonRJAccuracy pins R_J, including the arguments the third-kind +// identity hands it: a zero first argument and a small fourth one. +func TestCarlsonRJAccuracy(t *testing.T) { + cases := []struct { + x, y, z, p, want float64 + }{ + {1, 1, 1, 1, 1.0}, + {0, 0.5, 1, 0.5, 5.0832785087638745196}, + {0, 1e-12, 1, 1e-06, 22802707.378631573755}, + {1e-20, 1, 2, 3, 0.77688623771511264203}, + {0, 1, 1, 1e-16, 471238893.32608005743}, + {0.1, 0.2, 0.3, 0.9, 4.2398211507913140837}, + } + for _, tc := range cases { + got := carlsonRJ(tc.x, tc.y, tc.z, tc.p) + if rel := math.Abs(got-tc.want) / math.Abs(tc.want); rel > 1e-13 { + t.Errorf("R_J(%g, %g, %g, %g) = %.17g, mpmath says %.17g (relative %.2g)", + tc.x, tc.y, tc.z, tc.p, got, tc.want, rel) + } + } + // A negative fourth argument needs the principal value, which this + // implementation refuses rather than approximates. + if got := carlsonRJ(1, 1, 1, -1); !math.IsNaN(got) { + t.Errorf("R_J(1, 1, 1, −1) = %v, want NaN", got) + } +} + +// TestCarlsonRCAccuracy pins R_C, both argument orders and the Cauchy +// principal value for a negative second argument. +func TestCarlsonRCAccuracy(t *testing.T) { + cases := []struct { + x, y, want float64 + }{ + {1, 1, 1.0}, + {0, 1, 1.5707963267948966192}, + {1, 2, 0.78539816339744830962}, + {2, 1, 0.881373587019543025}, + {1.5, 0.5, 1.14621583478058884}, + {0.5, 1.5, 0.955316618124509278}, + {0.1, 10, 0.467396548061524389}, + {10, 0.1, 0.951308668352240085}, + {1, 0.25, 1.5206919926018927}, + {0.25, 1, 1.20919957615614523}, + {1e-30, 1, 1.5707963267948956192}, + {1, 1e-12, 14.508657738531223752}, + {1, -0.5, 0.93588131010357011049}, + {1, -3, 0.27465307216702742285}, + {0, -1, 0}, + } + for _, tc := range cases { + got := carlsonRC(tc.x, tc.y) + if tc.want == 0 { + if got != 0 { + t.Errorf("R_C(%g, %g) = %v, want 0", tc.x, tc.y, got) + } + continue + } + if rel := math.Abs(got-tc.want) / math.Abs(tc.want); rel > 1e-14 { + t.Errorf("R_C(%g, %g) = %.17g, mpmath says %.17g (relative %.2g)", + tc.x, tc.y, got, tc.want, rel) + } + } +} diff --git a/internal/core/elliptic_test.go b/internal/core/elliptic_test.go new file mode 100644 index 0000000..46b5d14 --- /dev/null +++ b/internal/core/elliptic_test.go @@ -0,0 +1,293 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// simpsonRef integrates f over [0, phi] with a dense composite +// Simpson rule: the independent oracle for the elliptic values. +func simpsonRef(f func(float64) float64, phi float64) float64 { + const n = 200000 + h := phi / float64(n) + s := f(0) + f(phi) + for i := 1; i < n; i++ { + if i%2 == 1 { + s += 4 * f(float64(i)*h) + } else { + s += 2 * f(float64(i)*h) + } + } + return s * h / 3 +} + +// TestEllipticK pins K against the closed form at m = 1/2, the AGM +// degenerate values, and Simpson. +func TestEllipticK(t *testing.T) { + k, err := EllipticK(mustFloats(t, []float64{0, 0.5, 0.9, -1.5}, 4)) + if err != nil { + t.Fatalf("EllipticK: %v", err) + } + if math.Abs(k.FloatAt(0)-math.Pi/2) > 1e-15 { + t.Fatalf("K(0) = %.15f, want π/2", k.FloatAt(0)) + } + // K(1/2) = Γ(1/4)² / (4√π). + want := math.Gamma(0.25) * math.Gamma(0.25) / (4 * math.Sqrt(math.Pi)) + if math.Abs(k.FloatAt(1)-want) > 1e-14 { + t.Fatalf("K(0.5) = %.15f, want %.15f", k.FloatAt(1), want) + } + // m = 0.9 against Simpson on the integrand. + f := func(th float64) float64 { return 1 / math.Sqrt(1-0.9*math.Sin(th)*math.Sin(th)) } + if math.Abs(k.FloatAt(2)-simpsonRef(f, math.Pi/2)) > 1e-11 { + t.Fatalf("K(0.9) = %.14f, Simpson says %.14f", k.FloatAt(2), simpsonRef(f, math.Pi/2)) + } + // Negative parameter against Simpson through its own integrand. + fn := func(th float64) float64 { return 1 / math.Sqrt(1+1.5*math.Sin(th)*math.Sin(th)) } + if math.Abs(k.FloatAt(3)-simpsonRef(fn, math.Pi/2)) > 1e-11 { + t.Fatalf("K(-1.5) = %.14f, Simpson says %.14f", k.FloatAt(3), simpsonRef(fn, math.Pi/2)) + } + inf, err := EllipticK(mustFloats(t, []float64{1}, 1)) + if err != nil || !math.IsInf(inf.FloatAt(0), 1) { + t.Fatalf("K(1) must be +Inf, got %v err %v", inf.FloatAt(0), err) + } +} + +// TestEllipticE pins E against Simpson and the AGM-series degenerate +// values, plus the negative-parameter transform. +func TestEllipticE(t *testing.T) { + e, err := EllipticE(mustFloats(t, []float64{0, 0.5, 0.99, -2.0, 1.0}, 5)) + if err != nil { + t.Fatalf("EllipticE: %v", err) + } + if math.Abs(e.FloatAt(0)-math.Pi/2) > 1e-15 { + t.Fatalf("E(0) = %.15f, want π/2", e.FloatAt(0)) + } + for i, m := range []float64{0.5, 0.99} { + f := func(th float64) float64 { return math.Sqrt(1 - m*math.Sin(th)*math.Sin(th)) } + if math.Abs(e.FloatAt(1+i)-simpsonRef(f, math.Pi/2)) > 1e-11 { + t.Fatalf("E(%g) = %.14f, Simpson says %.14f", m, e.FloatAt(1+i), simpsonRef(f, math.Pi/2)) + } + } + fn := func(th float64) float64 { return math.Sqrt(1 + 2.0*math.Sin(th)*math.Sin(th)) } + if math.Abs(e.FloatAt(3)-simpsonRef(fn, math.Pi/2)) > 1e-11 { + t.Fatalf("E(-2) = %.14f, Simpson says %.14f", e.FloatAt(3), simpsonRef(fn, math.Pi/2)) + } + if math.Abs(e.FloatAt(4)-1) > 1e-15 { + t.Fatalf("E(1) = %.15f, want 1", e.FloatAt(4)) + } +} + +// TestEllipticPi pins Pi against Simpson, including the Pi(0, m) = K +// identity. +func TestEllipticPi(t *testing.T) { + n := mustFloats(t, []float64{0, 0.5, -1.0}, 3) + m := mustFloats(t, []float64{0.5, 0.8, 0.3}, 3) + pi, err := EllipticPi(n, m) + if err != nil { + t.Fatalf("EllipticPi: %v", err) + } + k05 := EllipticKScalar(0.5) + if math.Abs(pi.FloatAt(0)-k05) > 1e-13 { + t.Fatalf("Π(0, 0.5) = %.14f, K(0.5) = %.14f", pi.FloatAt(0), k05) + } + for i, pair := range [][2]float64{{0.5, 0.8}, {-1.0, 0.3}} { + nn, mm := pair[0], pair[1] + f := func(th float64) float64 { + sq := math.Sin(th) + return 1 / ((1 - nn*sq*sq) * math.Sqrt(1-mm*sq*sq)) + } + if math.Abs(pi.FloatAt(1+i)-simpsonRef(f, math.Pi/2)) > 1e-11 { + t.Fatalf("Π(%g, %g) = %.14f, Simpson says %.14f", nn, mm, pi.FloatAt(1+i), simpsonRef(f, math.Pi/2)) + } + } +} + +// TestJacobiIdentities pins sn, cn, dn against their defining +// identities, degenerate parameters, the derivative relation and the +// inversion that defines them (F(am(u)) = u). +func TestJacobiIdentities(t *testing.T) { + us := mustFloats(t, []float64{0.3, 1.2, 2.7, -0.9, 5.5}, 5) + for _, m := range []float64{0.0, 0.37, 0.96} { + sn, err := JacobiSN(us, m) + if err != nil { + t.Fatalf("JacobiSN: %v", err) + } + cn, err := JacobiCN(us, m) + if err != nil { + t.Fatalf("JacobiCN: %v", err) + } + dn, err := JacobiDN(us, m) + if err != nil { + t.Fatalf("JacobiDN: %v", err) + } + for i := range 5 { + u := us.FloatAt(i) + s, c, d := sn.FloatAt(i), cn.FloatAt(i), dn.FloatAt(i) + if math.Abs(s*s+c*c-1) > 1e-12 { + t.Fatalf("m=%g u=%g: sn²+cn² = %.14f", m, u, s*s+c*c) + } + if math.Abs(m*s*s+d*d-1) > 1e-12 { + t.Fatalf("m=%g u=%g: m·sn²+dn² = %.14f", m, u, m*s*s+d*d) + } + if u < 0 && math.Abs(s+sn0(us, m, -u)) > 1e-13 { + t.Fatalf("m=%g: sn is not odd", m) + } + if u < 0 && math.Abs(d-dn0(us, m, -u)) > 1e-13 { + t.Fatalf("m=%g: dn is not even", m) + } + // d/du sn = cn·dn by central differences. + const h = 1e-6 + ups := mustFloats(t, []float64{u + h, u - h}, 2) + sp, _ := JacobiSN(ups, m) + der := (sp.FloatAt(0) - sp.FloatAt(1)) / (2 * h) + if math.Abs(der-c*d) > 1e-6 { + t.Fatalf("m=%g u=%g: sn' = %.8f, cn·dn = %.8f", m, u, der, c*d) + } + } + } + // Degenerate parameters: m = 0 circular, m = 1 hyperbolic. + us2 := mustFloats(t, []float64{0.4, 1.3}, 2) + sn, _ := JacobiSN(us2, 0) + cn, _ := JacobiCN(us2, 0) + if math.Abs(sn.FloatAt(0)-math.Sin(0.4)) > 1e-14 || math.Abs(cn.FloatAt(1)-math.Cos(1.3)) > 1e-14 { + t.Fatal("m = 0 must degenerate to sin/cos") + } + // The hyperbolic limit m tending to 1: sn approaches tanh from + // inside a boundary layer of width ~sqrt(1−m), hence the loose + // tolerance. + sn1, err := JacobiSN(us2, 1-1e-12) + if err != nil { + t.Fatalf("JacobiSN near m=1: %v", err) + } + if math.Abs(sn1.FloatAt(1)-math.Tanh(1.3)) > 1e-4 { + t.Fatalf("m tending to 1: sn(1.3) = %.12f, tanh = %.12f", sn1.FloatAt(1), math.Tanh(1.3)) + } + if _, err := JacobiSN(us2, 1); err == nil { + t.Fatal("m = 1 must be rejected (K diverges there)") + } + // The defining inversion: F(asin(sn(u, m)), m) = u. + // Valid below K(m) ~ 1.995: past it the amplitude passes π/2 and + // asin folds the branch away. + for _, u := range []float64{0.7, 1.5} { + ua := mustFloats(t, []float64{u}, 1) + snu, _ := JacobiSN(ua, 0.6) + phi := math.Asin(snu.FloatAt(0)) + if math.Abs(ellipticF(phi, 0.6)-u) > 1e-10 { + t.Fatalf("F(asin sn(%g)) = %.12f, want %g", u, ellipticF(phi, 0.6), u) + } + } + // Periodicity: sn(u + 4K) = sn(u). + k := EllipticKScalar(0.5) + uper := mustFloats(t, []float64{0.9 + 4*k}, 1) + uplain := mustFloats(t, []float64{0.9}, 1) + s1, _ := JacobiSN(uper, 0.5) + s2, _ := JacobiSN(uplain, 0.5) + if math.Abs(s1.FloatAt(0)-s2.FloatAt(0)) > 1e-10 { + t.Fatalf("periodicity broke: %.12f vs %.12f", s1.FloatAt(0), s2.FloatAt(0)) + } +} + +// sn0/dn0 are scalar convenience reads for the parity checks. +func sn0(us *Array, m, u float64) float64 { + one, _ := FromFloats([]float64{u}, 1) + s, err := JacobiSN(one, m) + if err != nil { + panic(err) + } + return s.FloatAt(0) +} + +func dn0(us *Array, m, u float64) float64 { + one, _ := FromFloats([]float64{u}, 1) + d, err := JacobiDN(one, m) + if err != nil { + panic(err) + } + return d.FloatAt(0) +} + +// TestHypergeometric2F1 pins 2F1 against its closed forms. +func TestHypergeometric2F1(t *testing.T) { + x := mustFloats(t, []float64{0, 0.3, 0.7, 0.97, -0.8, 1.0}, 6) + f, err := Hypergeometric2F1(1, 1, 2, x) + if err != nil { + t.Fatalf("Hypergeometric2F1: %v", err) + } + for i, v := range x.RawFloats() { + want := -math.Log(1-v) / v + if v == 0 { + want = 1 + } + if math.Abs(f.FloatAt(i)-want) > 1e-12 { + t.Fatalf("2F1(1,1;2;%g) = %.14f, want %.14f", v, f.FloatAt(i), want) + } + } + // 2F1(a, b; b; x) = (1−x)^{−a}. + g, err := Hypergeometric2F1(0.7, 3.5, 3.5, x) + if err != nil { + t.Fatalf("Hypergeometric2F1: %v", err) + } + for i, v := range x.RawFloats() { + if v == 1 { + continue + } + want := math.Pow(1-v, -0.7) + if math.Abs(g.FloatAt(i)-want) > 1e-12 { + t.Fatalf("2F1(a,b;b;%g) = %.14f, want %.14f", v, g.FloatAt(i), want) + } + } + // Terminating polynomial: 2F1(-3, 2; 1.5; x), built term by term below. + h, err := Hypergeometric2F1(-3, 2, 1.5, mustFloats(t, []float64{0.6}, 1)) + if err != nil { + t.Fatalf("Hypergeometric2F1: %v", err) + } + // k=0: 1; k=1: (−3·2/1.5)·0.6 = −2.4; k=2: (−3·−2·2·3/(1.5·2.5·2))·0.36; + // k=3: (−3·−2·−1·2·3·4/(1.5·2.5·3.5·6))·0.216. + poly := 1.0 + poly += (-3.0 * 2.0 / 1.5) * 0.6 + poly += (-3.0 * -2.0 * 2.0 * 3.0 / (1.5 * 2.5 * 2.0)) * 0.36 + poly += (-3.0 * -2.0 * -1.0 * 2.0 * 3.0 * 4.0 / (1.5 * 2.5 * 3.5 * 6.0)) * 0.216 + if math.Abs(h.FloatAt(0)-poly) > 1e-14 { + t.Fatalf("2F1(-3,2;1.5;0.6) = %.14f, want %.14f", h.FloatAt(0), poly) + } + // c a non-positive integer is an error. + if _, err := Hypergeometric2F1(1, 1, -1, x); err == nil { + t.Fatal("c = −1 accepted") + } +} + +// TestJacobiCDScalar pins the AGM evaluation against mpmath 3.16 +// reference values (ellipfun cn/dn at 30 digits) over the parameter +// range the Cauer filter design walks: m near both ends and u across +// several periods. +func TestJacobiCDScalar(t *testing.T) { + cases := []struct { + u, m, want float64 + }{ + {0.3, 0.5, 9.77250336444249856e-01}, + {1.0, 0.5, 7.24009721659370498e-01}, + {1.7, 0.9, 7.12053939073435838e-01}, + {2.5, 0.1, -7.69223750223807068e-01}, + {3.2, 0.7, -8.39726459768202815e-01}, + {5.0, 0.99, -8.64163207523496957e-01}, + {7.5, 0.6, 9.81944916682632396e-01}, + {12.0, 0.25, 1.98099136005263327e-01}, + {0.0, 0.4, 1}, + {1.0, 0.0, 5.40302305868139765e-01}, + } + for _, c := range cases { + got := JacobiCDScalar(c.u, c.m) + if math.Abs(got-c.want) > 5e-16*math.Max(1, math.Abs(c.want)) { + t.Fatalf("cd(%.3g, %.3g) = %.16g, want %.16g", c.u, c.m, got, c.want) + } + } + if v := JacobiCDScalar(1, math.NaN()); !math.IsNaN(v) { + t.Fatal("NaN parameter accepted") + } + if v := JacobiCDScalar(1, 1); !math.IsNaN(v) { + t.Fatal("m = 1 accepted") + } +} diff --git a/internal/core/expint.go b/internal/core/expint.go new file mode 100644 index 0000000..0dfffa0 --- /dev/null +++ b/internal/core/expint.go @@ -0,0 +1,164 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "math" + +// The exponential-integral and polygamma family. Element-wise over +// real arrays; ints and float32 promote, complex inputs are refused. + +// ExpIntegralE1 returns the exponential integral E1(x) = +// ∫ₓ^∞ e^(−t)/t dt of each element, defined for x > 0. A power series +// runs for x ≤ 2 and a continued fraction above, both accurate to +// near machine precision. +func ExpIntegralE1(a *Array) (*Array, error) { + return a.realFunc("ExpIntegralE1", expIntegralE1) +} + +func expIntegralE1(x float64) float64 { + if x <= 0 { + return math.NaN() + } + if x <= 2 { + // E1(x) = −γ − ln x − Σ_{k≥1} (−x)^k/(k·k!). + sum := 0.0 + term := 1.0 + for k := 1; k < 200; k++ { + term = -term * x / float64(k) + add := term / float64(k) + sum += add + if math.Abs(add) < 1e-18*math.Abs(sum) { + break + } + } + return -eulerGamma - math.Log(x) - sum + } + // Gauss continued fraction of the incomplete gamma function at + // a = 0, evaluated bottom-up from a generous depth: E1 = e^(-x)/f + // with f = x+1 - 1^2/(x+3 - 2^2/(x+5 - 3^2/(x+7 - ...))). The + // numerators grow quadratically, so a hundred levels leave the tail + // far below double rounding for every x in this branch. + const cfDepth = 100 + f := x + 2*float64(cfDepth) + 1 + for i := cfDepth; i >= 1; i-- { + f = (x + 2*float64(i) - 1) - float64(i*i)/f + } + return math.Exp(-x) / f +} + +// ExpIntegralEi returns the exponential integral Ei(x) of each +// element, defined for real x ≠ 0 by the Cauchy principal value. +// Negative x routes through E1; positive x uses the convergent series +// with the logarithmic singularity subtracted. +func ExpIntegralEi(a *Array) (*Array, error) { + return a.realFunc("ExpIntegralEi", func(x float64) float64 { + if x < 0 { + return -expIntegralE1(-x) + } + if x == 0 { + return math.Inf(-1) + } + if x >= 40 { + // Asymptotic expansion summed to its smallest term at + // k ≈ x. The convergent series needs k proportional to x + // (already past 300 terms at x = 250), so the old fixed + // cap silently truncated it; from x = 40 on, the omitted + // asymptotic tail sits below double rounding. The exp/log + // shift evaluates e^x/x without overflowing before the + // true value does. + t, sum := 1.0, 1.0 + for k := 1; k <= int(x); k++ { + t *= float64(k) / x + sum += t + } + return math.Exp(x-math.Log(x)) * sum + } + sum := 0.0 + term := 1.0 + for k := 1; k < 300; k++ { + term *= x / float64(k) + add := term / float64(k) + sum += add + if math.Abs(add) < 1e-18*math.Abs(sum) { + break + } + } + return eulerGamma + math.Log(x) + sum + }) +} + +// Digamma returns the digamma function ψ(x) = d/dx ln Γ(x) of each +// element. A recurrence lifts the argument above 20, then the +// asymptotic series applies; the reflection formula covers x < 0. +func Digamma(a *Array) (*Array, error) { + return a.realFunc("Digamma", digamma) +} + +func digamma(x float64) float64 { + if math.IsNaN(x) || math.IsInf(x, 0) { + return math.NaN() + } + result := 0.0 + if x <= 0 && x == math.Floor(x) { + return math.NaN() + } + if x < 0 { + // Reflection: ψ(1−x) = ψ(x) + π·cot(πx). + result = -math.Pi / math.Tan(math.Pi*x) + x = 1 - x + } + for x < 20 { + result -= 1 / x + x++ + } + inv := 1 / x + inv2 := inv * inv + result += math.Log(x) - 0.5*inv - inv2*(1.0/12-inv2*(1.0/120-inv2*(1.0/252-inv2*(1.0/240)))) + return result +} + +// Trigamma returns the trigamma function ψ′(x), the derivative of the +// digamma function, of each element. The same recurrence and +// asymptotic strategy as Digamma applies. +func Trigamma(a *Array) (*Array, error) { + return a.realFunc("Trigamma", trigamma) +} + +func trigamma(x float64) float64 { + if math.IsNaN(x) || math.IsInf(x, 0) { + return math.NaN() + } + result := 0.0 + reflected := false + if x <= 0 && x == math.Floor(x) { + return math.NaN() + } + if x < 0 { + // Reflection: ψ′(x) + ψ′(1−x) = π²/sin²(πx), so ψ′(x) is the + // constant minus the lifted value ψ′(1−x): the recurrence and + // the asymptotic below subtract instead of add, mirroring the + // digamma structure above. + result = math.Pi * math.Pi / (math.Sin(math.Pi*x) * math.Sin(math.Pi*x)) + x = 1 - x + reflected = true + } + for x < 20 { + if reflected { + result -= 1 / (x * x) + } else { + result += 1 / (x * x) + } + x++ + } + inv := 1 / x + inv2 := inv * inv + tail := inv*(1+0.5*inv) + inv*inv2*(1.0/6-inv2*(1.0/30-inv2*(1.0/42-inv2*(1.0/30)))) + if reflected { + return result - tail + } + return result + tail +} + +// eulerGamma is the Euler-Mascheroni constant. +const eulerGamma = 0.57721566490153286060651209008240243 diff --git a/internal/core/expint_test.go b/internal/core/expint_test.go new file mode 100644 index 0000000..fe2279a --- /dev/null +++ b/internal/core/expint_test.go @@ -0,0 +1,155 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// TestExpIntegralE1 checks E1 against tabulated values on both the +// series branch (x ≤ 2) and the continued-fraction branch (x > 2). +func TestExpIntegralE1(t *testing.T) { + x := mustFloats(t, []float64{0.1, 0.5, 1, 2, 5, 10}) + got, err := ExpIntegralE1(x) + if err != nil { + t.Fatalf("ExpIntegralE1: %v", err) + } + want := []float64{ + 1.8229239584193906, + 0.5597735947761609, + 0.21938393439552027, + 0.04890051070806112, + 0.0011482955912753257, + 4.156968929685324e-06, + } + for i := range want { + if math.Abs(got.FloatAt(i)-want[i]) > 1e-13*(1+math.Abs(want[i])) { + t.Fatalf("E1(%v) = %.16g, want %.16g", x.FloatAt(i), got.FloatAt(i), want[i]) + } + } + neg, nerr := ExpIntegralE1(mustFloats(t, []float64{-1})) + if nerr != nil { + t.Fatalf("ExpIntegralE1(−1): %v", nerr) + } + if !math.IsNaN(neg.FloatAt(0)) { + t.Fatalf("E1(−1) = %v, want NaN outside the domain", neg.FloatAt(0)) + } +} + +// TestExpIntegralEi checks Ei on both branches, including the +// reflection Ei(−x) = −E1(x) for x < 0. +func TestExpIntegralEi(t *testing.T) { + x := mustFloats(t, []float64{-1, -0.5, 0.5, 1, 5}) + got, err := ExpIntegralEi(x) + if err != nil { + t.Fatalf("ExpIntegralEi: %v", err) + } + want := []float64{ + -0.21938393439552027, + -0.5597735947761609, + 0.4542199048631725, + 1.8951178163559368, + 40.18527535580318, + } + for i := range want { + if math.Abs(got.FloatAt(i)-want[i]) > 1e-13*(1+math.Abs(want[i])) { + t.Fatalf("Ei(%v) = %.16g, want %.16g", x.FloatAt(i), got.FloatAt(i), want[i]) + } + } +} + +// TestDigamma checks the polygamma family against exact values: +// ψ(1) = −γ, ψ(2) = 1 − γ, ψ(½) = −γ − 2 ln 2, and the same for ψ′. +func TestDigamma(t *testing.T) { + x := mustFloats(t, []float64{0.5, 1, 2, 5}) + got, err := Digamma(x) + if err != nil { + t.Fatalf("Digamma: %v", err) + } + want := []float64{ + -eulerGamma - 2*math.Log(2), + -eulerGamma, + 1 - eulerGamma, + 1.5061176684318005, + } + for i := range want { + if math.Abs(got.FloatAt(i)-want[i]) > 1e-11*(1+math.Abs(want[i])) { + t.Fatalf("psi(%v) = %.16g, want %.16g", x.FloatAt(i), got.FloatAt(i), want[i]) + } + } + // Negative non-integer via the reflection formula: ψ(−½) = + // ψ(1.5) + π·cot(π/2) = ψ(1.5) = 0.03648997397857652. + neg := mustFloats(t, []float64{-0.5}) + got, err = Digamma(neg) + if err != nil { + t.Fatalf("Digamma(−½): %v", err) + } + if want := 0.03648997397857652; math.Abs(got.FloatAt(0)-want) > 1e-11 { + t.Fatalf("psi(−½) = %.16g, want %.16g", got.FloatAt(0), want) + } + // Poles at non-positive integers return NaN, the same IEEE + // convention the gamma family applies. + pole, perr := Digamma(mustFloats(t, []float64{-2})) + if perr != nil { + t.Fatalf("Digamma(−2): %v", perr) + } + if !math.IsNaN(pole.FloatAt(0)) { + t.Fatalf("psi(−2) = %v, want NaN at the pole", pole.FloatAt(0)) + } +} + +// TestTrigamma checks ψ′ against exact values: ψ′(1) = π²/6, +// ψ′(½) = π²/2, ψ′(2) = 1 − π²/6. +func TestTrigamma(t *testing.T) { + x := mustFloats(t, []float64{0.5, 1, 2, 5}) + got, err := Trigamma(x) + if err != nil { + t.Fatalf("Trigamma: %v", err) + } + want := []float64{ + math.Pi * math.Pi / 2, + math.Pi * math.Pi / 6, + math.Pi*math.Pi/6 - 1, + 0.22132295573718011, // π²/6 − (1 + ¼ + ¹⁄₉ + ¹⁄₁₆) + } + for i := range want { + if math.Abs(got.FloatAt(i)-want[i]) > 1e-11*(1+math.Abs(want[i])) { + t.Fatalf("psi'(%v) = %.16g, want %.16g", x.FloatAt(i), got.FloatAt(i), want[i]) + } + } +} + +// TestFresnel checks C and S against high-precision series values, +// the odd symmetry and the ½ limits at large argument. +func TestFresnel(t *testing.T) { + x := mustFloats(t, []float64{0.5, 1, 2, 4, 6}) + c, err := FresnelC(x) + if err != nil { + t.Fatalf("FresnelC: %v", err) + } + s, err := FresnelS(x) + if err != nil { + t.Fatalf("FresnelS: %v", err) + } + // References at 60-digit precision, all pinned at 1e-12: the plain + // power series covers the first three (x = 4 sits at its boundary) + // and x = 6 sums the same series in extended precision. + wantC := []float64{0.4923442258714464, 0.7798934003768228, 0.4882534060753408, 0.4984260330381776, 0.4995314678555011} + wantS := []float64{0.06473243286000028, 0.4382591473903548, 0.3434156783636982, 0.4205157542469284, 0.4469607612369303} + for i := range wantC { + if math.Abs(c.FloatAt(i)-wantC[i]) > 1e-12 { + t.Fatalf("C(%v) = %.16g, want %.16g", x.FloatAt(i), c.FloatAt(i), wantC[i]) + } + if math.Abs(s.FloatAt(i)-wantS[i]) > 1e-12 { + t.Fatalf("S(%v) = %.16g, want %.16g", x.FloatAt(i), s.FloatAt(i), wantS[i]) + } + } + // Odd symmetry. + nx := mustFloats(t, []float64{-1}) + nc, _ := FresnelC(nx) + if math.Abs(nc.FloatAt(0)+0.7798934003768228) > 1e-12 { + t.Fatalf("C(−1) = %v, want −C(1)", nc.FloatAt(0)) + } +} diff --git a/internal/core/extrema.go b/internal/core/extrema.go new file mode 100644 index 0000000..c166065 --- /dev/null +++ b/internal/core/extrema.go @@ -0,0 +1,213 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +// Element-wise extrema and clipping, plus the Go 1.27 generic +// payload accessor. Minimum and Maximum follow IEEE NaN propagation: a +// NaN operand poisons the result element; complex arrays have no +// ordering and error. + +// Minimum returns the element-wise smaller of two same-shape arrays; NaN +// propagates. +func Minimum(a, b *Array) (*Array, error) { + return elementwise(a, b, "Minimum", + func(x, y int64) int64 { return min(x, y) }, + func(x, y float64) float64 { return min(x, y) }, + nil) +} + +// Maximum returns the element-wise larger of two same-shape arrays; NaN +// propagates. +func Maximum(a, b *Array) (*Array, error) { + return elementwise(a, b, "Maximum", + func(x, y int64) int64 { return max(x, y) }, + func(x, y float64) float64 { return max(x, y) }, + nil) +} + +// ClipI clamps each element of a real array into [lo, hi]; the dtype +// keeps its kind. lo > hi is an error. +func ClipI(a *Array, lo, hi int64) (*Array, error) { + if lo > hi { + return nil, errf("ClipI: lo must be at most hi, got %d and %d", lo, hi) + } + switch a.dt { + case Int: + out := &Array{shape: a.Shape(), dt: Int, ints: make([]int64, a.Len())} + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.ints[s:e], out.ints[s:e] + for i := range os { + os[i] = min(max(as[i], lo), hi) + } + }) + return out, nil + case Float16: + flo, fhi := float64(lo), float64(hi) + out := &Array{shape: a.Shape(), dt: Float16, halves: make([]uint16, a.Len())} + if a.strides == nil { + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + // Clamped in float64, where the half values compare + // exactly, then narrowed once: half bit patterns do not + // order as uint16. + os[i] = HalfFromFloat64(min(max(HalfToFloat64(as[i]), flo), fhi)) + } + }) + return out, nil + } + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.halves[i] = HalfFromFloat64(min(max(a.halfAt(i), flo), fhi)) + } + }) + return out, nil + case Float32: + flo, fhi := float32(lo), float32(hi) + out := &Array{shape: a.Shape(), dt: Float32, floats32: make([]float32, a.Len())} + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = min(max(as[i], flo), fhi) + } + }) + return out, nil + case Float: + flo, fhi := float64(lo), float64(hi) + out := &Array{shape: a.Shape(), dt: Float, floats: make([]float64, a.Len())} + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats[s:e], out.floats[s:e] + for i := range os { + os[i] = min(max(as[i], flo), fhi) + } + }) + return out, nil + case Complex: + return nil, errf("ClipI: complex arrays have no ordering") + default: + // Bool and the narrow integer widths carry no clamp kernel here; + // the message names the actual dtype rather than complex. + return nil, errf("ClipI: dtype %s is not supported; convert with Astype", a.dt) + } +} + +// ClipF clamps each element of a real array into [lo, hi]; int arrays +// become float, float16 and float32 arrays keep their dtype with the +// clamp run in float64 and narrowed once. lo > hi is an error. +func ClipF(a *Array, lo, hi float64) (*Array, error) { + if lo > hi { + return nil, errf("ClipF: lo must be at most hi, got %v and %v", lo, hi) + } + if a.dt == Complex { + return nil, errf("ClipF: complex arrays have no ordering") + } + if narrowRefused(a.dt) { + // Bool and the narrow integer widths carry no clamp kernel here; + // the refusal names the dtype and the conversion. + return nil, errf("ClipF: dtype %s is not supported; convert with Astype", a.dt) + } + out := &Array{shape: a.Shape(), dt: a.dt} + if a.dt == Int { + out.dt = Float + } + out.alloc(a.Len()) + switch a.dt { + case Float16: + if a.strides == nil { + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(min(max(HalfToFloat64(as[i]), lo), hi)) + } + }) + } else { + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.halves[i] = HalfFromFloat64(min(max(a.halfAt(i), lo), hi)) + } + }) + } + case Float32: + flo, fhi := float32(lo), float32(hi) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = min(max(as[i], flo), fhi) + } + }) + case Int: + if a.strides == nil { + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.ints[s:e], out.floats[s:e] + for i := range os { + os[i] = min(max(float64(as[i]), lo), hi) + } + }) + } else { + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.floats[i] = min(max(a.floatAt(i), lo), hi) + } + }) + } + default: + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats[s:e], out.floats[s:e] + for i := range os { + os[i] = min(max(as[i], lo), hi) + } + }) + } + return out, nil +} + +// Elements returns the flat payload converted to the caller's type, a +// Go 1.27 generic method. E may hold anything the payload ladder +// widens to: the ladder runs int64 ← float16 ← float32 ← float64 ← +// complex128, and requesting a rung wider than the array's converts +// while a narrower one is an error naming both dtypes. E = int64 +// follows the rule IntAt carries: the whole integer class, bool +// included, widens exactly through intAt, and a float or complex +// source has no exact int64 image, so it keeps the narrowing refusal. +// The returned slice is a copy; changing it never reaches the array. +func (a *Array) Elements[E int64 | float32 | float64 | complex128]() ([]E, error) { + out := make([]E, a.Len()) + var zero E + switch any(zero).(type) { + case complex128: + // Every dtype converts to complex. + for i := range a.Len() { + out[i] = any(a.complexAt(i)).(E) + } + case float64: + // int, float32 and float64 convert; complex is a narrowing. + if a.dt == Complex { + return nil, errf("Elements: cannot narrow %s to float64", a.dt) + } + for i := range a.Len() { + out[i] = any(a.floatAt(i)).(E) + } + case float32: + // int and float32 convert; float64 and complex are narrowings. + if a.dt == Float || a.dt == Complex { + return nil, errf("Elements: cannot narrow %s to float32", a.dt) + } + for i := range a.Len() { + out[i] = any(a.float32At(i)).(E) + } + default: + // int64: the integer class widens exactly through intAt, the + // same widening IntAt answers; float and complex sources are + // genuine narrowings and keep the refusal. intAt gathers a + // strided array through physIndex, exactly as the float32 + // branch's accessor does, instead of reading the payload raw. + if !intClass(a.dt) { + return nil, errf("Elements: cannot narrow %s to int64", a.dt) + } + for i := range a.Len() { + out[i] = any(a.intAt(i)).(E) + } + } + return out, nil +} diff --git a/internal/core/extrema_test.go b/internal/core/extrema_test.go new file mode 100644 index 0000000..7e5b4c3 --- /dev/null +++ b/internal/core/extrema_test.go @@ -0,0 +1,187 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +func TestMinimumMaximum(t *testing.T) { + a := mustFromInts(t, []int64{1, 5, 3}, 3) + b := mustFromInts(t, []int64{4, 2, 3}, 3) + + mn, err := Minimum(a, b) + if err != nil { + t.Fatalf("Minimum: %v", err) + } + if !Equal(mustFromInts(t, []int64{1, 2, 3}, 3), mn) { + t.Fatalf("Minimum: %s", mn) + } + mx, err := Maximum(a, b) + if err != nil { + t.Fatalf("Maximum: %v", err) + } + if !Equal(mustFromInts(t, []int64{4, 5, 3}, 3), mx) { + t.Fatalf("Maximum: %s", mx) + } + + // Promotion and NaN propagation. + f := mustFromFloats(t, []float64{1.0, math.NaN()}, 2) + fb := mustFromFloats(t, []float64{0.5, 1.0}, 2) + fmin, _ := Minimum(f, fb) + if v, _ := FloatAt(fmin, 0); v != 0.5 { + t.Fatalf("Minimum promote: %v", v) + } + if v, _ := FloatAt(fmin, 1); !math.IsNaN(v) { + t.Fatalf("Minimum NaN must propagate: %v", v) + } + i := mustFromInts(t, []int64{1}, 1) + mixed, _ := Minimum(i, mustFromFloats(t, []float64{0.5}, 1)) + if mixed.Dtype() != Float { + t.Fatalf("Minimum promote dtype: %s", mixed.Dtype()) + } + + c := mustFromComplexes(t, []complex128{1}, 1) + if _, err := Minimum(c, c); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("Minimum complex: %v", err) + } + if _, err := Maximum(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "shape mismatch") { + t.Fatalf("Maximum shape: %v", err) + } +} + +func TestClip(t *testing.T) { + a := mustFromInts(t, []int64{-5, 3, 99}, 3) + + cli, err := ClipI(a, 0, 10) + if err != nil { + t.Fatalf("ClipI: %v", err) + } + if !Equal(mustFromInts(t, []int64{0, 3, 10}, 3), cli) { + t.Fatalf("ClipI: %s", cli) + } + if cli.Dtype() != Int { + t.Fatalf("ClipI keeps int: %s", cli.Dtype()) + } + + clf, err := ClipF(a, -1.5, 1.5) + if err != nil { + t.Fatalf("ClipF: %v", err) + } + if clf.Dtype() != Float { + t.Fatalf("ClipF dtype: %s", clf.Dtype()) + } + if v, _ := FloatAt(clf, 0); v != -1.5 { + t.Fatalf("ClipF lo: %v", v) + } + if v, _ := FloatAt(clf, 2); v != 1.5 { + t.Fatalf("ClipF hi: %v", v) + } + + f := mustFromFloats(t, []float64{0.5, 2.5}, 2) + fc, _ := ClipI(f, 1, 2) + if v, _ := FloatAt(fc, 0); v != 1 { + t.Fatalf("ClipI on float: %v", v) + } + + if _, err := ClipI(a, 5, 0); err == nil || !strings.Contains(err.Error(), "lo must be at most hi") { + t.Fatalf("ClipI range: %v", err) + } + if _, err := ClipF(a, 2, 1); err == nil || !strings.Contains(err.Error(), "lo must be at most hi") { + t.Fatalf("ClipF range: %v", err) + } + c := mustFromComplexes(t, []complex128{1}, 1) + if _, err := ClipI(c, 0, 1); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("ClipI complex: %v", err) + } + if _, err := ClipF(c, 0, 1); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("ClipF complex: %v", err) + } +} + +func TestElementsGeneric(t *testing.T) { + i := mustFromInts(t, []int64{1, 2}, 2) + + // Same-type access returns the values unchanged. + ints, err := i.Elements[int64]() + if err != nil || ints[0] != 1 || ints[1] != 2 { + t.Fatalf("Elements[int64]: %v %v", ints, err) + } + + // Widening converts along the ladder. + floats, err := i.Elements[float64]() + if err != nil || floats[0] != 1 || floats[1] != 2 { + t.Fatalf("Elements[float64]: %v %v", floats, err) + } + complexes, err := i.Elements[complex128]() + if err != nil || complexes[1] != complex(2, 0) { + t.Fatalf("Elements[complex128]: %v %v", complexes, err) + } + + // Float arrays widen to complex; narrowing errors. + f := mustFromFloats(t, []float64{2.5}, 1) + fc, err := f.Elements[complex128]() + if err != nil || fc[0] != complex(2.5, 0) { + t.Fatalf("Elements float to complex: %v %v", fc, err) + } + if _, err := f.Elements[int64](); err == nil || !strings.Contains(err.Error(), "cannot narrow float to int64") { + t.Fatalf("Elements float to int64: %v", err) + } + + c := mustFromComplexes(t, []complex128{complex(1, 2)}, 1) + if _, err := c.Elements[float64](); err == nil || !strings.Contains(err.Error(), "cannot narrow complex to float64") { + t.Fatalf("Elements complex to float64: %v", err) + } + cc, err := c.Elements[complex128]() + if err != nil || cc[0] != complex(1, 2) { + t.Fatalf("Elements[complex128] on complex: %v %v", cc, err) + } + + // The returned slice is a copy. + vals, _ := i.Elements[int64]() + vals[0] = 99 + if v, _ := IntAt(i, 0); v != 1 { + t.Fatalf("Elements must copy: %d", v) + } + + // The integer class widens to int64 exactly, the rule IntAt + // carries: a narrow source is an exact widening, never a + // narrowing refusal. + n8, err := FromInt8s([]int8{-128, -1, 0, 1, 127}, 5) + if err != nil { + t.Fatalf("FromInt8s: %v", err) + } + nv, err := n8.Elements[int64]() + if err != nil { + t.Fatalf("Elements[int64] on int8: %v", err) + } + for j, w := range []int64{-128, -1, 0, 1, 127} { + if nv[j] != w { + t.Fatalf("Elements[int64] int8[%d] = %d, want %d", j, nv[j], w) + } + } + bs, err := FromBools([]bool{true, false, true}, 3) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + bi, err := bs.Elements[int64]() + if err != nil || bi[0] != 1 || bi[1] != 0 || bi[2] != 1 { + t.Fatalf("Elements[int64] on bool = %v, %v; want [1 0 1]", bi, err) + } + u32, err := FromUint32s([]uint32{0, 4294967295}, 2) + if err != nil { + t.Fatalf("FromUint32s: %v", err) + } + ui, err := u32.Elements[int64]() + if err != nil || ui[0] != 0 || ui[1] != 4294967295 { + t.Fatalf("Elements[int64] on uint32 = %v, %v; want [0 4294967295] exact", ui, err) + } + // A complex source has no exact int64 image: the genuine + // narrowing keeps its refusal. + if _, err := c.Elements[int64](); err == nil || !strings.Contains(err.Error(), "cannot narrow complex to int64") { + t.Fatalf("Elements complex to int64: %v", err) + } +} diff --git a/internal/core/float16.go b/internal/core/float16.go new file mode 100644 index 0000000..d8b412a --- /dev/null +++ b/internal/core/float16.go @@ -0,0 +1,183 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/bits" +) + +// IEEE 754 binary16 (half precision) as a fifth dtype. The payload is a +// []uint16 holding the half bit patterns, exactly as the design note +// prescribes: a payload per dtype, and kernels that widen on read. The +// half dtype sits on the promotion ladder one step below float32 and +// behaves like it: the arithmetic runs in float64 (where every half +// value is exact) and narrows once per element. +// +// Conversion contract, for both directions: +// - widening (HalfToFloat64) is exact: every finite half value, +// subnormals included, is a float64 value; infinities widen to +// infinities; a half NaN widens to a float64 NaN carrying the same +// sign with its ten payload bits at the top of the mantissa. +// - narrowing (HalfFromFloat64) rounds to nearest, ties to even +// (IEEE 754 round-to-nearest-even), with gradual underflow through +// the half subnormals. |x| >= 65520 overflows to the signed +// infinity (65520 is the midpoint between the largest finite half +// 65504 and the next binade at 65536, and the tie rounds away from +// the odd mantissa of 65504); anything smaller rounds to a finite +// half. A float64 NaN, quiet or signalling, narrows to the +// canonical half NaN 0x7E00 with its sign preserved: the payload is +// deliberately not carried across, which keeps the conversion a +// pure function of the value and makes the constructor's NaN +// behaviour trivial to reason about. Signed zero is preserved. +// +// Constructors built on the narrowing (FromFloat16s, FullF16, SetFloatAt +// on a half array) inherit exactly this contract. + +// Half bit patterns the package writes directly. +const ( + // halfOne is the bit pattern of 1.0. + halfOne uint16 = 0x3C00 + // halfNaN is the canonical quiet NaN. + halfNaN uint16 = 0x7E00 + // halfInf is +Inf; OR-ing the sign bit turns it into -Inf. + halfInf uint16 = 0x7C00 +) + +// HalfToFloat64 widens the half bit pattern h to float64 exactly: no +// rounding, no range loss, NaN payloads preserved in the top mantissa +// bits. +func HalfToFloat64(h uint16) float64 { + sign := uint64(h&0x8000) << 48 + exp := uint64(h>>10) & 0x1F + frac := uint64(h & 0x03FF) + switch { + case exp == 0x1F: + if frac == 0 { + return math.Float64frombits(sign | 0x7FF0_0000_0000_0000) + } + // NaN: the ten payload bits land above the float64 mantissa's + // low 42 bits, so the quiet bit travels with the payload. + return math.Float64frombits(sign | 0x7FF0_0000_0000_0000 | frac<<42) + case exp == 0: + if frac == 0 { + return math.Float64frombits(sign) // ±0 + } + // Subnormal: the value is frac × 2^-24, renormalised into the + // float64 field. k is the zero-based position of frac's leading + // bit, so the value is 2^(k-24) × (1 + r/2^k) with r below. + k := uint(bits.Len64(frac) - 1) + mant := (frac - 1<= 65520, NaN canonicalised to +// 0x7E00 (sign preserved), signed zero preserved. +func HalfFromFloat64(f float64) uint16 { + b := math.Float64bits(f) + sign := uint16(b>>63) << 15 + exp := int64(b>>52) & 0x7FF + frac := b & (1<<52 - 1) + if exp == 0x7FF { + if frac != 0 { + return sign | halfNaN + } + return sign | halfInf + } + // The half exponent field the value would carry if it fit: float64 + // bias 1023 against half bias 15. + e := exp - 1008 + if e >= 0x1F { + // Above every finite half: ±Inf. + return sign | halfInf + } + if e >= 1 { + // A normal half: round the 52-bit mantissa to 10 bits, ties to + // even, then carry a mantissa overflow into the exponent (which + // may itself overflow into the infinity boundary). + hf := frac >> 42 + rem := frac & (1<<42 - 1) + if rem > 1<<41 || (rem == 1<<41 && hf&1 == 1) { + hf++ + } + if hf == 1<<10 { + hf = 0 + e++ + if e >= 0x1F { + return sign | halfInf + } + } + return sign | uint16(e)<<10 | uint16(hf) + } + // Subnormal (or zero): the half mantissa m counts 2^-24 units. The + // value is num × 2^(exp-1075) with num the mantissa plus its + // implicit leading bit, so m is num shifted right by 1051-exp, the + // guard and sticky bits deciding the tie. A shift of 53 is the + // smallest that can still round up: f = 2^-25 shifts by exactly 53 + // and ties to the even zero. Anything further is below the tie and + // drops out; float64 subnormals (exp 0) never reach even that. + sh := uint(1051 - exp) + if sh > 54 { + return sign + } + num := 1<<52 | frac + m := num >> sh + rem := num & (1< halfPoint || (rem == halfPoint && m&1 == 1) { + m++ + } + if m == 1<<10 { + // The subnormal rounding crossed into the smallest normal. + return sign | 1<<10 + } + return sign | uint16(m) +} + +// FromFloat16s builds a float16 array of the given shape from vals, +// copying them and narrowing each value to the nearest half +// (round-to-nearest-even, overflow to ±Inf, NaN canonicalised to +// 0x7E00, signed zero preserved: the HalfFromFloat64 contract). +func FromFloat16s(vals []float64, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + halves := make([]uint16, len(vals)) + for i, v := range vals { + halves[i] = HalfFromFloat64(v) + } + return &Array{shape: sh, dt: Float16, halves: halves}, nil +} + +// HalvesFromArray builds a float16 array that takes ownership of the +// raw half bit patterns in vals: no copy is made and no value is +// converted, so the caller must not touch the slice afterwards. The +// value count must fill the shape exactly. This is the bits-taking +// route for callers that already hold IEEE 754 binary16 data. +func HalvesFromArray(vals []uint16, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Float16, halves: vals}, nil +} + +// RawHalves returns the float16 payload (the raw IEEE 754 binary16 bit +// patterns) with the RawFloats contract. +func (a *Array) RawHalves() []uint16 { return a.halves } + +// halfAt returns element i widened from the half payload, the float16 +// twin of floatAt. +func (a *Array) halfAt(i int) float64 { + if a.strides != nil { + i = a.physIndex(i) + } + return HalfToFloat64(a.halves[i]) +} diff --git a/internal/core/float16_test.go b/internal/core/float16_test.go new file mode 100644 index 0000000..c37254b --- /dev/null +++ b/internal/core/float16_test.go @@ -0,0 +1,978 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/bits" + "strings" + "testing" +) + +func mustFromFloat16s(t *testing.T, vals []float64, shape ...int) *Array { + t.Helper() + a, err := FromFloat16s(vals, shape...) + if err != nil { + t.Fatalf("FromFloat16s(%v, %v): %v", vals, shape, err) + } + return a +} + +func mustFromHalves(t *testing.T, halves []uint16, shape ...int) *Array { + t.Helper() + a, err := HalvesFromArray(append([]uint16(nil), halves...), shape...) + if err != nil { + t.Fatalf("HalvesFromArray(%v, %v): %v", halves, shape, err) + } + return a +} + +// TestHalfConversionTable pins the exact bit patterns of the narrowing +// under IEEE 754 round-to-nearest-even. Each expectation is derived from +// the standard: the half format has a 5-bit exponent (bias 15) and a +// 10-bit fraction, the largest finite half is 65504 (0x7BFF), the +// overflow boundary for RNE is |x| = 65520 (the midpoint between 65504 +// and the next binade at 65536; the tie rounds away from 65504's odd +// mantissa), and the smallest subnormal is 2^-24 (0x0001). +func TestHalfConversionTable(t *testing.T) { + cases := []struct { + in float64 + want uint16 + why string + }{ + {0, 0x0000, "zero"}, + {math.Copysign(0, -1), 0x8000, "negative zero is preserved"}, + {0.5, 0x3800, "exponent field 14, fraction 0"}, + {-0.5, 0xB800, "sign preserved"}, + {1, 0x3C00, "exponent field 15, fraction 0"}, + {2, 0x4000, "exponent field 16"}, + {-2, 0xC000, "sign preserved"}, + {0.1, 0x2E66, "1.6 × 2^-4: fraction 0.6 × 1024 = 614.4 rounds down"}, + {1 + math.Ldexp(1, -10), 0x3C01, "one ulp above 1"}, + {1 + math.Ldexp(1, -11), 0x3C00, "exact tie between 1 and 1+2^-10: even mantissa 0 wins"}, + {2048, 0x6800, "exponent field 23"}, + {2049, 0x6800, "exact tie between 2048 and 2050: even mantissa 0 wins"}, + {2049.9, 0x6801, "closer to 2050"}, + {math.Ldexp(1, -14), 0x0400, "the smallest normal"}, + {math.Ldexp(1, -15), 0x0200, "subnormal with fraction field 512"}, + {math.Ldexp(1, -24), 0x0001, "the smallest subnormal"}, + {math.Ldexp(1, -25), 0x0000, "halfway between 0 and 2^-24: ties to the even zero"}, + {1.5 * math.Ldexp(1, -24), 0x0002, "halfway between 2^-24 and 2^-23: even mantissa 2 wins"}, + {0.9999 * math.Ldexp(1, -14), 0x0400, "subnormal rounding crosses into the smallest normal"}, + {65504, 0x7BFF, "the largest finite half"}, + {65512, 0x7BFF, "below the overflow midpoint"}, + {65519.999, 0x7BFF, "just below the midpoint still rounds to 65504"}, + {65520, 0x7C00, "the midpoint itself: the tie leaves the odd mantissa 1023 for infinity"}, + {65521, 0x7C00, "beyond the midpoint"}, + {-65504, 0xFBFF, "sign preserved at the maximum"}, + {-65520, 0xFC00, "negative overflow to -Inf"}, + {math.Inf(1), 0x7C00, "+Inf"}, + {math.Inf(-1), 0xFC00, "-Inf"}, + {math.NaN(), 0x7E00, "every NaN narrows to the canonical quiet NaN"}, + } + for _, c := range cases { + if got := HalfFromFloat64(c.in); got != c.want { + t.Errorf("HalfFromFloat64(%v) = 0x%04X, want 0x%04X (%s)", c.in, got, c.want, c.why) + } + } +} + +// TestHalfWidening checks the exact widening: every half value is a +// float32 value too, so the classic half to float32 to float64 route is +// an independent oracle for HalfToFloat64 across all 65536 patterns, +// NaNs included. +func TestHalfWidening(t *testing.T) { + for b := range 65536 { + h := uint16(b) + sign := uint32(h&0x8000) << 16 + e := uint32(h>>10) & 0x1F + fr := uint32(h & 0x3FF) + var b32 uint32 + switch { + case e == 0: + if fr == 0 { + b32 = sign + break + } + // frac × 2^-24 renormalised: the leading bit at position k + // makes the float32 exponent k-24. + k := uint(bits.Len32(fr) - 1) + b32 = sign | uint32(127-24+k)<<23 | (fr-1< 1e-15 { + t.Fatalf("Norm = %s %v", nrm.Dtype(), nrm.FloatAt(0)) + } + // TopK keeps half and skips NaN. + tv, ti, err := TopK(nan, 2, 0) + if err != nil { + t.Fatal(err) + } + if tv.Dtype() != Float16 || tv.FloatAt(0) != 2 || tv.FloatAt(1) != 1 { + t.Fatalf("TopK values = %v %v", tv.FloatAt(0), tv.FloatAt(1)) + } + if ti.ints[0] != 2 || ti.ints[1] != 1 { + t.Fatalf("TopK indices = %v %v", ti.ints[0], ti.ints[1]) + } +} + +func TestFloat16MathAndShape(t *testing.T) { + h := mustFromFloat16s(t, []float64{4, 0.25}, 2) + // Element-wise math keeps the half dtype and narrows once. + sq, err := Sqrt(h) + if err != nil { + t.Fatal(err) + } + if sq.Dtype() != Float16 || sq.RawHalves()[0] != 0x4000 || sq.RawHalves()[1] != 0x3800 { + t.Fatalf("Sqrt = %v", sq.RawHalves()) + } + ex, err := Exp(h) + if err != nil { + t.Fatal(err) + } + if ex.Dtype() != Float16 || ex.RawHalves()[0] != HalfFromFloat64(math.Exp(4)) { + t.Fatalf("Exp = 0x%04X", ex.RawHalves()[0]) + } + fl, err := Floor(mustFromFloat16s(t, []float64{1.5}, 1)) + if err != nil { + t.Fatal(err) + } + if fl.Dtype() != Float16 || fl.FloatAt(0) != 1 { + t.Fatalf("Floor = %s %v", fl.Dtype(), fl.FloatAt(0)) + } + pi, err := PowI(h, 2) + if err != nil { + t.Fatal(err) + } + if pi.Dtype() != Float16 || pi.FloatAt(0) != 16 { + t.Fatalf("PowI = %s %v", pi.Dtype(), pi.FloatAt(0)) + } + ab := Abs(mustFromFloat16s(t, []float64{-2}, 1)) + if ab.Dtype() != Float16 || ab.FloatAt(0) != 2 { + t.Fatalf("Abs = %s %v", ab.Dtype(), ab.FloatAt(0)) + } + // Clipping keeps half, in float64 order. + cl, err := ClipF(h, 0.5, 3) + if err != nil { + t.Fatal(err) + } + if cl.Dtype() != Float16 || cl.FloatAt(0) != 3 || cl.FloatAt(1) != 0.5 { + t.Fatalf("ClipF = %v %v", cl.FloatAt(0), cl.FloatAt(1)) + } + cli, err := ClipI(h, 1, 3) + if err != nil { + t.Fatal(err) + } + if cli.Dtype() != Float16 || cli.FloatAt(0) != 3 || cli.FloatAt(1) != 1 { + t.Fatalf("ClipI = %v %v", cli.FloatAt(0), cli.FloatAt(1)) + } + // Shape and layout operations are dtype agnostic. + r, err := Reshape(h, 1, 2) + if err != nil { + t.Fatal(err) + } + if r.Dtype() != Float16 || r.RawHalves()[1] != 0x3400 { + t.Fatalf("Reshape = %v", r.RawHalves()) + } + if tr := Transpose(r); tr.NDim() != 2 || tr.RawHalves()[0] != 0x4400 || tr.RawHalves()[1] != 0x3400 { + t.Fatalf("Transpose = %v", tr.RawHalves()) + } + row, err := Row(r, 0) + if err != nil { + t.Fatal(err) + } + if !slicesEqualU16(row.RawHalves(), []uint16{0x4400, 0x3400}) { + t.Fatalf("Row = %v", row.RawHalves()) + } + col, err := Col(r, 1) + if err != nil { + t.Fatal(err) + } + if !slicesEqualU16(col.RawHalves(), []uint16{0x3400}) { + t.Fatalf("Col = %v", col.RawHalves()) + } + id, err := Identity(Float16, 2) + if err != nil { + t.Fatal(err) + } + if id.FloatAt(0) != 1 || id.FloatAt(1) != 0 || id.FloatAt(3) != 1 { + t.Fatalf("Identity = %s", id) + } + zl := ZerosLike(h) + if zl.Dtype() != Float16 || zl.FloatAt(0) != 0 { + t.Fatal("ZerosLike must keep float16") + } + // Concat promotes through the ladder. + c, err := Concat(h, mustFromFloat32s(t, []float32{1}, 1), 0) + if err != nil { + t.Fatal(err) + } + if c.Dtype() != Float32 { + t.Fatalf("Concat(float16, float32) dtype = %s", c.Dtype()) + } + // Diff widens to float, like every non-int, non-complex dtype. + d, err := Diff(h, 1, 0) + if err != nil { + t.Fatal(err) + } + if d.Dtype() != Float || d.FloatAt(0) != -3.75 { + t.Fatalf("Diff = %s %v", d.Dtype(), d.FloatAt(0)) + } +} + +// TestFloat16AstypeAndAccess walks the conversion and accessor routes. +func TestFloat16AstypeAndAccess(t *testing.T) { + h := mustFromFloat16s(t, []float64{1.5, -0.5, 65504}, 3) + // To Int truncates like float32 to Int does. + i, err := Astype(h, Int) + if err != nil { + t.Fatal(err) + } + if i.ints[0] != 1 || i.ints[1] != 0 || i.ints[2] != 65504 { + t.Fatalf("Astype to int = %v", i.RawInts()) + } + // Up the ladder is exact. + f32, err := Astype(h, Float32) + if err != nil { + t.Fatal(err) + } + if f32.Dtype() != Float32 || f32.FloatAt(2) != 65504 { + t.Fatalf("Astype to float32 = %s", f32) + } + f64, err := Astype(h, Float) + if err != nil { + t.Fatal(err) + } + if f64.FloatAt(0) != 1.5 { + t.Fatalf("Astype to float = %v", f64.FloatAt(0)) + } + // Down from float narrows under the RNE contract. + down, err := Astype(mustFromFloats(t, []float64{0.1, 1e10}, 2), Float16) + if err != nil { + t.Fatal(err) + } + if got := down.RawHalves(); !slicesEqualU16(got, []uint16{0x2E66, 0x7C00}) { + t.Fatalf("Astype float to float16 = %v", got) + } + // Same dtype copies, views included. + sl, err := Slice(h, 0, 1, 3) + if err != nil { + t.Fatal(err) + } + cp, err := Astype(sl, Float16) + if err != nil { + t.Fatal(err) + } + if !slicesEqualU16(cp.RawHalves(), []uint16{0xB800, 0x7BFF}) { + t.Fatalf("Astype same dtype on a view = %v", cp.RawHalves()) + } + // Complex to float16 is a loud refusal, like complex to float32. + if _, err := Astype(mustFromComplexes(t, []complex128{1}, 1), Float16); err == nil || + !strings.Contains(err.Error(), "cannot narrow complex to float16") { + t.Fatalf("Astype complex to float16: %v", err) + } + // WithFloat refuses a float16 array by name, as it refuses float32. + if _, err := WithFloat(h, 1, 0); err == nil || !strings.Contains(err.Error(), "float16") { + t.Fatalf("WithFloat on float16: %v", err) + } + // Scatter refuses a complex source into a half array (canStore). + dst := mustFromHalves(t, []uint16{0, 0}, 2) + _, err = Scatter(dst, 0, mustFromInts(t, []int64{0}, 1), mustFromComplexes(t, []complex128{1i}, 1)) + if err == nil || !strings.Contains(err.Error(), "cannot store") { + t.Fatalf("Scatter complex into float16: %v", err) + } + // Pad's constant mode narrows the fill value. + p, err := Pad(h, []int{0, 1}, "constant", 0.5) + if err != nil { + t.Fatal(err) + } + if got := p.RawHalves(); !slicesEqualU16(got, []uint16{0x3E00, 0xB800, 0x7BFF, 0x3800}) { + t.Fatalf("Pad = %v", got) + } +} + +// TestFloat16AxisAndConversions pins the remaining half paths: the +// strided axis folds, the axis arg-extremes, the sort-based TopK path, +// the scalar maps and the Interpolate2D half grid. +func TestFloat16AxisAndConversions(t *testing.T) { + m2 := mustFromFloat16s(t, []float64{1, 2, 3, 4}, 2, 2) + // SumAxis along dim 0 streams the strided half fold. + sa, err := SumAxis(m2, 0) + if err != nil { + t.Fatal(err) + } + if sa.Dtype() != Float16 || sa.FloatAt(0) != 4 || sa.FloatAt(1) != 6 { + t.Fatalf("SumAxis dim 0 = %v %v", sa.FloatAt(0), sa.FloatAt(1)) + } + // MaxAxis with a NaN in the line: NaN never wins. + nan := mustFromHalves(t, []uint16{0x3C00, 0x7E00, 0x4400, 0x4200}, 2, 2) + ma, err := MaxAxis(nan, 0) + if err != nil { + t.Fatal(err) + } + if ma.Dtype() != Float16 || ma.FloatAt(0) != 4 || ma.FloatAt(1) != 3 { + t.Fatalf("MaxAxis with NaN = %v %v", ma.FloatAt(0), ma.FloatAt(1)) + } + // The axis arg-extremes walk the half payload widened. + am, err := ArgMaxAxis(m2, 0) + if err != nil { + t.Fatal(err) + } + if !slicesEqualI64(am.RawInts(), []int64{1, 1}) { + t.Fatalf("ArgMaxAxis = %v", am.RawInts()) + } + an, err := ArgMinAxis(m2, 0) + if err != nil { + t.Fatal(err) + } + if !slicesEqualI64(an.RawInts(), []int64{0, 0}) { + t.Fatalf("ArgMinAxis = %v", an.RawInts()) + } + // k*8 > n switches TopK to its sort path; the NaN lands last. The + // output is (k, cols): row 0 holds each column's top value. + tv, ti, err := TopK(nan, 2, 0) + if err != nil { + t.Fatal(err) + } + if tv.Dtype() != Float16 || tv.FloatAt(0) != 4 || tv.FloatAt(1) != 3 || + tv.FloatAt(2) != 1 || !math.IsNaN(tv.FloatAt(3)) { + t.Fatalf("TopK values = %s %v", tv.Dtype(), tv.RawHalves()) + } + if !slicesEqualI64(ti.RawInts(), []int64{1, 1, 0, 0}) { + t.Fatalf("TopK indices = %v", ti.RawInts()) + } + // Pow against an int exponent promotes to half and narrows once. + pw, err := Pow(m2, mustFromInts(t, []int64{2, 2, 2, 2}, 2, 2)) + if err != nil { + t.Fatal(err) + } + if pw.Dtype() != Float16 || pw.FloatAt(3) != 16 { + t.Fatalf("Pow = %s %v", pw.Dtype(), pw.FloatAt(3)) + } + // AddC forces complex; the half values widen exactly. + ac := AddC(mustFromFloat16s(t, []float64{1.5}, 1), 2i) + if ac.Dtype() != Complex || ac.ComplexAt(0) != complex(1.5, 2) { + t.Fatalf("AddC on float16 = %s %v", ac.Dtype(), ac.ComplexAt(0)) + } + // Concat of int and half lands on half through setConverted. + ci, err := Concat(mustFromInts(t, []int64{1, 2}, 2), mustFromFloat16s(t, []float64{3, 4}, 2), 0) + if err != nil { + t.Fatal(err) + } + if ci.Dtype() != Float16 || ci.FloatAt(3) != 4 { + t.Fatalf("Concat(int, float16) = %s %v", ci.Dtype(), ci.FloatAt(3)) + } + // Astype half to complex keeps the real route. + cx, err := Astype(mustFromFloat16s(t, []float64{1.5}, 1), Complex) + if err != nil { + t.Fatal(err) + } + if cx.Dtype() != Complex || cx.ComplexAt(0) != complex(1.5, 0) { + t.Fatalf("Astype to complex = %s %v", cx.Dtype(), cx.ComplexAt(0)) + } + // The narrowing constructors refuse a shape their values cannot fill. + if _, err := FromFloat16s([]float64{1}, 2); err == nil { + t.Fatal("FromFloat16s shape mismatch must error") + } + if _, err := HalvesFromArray([]uint16{1}, 2); err == nil { + t.Fatal("HalvesFromArray shape mismatch must error") + } + // A half grid interpolates through the exact widening: a linear + // field is reproduced exactly. + grid := mustFromFloat16s(t, []float64{0, 1, 2, 0, 1, 2}, 2, 3) + xs := mustFromFloats(t, []float64{0.5}, 1) + ys := mustFromFloats(t, []float64{0.5}, 1) + ip, err := Interpolate2D(grid, xs, ys, 0, 0, 1, 1) + if err != nil { + t.Fatal(err) + } + if ip.FloatAt(0) != 0.5 { + t.Fatalf("Interpolate2D half grid = %v", ip.FloatAt(0)) + } +} + +func slicesEqualI64(a, b []int64) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +func TestFloat16StringAndEqual(t *testing.T) { + h := mustFromFloat16s(t, []float64{1, 0.5}, 2) + if got := h.String(); !strings.HasPrefix(got, "float16 (2) [") || !strings.Contains(got, "0.5") { + t.Fatalf("String = %q", got) + } + if Dtype(0).String() != "int" || Float16.String() != "float16" { + t.Fatal("dtype names") + } + // Equal compares values, so -0.0 and +0.0 compare equal and NaN + // never does, exactly as for the other float dtypes. + zero := mustFromHalves(t, []uint16{0x0000, 0x3C00}, 2) + nzero := mustFromHalves(t, []uint16{0x8000, 0x3C00}, 2) + if !Equal(zero, nzero) { + t.Fatal("±0.0 halves must compare equal") + } + nan1 := mustFromHalves(t, []uint16{0x7E00}, 1) + nan2 := mustFromHalves(t, []uint16{0x7E01}, 1) + if Equal(nan1, nan2) { + t.Fatal("NaN halves must never compare equal") + } + // Equal on the very same slice is the documented pointer shortcut + // and answers true before any NaN logic, for every dtype. + // The dtype is part of the identity. + f32 := mustFromFloat32s(t, []float32{1, 0.5}, 2) + if Equal(h, f32) { + t.Fatal("float16 must not equal float32") + } +} + +// TestFloat16ViewsAndMisc exercises the raw-payload paths on strided +// views: every kernel must see exactly the view's own elements. +func TestFloat16ViewsAndMisc(t *testing.T) { + h := mustFromFloat16s(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + // An interior column slice copies; a leading slice views. + // An interior column slice copies: rows [1..2] of columns 1, 2. + col, err := Slice(h, 1, 1, 3) + if err != nil { + t.Fatal(err) + } + if !slicesEqualU16(col.RawHalves(), []uint16{0x4000, 0x4200, 0x4500, 0x4600}) { + t.Fatalf("column slice = %v", col.RawHalves()) + } + if got := Sum(col).Float(); got != 16 { + t.Fatalf("Sum over the slice = %v", got) + } + // A leading-dimension slice is a view; ArgMax materialises it and + // still finds the right element. + view, err := Slice(mustFromFloat16s(t, []float64{1, 5, 2}, 3), 0, 0, 3) + if err != nil { + t.Fatal(err) + } + if got, err := ArgMax(view); err != nil || got != 1 { + t.Fatalf("ArgMax over the view = %d, %v", got, err) + } + s, err := Sort(col) + if err != nil { + t.Fatal(err) + } + if !slicesEqualU16(s.RawHalves(), []uint16{0x4000, 0x4200, 0x4500, 0x4600}) { + t.Fatalf("Sort over the slice = %v", s.RawHalves()) + } + // Unique keeps half and dedupes by value. + u, err := Unique(mustFromHalves(t, []uint16{0x3C00, 0x3C00, 0x8000}, 3)) + if err != nil { + t.Fatal(err) + } + if u.Dtype() != Float16 || u.Len() != 2 { + t.Fatalf("Unique = %s %d", u.Dtype(), u.Len()) + } + // Gather gathers per coordinate along dim 0: index {1, 0, 1} picks + // row 1, row 0, row 1 at the three column positions. + g, err := Gather(h, 0, mustFromInts(t, []int64{1, 0, 1}, 1, 3)) + if err != nil { + t.Fatal(err) + } + if !slicesEqualU16(g.RawHalves(), []uint16{0x4400, 0x4000, 0x4600}) { + t.Fatalf("Gather = %v", g.RawHalves()) + } + tk, err := Take(h, mustFromInts(t, []int64{2, 0}, 2)) + if err != nil { + t.Fatal(err) + } + if !slicesEqualU16(tk.RawHalves(), []uint16{0x4200, 0x3C00}) { + t.Fatalf("Take = %v", tk.RawHalves()) + } + rev := Reverse(h) + if rev.FloatAt(0) != 6 || rev.FloatAt(5) != 1 { + t.Fatalf("Reverse = %s", rev) + } + // SparseFrom and Dense round-trip the bits. + sp, err := SparseFrom(mustFromHalves(t, []uint16{0x3C00, 0x0000, 0x4000}, 3)) + if err != nil { + t.Fatal(err) + } + if sp.Values.Dtype() != Float16 || sp.NNZ() != 2 { + t.Fatalf("SparseFrom dtype %s nnz %d", sp.Values.Dtype(), sp.NNZ()) + } + dn, err := sp.Dense() + if err != nil { + t.Fatal(err) + } + if dn.Dtype() != Float16 || !slicesEqualU16(dn.RawHalves(), []uint16{0x3C00, 0x0000, 0x4000}) { + t.Fatalf("Sparse Dense = %v", dn.RawHalves()) + } + // SpMul promotes and narrows once. + sm, err := SpMul(sp, mustFromFloat16s(t, []float64{2, 0, 0.5}, 3)) + if err != nil { + t.Fatal(err) + } + if sm.Dtype() != Float16 || sm.FloatAt(0) != 2 || sm.FloatAt(2) != 1 { + t.Fatalf("SpMul = %s %v %v", sm.Dtype(), sm.FloatAt(0), sm.FloatAt(2)) + } +} + +// TestFloat16Refusals pins the loud refusals: the matmul and einsum +// kernels are not offered for the half dtype, and conversion with +// Astype is the documented route. +func TestFloat16Refusals(t *testing.T) { + h2 := mustFromFloat16s(t, []float64{1, 2, 3, 4}, 2, 2) + f2 := mustFromFloats(t, []float64{1, 0, 0, 1}, 2, 2) + if _, err := MatMul2D(h2, f2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") { + t.Fatalf("MatMul2D float16: %v", err) + } + if _, err := MatMul2D(f2, h2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") { + t.Fatalf("MatMul2D float16 on the right: %v", err) + } + if _, err := Einsum("ij,jk->ik", h2, f2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") { + t.Fatalf("Einsum float16: %v", err) + } + if _, err := Einsum("ij->", h2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") { + t.Fatalf("Einsum reduce-all float16: %v", err) + } + sp := &SparseCOO{Indices: mustFromInts(t, []int64{0, 0}, 1, 2), Values: mustFromFloat16s(t, []float64{1}, 1), Shape: []int{2, 2}} + if _, err := SpMatMul(sp, f2); err == nil || !strings.Contains(err.Error(), "float16 is not supported") { + t.Fatalf("SpMatMul float16: %v", err) + } + // The float64 side is unchanged: float64 by float64 still computes. + got, err := MatMul2D(f2, f2) + if err != nil || got.Dtype() != Float || got.FloatAt(0) != 1 { + t.Fatalf("MatMul2D float64 baseline moved: %v %s", err, got) + } + // Kron promotes to half and narrows once. + k, err := Kron(h2, mustFromInts(t, []int64{1, 1, 1, 1}, 2, 2)) + if err != nil { + t.Fatal(err) + } + if k.Dtype() != Float16 || k.FloatAt(0) != 1 || k.FloatAt(3) != 2 || k.FloatAt(15) != 4 { + t.Fatalf("Kron = %s %v %v %v", k.Dtype(), k.FloatAt(0), k.FloatAt(3), k.FloatAt(15)) + } + // Elements refuses only the genuine narrowings. + h1 := mustFromFloat16s(t, []float64{1.5}, 1) + if _, err := h1.Elements[int64](); err == nil || !strings.Contains(err.Error(), "cannot narrow float16 to int64") { + t.Fatalf("Elements int64 on float16: %v", err) + } + if v, err := h1.Elements[float64](); err != nil || v[0] != 1.5 { + t.Fatalf("Elements float64 on float16: %v %v", v, err) + } + if v, err := h1.Elements[float32](); err != nil || v[0] != 1.5 { + t.Fatalf("Elements float32 on float16: %v %v", v, err) + } +} + +// The ordering trap: half bits do not order as uint16 for negatives +// (−1.0 is 0xBC00, which bit order places after 2.0), so the kernels +// that order or clamp must widen to float64 first. Each pin below +// fails if the glue ever sorts or clamps the raw payload, and the +// overflow pin holds the compute-in-float64-narrow-once contract's +// sticky infinity. +func TestFloat16NegativeOrdering(t *testing.T) { + x, xerr := FromFloat16s([]float64{2, -1, 0.5}, 3) + if xerr != nil { + t.Fatal(xerr) + } + sorted, serr := Sort(x) + if serr != nil { + t.Fatal(serr) + } + want := []uint16{HalfFromFloat64(-1), HalfFromFloat64(0.5), HalfFromFloat64(2)} + for i, h := range want { + if got := sorted.RawHalves()[i]; got != h { + t.Fatalf("Sort ascending [%d] = %#04x, want %#04x", i, got, h) + } + } + order, oerr := ArgSort(x) + if oerr != nil { + t.Fatal(oerr) + } + if got0, got1, got2 := order.FloatAt(0), order.FloatAt(1), order.FloatAt(2); got0 != 1 || got1 != 2 || got2 != 0 { + t.Fatalf("ArgSort ascending = [%g %g %g], want [1 2 0]", got0, got1, got2) + } + fIn, ferr := FromFloat16s([]float64{-0.5, 0.25}, 2) + if ferr != nil { + t.Fatal(ferr) + } + clipped, cerr := ClipF(fIn, -0.25, 0.5) + if cerr != nil { + t.Fatal(cerr) + } + if got := clipped.RawHalves()[0]; got != HalfFromFloat64(-0.25) { + t.Fatalf("ClipF on a negative = %#04x, want %#04x", got, HalfFromFloat64(-0.25)) + } + clipIn, clipErr := FromFloat16s([]float64{-5, 0.25}, 2) + if clipErr != nil { + t.Fatal(clipErr) + } + clamped, ierr := ClipI(clipIn, -1, 1) + if ierr != nil { + t.Fatal(ierr) + } + if math.Abs(clamped.FloatAt(0)+1) > 1e-12 { + t.Fatalf("ClipI on a negative past the wall = %g, want -1", clamped.FloatAt(0)) + } + lhs, lerr := FromFloat16s([]float64{65504, 65504}, 2) + if lerr != nil { + t.Fatal(lerr) + } + rhs, rerr := FromFloat16s([]float64{2, 2}, 2) + if rerr != nil { + t.Fatal(rerr) + } + overflow, merr := Mul(lhs, rhs) + if merr != nil { + t.Fatal(merr) + } + if h := overflow.RawHalves()[0]; h != 0x7C00 { + t.Fatalf("65504*2 in half = %#04x, want the +Inf half 0x7C00", h) + } + back, serr2 := Sub(overflow, rhs) + if serr2 != nil { + t.Fatal(serr2) + } + if h := back.RawHalves()[0]; h != 0x7C00 { + t.Fatalf("+Inf minus 65504 in half = %#04x, want the sticky +Inf 0x7C00", h) + } +} diff --git a/internal/core/float32_test.go b/internal/core/float32_test.go new file mode 100644 index 0000000..ad25255 --- /dev/null +++ b/internal/core/float32_test.go @@ -0,0 +1,334 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +func mustFromFloat32s(t *testing.T, vals []float32, shape ...int) *Array { + t.Helper() + a, err := FromFloat32s(vals, shape...) + if err != nil { + t.Fatalf("FromFloat32s(%v, %v): %v", vals, shape, err) + } + return a +} + +func TestFloat32Constructors(t *testing.T) { + a := mustFromFloat32s(t, []float32{1.5, -2.5}, 2) + if a.Dtype() != Float32 || a.Len() != 2 { + t.Fatalf("float32 array: %s len %d", a.Dtype(), a.Len()) + } + // floatAt reads exactly. + if v := a.FloatAt(0); v != 1.5 { + t.Fatalf("floatAt: %v", v) + } + + z, _ := Zeros(Float32, 3) + if z.Dtype() != Float32 || z.FloatAt(1) != 0 { + t.Fatalf("Zeros float32: %s", z) + } + o, _ := Ones(Float32, 2) + if o.FloatAt(0) != 1 { + t.Fatalf("Ones float32: %s", o) + } + f, _ := FullF32s(0.5, 2) + if f.FloatAt(1) != 0.5 { + t.Fatalf("FullF32s: %s", f) + } + // The dtype is part of the identity: float never equals float32. + wide := mustFromFloats(t, []float64{1.5, -2.5}, 2) + if Equal(wide, mustFromFloat32s(t, []float32{1.5, -2.5}, 2)) { + t.Fatal("float and float32 must not compare equal") + } + if got := a.String(); !strings.Contains(got, "float32") { + t.Fatalf("String: %q", got) + } +} + +func TestFloat32PromotionMatrix(t *testing.T) { + f32 := mustFromFloat32s(t, []float32{1.5}, 1) + i := mustFromInts(t, []int64{1}, 1) + f64 := mustFromFloats(t, []float64{1.5}, 1) + c := mustFromComplexes(t, []complex128{1}, 1) + + cases := []struct { + b *Array + want Dtype + }{ + {f32, Float32}, + {i, Float32}, + {f64, Float}, + {c, Complex}, + } + for _, tc := range cases { + sum, err := Add(f32, tc.b) + if err != nil { + t.Fatalf("Add(%s): %v", tc.b.Dtype(), err) + } + if sum.Dtype() != tc.want { + t.Fatalf("float32 + %s = %s, want %s", tc.b.Dtype(), sum.Dtype(), tc.want) + } + } +} + +func TestFloat32Arithmetic(t *testing.T) { + a := mustFromFloat32s(t, []float32{0.1, 2}, 2) + b := mustFromFloat32s(t, []float32{0.2, 4}, 2) + + sum, err := Add(a, b) + if err != nil || sum.Dtype() != Float32 { + t.Fatalf("Add: %s %v", sum, err) + } + // 0.1+0.2 computed in float64 and rounded once, the float32 result. + if v := sum.FloatAt(0); v != 0.30000001192092896 { + t.Fatalf("Add value: %v", v) + } + + prod, _ := Mul(a, b) + if v := prod.FloatAt(1); v != 8 { + t.Fatalf("Mul value: %v", v) + } + + // Division keeps float32. + q, _ := Div(b, a) + if q.Dtype() != Float32 { + t.Fatalf("Div dtype: %s", q.Dtype()) + } + if v := q.FloatAt(1); v != 2 { + t.Fatalf("Div value: %v", v) + } + + // Weak scalars keep float32 (NEP-50 style). + if v := AddI(a, 1).FloatAt(1); v != 3 { + t.Fatalf("AddI: %v", v) + } + if got := AddI(a, 1).Dtype(); got != Float32 { + t.Fatalf("AddI dtype: %s", got) + } + if got := AddF(a, 0.5).Dtype(); got != Float32 { + t.Fatalf("AddF keeps float32: %s", got) + } + if v := AddF(a, 0.25).FloatAt(1); v != 2.25 { + t.Fatalf("AddF value: %v", v) + } + // int arrays still promote to float64 under F scalars. + if got := AddF(mustFromInts(t, []int64{1}, 1), 0.5).Dtype(); got != Float { + t.Fatalf("int AddF dtype: %s", got) + } +} + +func TestFloat32MathFuncs(t *testing.T) { + a := mustFromFloat32s(t, []float32{4}, 1) + + sqrt, err := Sqrt(a) + if err != nil || sqrt.Dtype() != Float32 { + t.Fatalf("Sqrt: %s %v", sqrt, err) + } + if v := sqrt.FloatAt(0); v != 2 { + t.Fatalf("Sqrt value: %v", v) + } + + // Rounding keeps the width. + r, _ := Floor(mustFromFloat32s(t, []float32{2.7}, 1)) + if r.Dtype() != Float32 || r.FloatAt(0) != 2 { + t.Fatalf("Floor: %s %v", r, r.FloatAt(0)) + } + + // Abs keeps float32; complex magnitudes still go to float64. + ab := Abs(mustFromFloat32s(t, []float32{-2.5}, 1)) + if ab.Dtype() != Float32 || ab.FloatAt(0) != 2.5 { + t.Fatalf("Abs: %s %v", ab, ab.FloatAt(0)) + } + + tanh, _ := Tanh(mustFromFloat32s(t, []float32{0}, 1)) + if tanh.Dtype() != Float32 || tanh.FloatAt(0) != 0 { + t.Fatalf("Tanh: %s", tanh) + } +} + +func TestFloat32ReductionsAndMatMul(t *testing.T) { + m := mustFromFloat32s(t, []float32{1, 2, 3, 4, 5, 6}, 2, 3) + + // Scalar reductions keep the float scalar kind for float32 arrays. + sum := Sum(m) + if !sum.IsFloat() || sum.Float() != 21 { + t.Fatalf("Sum float32: %s", sum) + } + mn, err := Min(m) + if err != nil || !mn.IsFloat() || mn.Float() != 1 { + t.Fatalf("Min float32: %s %v", mn, err) + } + mx, err := Max(m) + if err != nil || !mx.IsFloat() || mx.Float() != 6 { + t.Fatalf("Max float32: %s %v", mx, err) + } + mean, err := Mean(m) + if err != nil || mean != 3.5 { + t.Fatalf("Mean float32: %v %v", mean, err) + } + f32a := mustFromFloat32s(t, []float32{1, 2}, 2) + f32b := mustFromFloat32s(t, []float32{3, 4}, 2) + d, err := Dot(f32a, f32b) + if err != nil || !d.IsFloat() || d.Float() != 11 { + t.Fatalf("Dot float32: %s %v", d, err) + } + + sums, err := SumAxis(m, 1) + if err != nil { + t.Fatalf("SumAxis: %v", err) + } + if sums.Dtype() != Float32 { + t.Fatalf("SumAxis dtype: %s", sums.Dtype()) + } + if v := sums.FloatAt(0); v != 6 { + t.Fatalf("SumAxis value: %v", v) + } + + maxes, _ := MaxAxis(m, 1) + if maxes.Dtype() != Float32 || maxes.FloatAt(1) != 6 { + t.Fatalf("MaxAxis: %s", maxes) + } + + // Mean is float64 regardless. + axisMean, _ := MeanAxis(m, 1) + if axisMean.Dtype() != Float { + t.Fatalf("MeanAxis dtype: %s", axisMean.Dtype()) + } + + // MatMul accumulates in float64 and rounds once. + w := mustFromFloat32s(t, []float32{1, 2, 3, 4, 5, 6}, 3, 2) + p, err := MatMul2D(m, w) + if err != nil { + t.Fatalf("MatMul: %v", err) + } + if p.Dtype() != Float32 { + t.Fatalf("MatMul dtype: %s", p.Dtype()) + } + // [[1,2,3],[4,5,6]]·[[1,2],[3,4],[5,6]] = [[22,28],[49,64]] + if v := p.FloatAt(0); v != 22 { + t.Fatalf("MatMul value: %v", v) + } + if v := p.FloatAt(3); v != 64 { + t.Fatalf("MatMul value: %v", v) + } + + // Vector shapes keep float32 too. + v1 := mustFromFloat32s(t, []float32{1, 2}, 2) + mv, err := MatMul2D(m, mustFromFloat32s(t, []float32{1, 1, 1}, 3)) + if err != nil || mv.Dtype() != Float32 { + t.Fatalf("matrix × vector: %s %v", mv, err) + } + if v := mv.FloatAt(0); v != 6 { + t.Fatalf("matrix × vector value: %v", v) + } + _ = v1 +} + +func TestFloat32Machinery(t *testing.T) { + f := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 2, 2) + + // Identity works for every dtype. + id, err := Identity(Float32, 3) + if err != nil { + t.Fatalf("Identity float32: %v", err) + } + if id.Dtype() != Float32 || id.FloatAt(0) != 1 || id.FloatAt(4) != 1 { + t.Fatalf("Identity float32: %s", id) + } + // int × int Pow stays int; a negative exponent is an error. + pi := mustFromInts(t, []int64{2, 3}, 2) + pe := mustFromInts(t, []int64{3, 2}, 2) + pow, err := Pow(pi, pe) + if err != nil || pow.Dtype() != Int || pow.RawInts()[0] != 8 || pow.RawInts()[1] != 9 { + t.Fatalf("Pow int: %s %v", pow, err) + } + if _, err := Pow(pi, mustFromInts(t, []int64{1, -1}, 2)); err == nil || + !strings.Contains(err.Error(), "negative exponent") { + t.Fatalf("Pow negative exponent: %v", err) + } + // PowI keeps float32. + pf, err := PowI(mustFromFloat32s(t, []float32{2, 3}, 2), 2) + if err != nil || pf.Dtype() != Float32 || pf.FloatAt(1) != 9 { + t.Fatalf("PowI float32: %s %v", pf, err) + } + + // Mask selection, Where, comparisons and broadcasting all carry the + // dtype. + mask, err := GtI(f, 2) + if err != nil { + t.Fatalf("GtI: %v", err) + } + sel, _ := Select(f, mask) + if sel.Dtype() != Float32 || sel.Len() != 2 { + t.Fatalf("Mask: %s", sel) + } + w, _ := Where(mask, f, mustFromFloat32s(t, []float32{9, 9, 9, 9}, 2, 2)) + if w.Dtype() != Float32 || w.FloatAt(0) != 9 { + t.Fatalf("Where: %s", w) + } + lt, _ := LtF(f, 3) + if got := lt.FloatAt(2); got != 0 { + t.Fatalf("LtF on float32: %v", got) + } + b, _ := Slice(f, 0, 1, 2) + if b.Dtype() != Float32 || b.FloatAt(0) != 3 { + t.Fatalf("Slice: %s", b) + } + tt := Transpose(f) + if tt.FloatAt(1) != 3 { + t.Fatalf("Transpose: %s", tt) + } + cat, _ := Concat(f, f, 1) + if cat.Dtype() != Float32 || cat.Shape()[1] != 4 { + t.Fatalf("Concat: %s", cat) + } + + // Sort, Clip and elements. + s, _ := Sort(mustFromFloat32s(t, []float32{3, 1, 2}, 3)) + if s.FloatAt(0) != 1 { + t.Fatalf("Sort: %s", s) + } + clip, _ := ClipI(mustFromFloat32s(t, []float32{-5, 3}, 2), 0, 1) + if clip.Dtype() != Float32 || clip.FloatAt(0) != 0 || clip.FloatAt(1) != 1 { + t.Fatalf("ClipI: %s", clip) + } + clipF, _ := ClipF(mustFromFloat32s(t, []float32{-5, 3}, 2), 0, 1) + if clipF.Dtype() != Float32 { + t.Fatalf("ClipF dtype: %s", clipF) + } + + // Elements conversions along the ladder. + vals, err := f.Elements[float32]() + if err != nil || vals[3] != 4 { + t.Fatalf("Elements[float32]: %v %v", vals, err) + } + widened, err := f.Elements[float64]() + if err != nil || widened[0] != 1 { + t.Fatalf("Elements[float64] from float32: %v %v", widened, err) + } + if _, err := mustFromFloats(t, []float64{1}, 1).Elements[float32](); err == nil || + !strings.Contains(err.Error(), "cannot narrow float to float32") { + t.Fatalf("Elements float to float32: %v", err) + } + + // Generator.Float32s stays in range. + g := NewGenerator(21) + r, err := Float32s(g, 500) + if err != nil || r.Dtype() != Float32 { + t.Fatalf("Float32s: %s %v", r, err) + } + for i := range r.Len() { + v := r.FloatAt(i) + if v < 0 || v >= 1 { + t.Fatalf("Float32s out of [0,1): %v", v) + } + } + if _, ok := any(math.NaN()).(float32); ok { + _ = ok // keep math imported if unused paths change + } +} diff --git a/internal/core/fresnel.go b/internal/core/fresnel.go new file mode 100644 index 0000000..e4ab20e --- /dev/null +++ b/internal/core/fresnel.go @@ -0,0 +1,199 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" +) + +// Fresnel integrals, the diffraction pair C(x) = ∫₀^x cos(πt²/2) dt +// and S(x) = ∫₀^x sin(πt²/2) dt. +// +// The power series converge for every real x, but their intermediate +// terms grow like e^(πx²/2) before cancelling down to the ½ limit. +// Within |x| ≤ 4 the cancellation is mild and plain float64 holds full +// accuracy. Above that the same series is summed in extended +// precision (math/big), with the working bit size scaled to the +// largest intermediate term, so the float64 result stays correct +// instead of drowning in cancellation noise. Both integrals are odd +// functions. + +// FresnelC returns the Fresnel cosine integral C(x) of each element. +func FresnelC(a *Array) (*Array, error) { return a.realFunc("FresnelC", fresnelC) } + +// FresnelS returns the Fresnel sine integral S(x) of each element. +func FresnelS(a *Array) (*Array, error) { return a.realFunc("FresnelS", fresnelS) } + +const fresnelSeriesCutoff = 3 + +// fresnelAsymptoticFrom is the argument above which the erfc +// asymptotic series answers the Fresnel pair. There the terms of that +// series fall by a factor 2πx² per order, so a handful of them reach +// double precision, where the power series would need a working +// precision growing with x² (and its term count with x²). +const fresnelAsymptoticFrom = 10 + +func fresnelC(x float64) float64 { + if x < 0 { + return -fresnelC(-x) + } + if x > fresnelAsymptoticFrom { + return fresnelAsym(x, false) + } + if x > fresnelSeriesCutoff { + return fresnelBig(x, false) + } + return fresnelCSeries(x) +} + +func fresnelS(x float64) float64 { + if x < 0 { + return -fresnelS(-x) + } + if x > fresnelAsymptoticFrom { + return fresnelAsym(x, true) + } + if x > fresnelSeriesCutoff { + return fresnelBig(x, true) + } + return fresnelSSeries(x) +} + +// fresnelCSeries sums C(x) = Σ (−1)^k (π/2)^{2k} x^{4k+1}/((2k)!(4k+1)) +// by the term-ratio recurrence. Accurate while the largest +// intermediate term stays comfortably inside float64, which holds +// through |x| ≈ 4. +func fresnelCSeries(x float64) float64 { + sum := 0.0 + term := x // the k = 0 term + for k := range 200 { + sum += term + next := -term * (math.Pi / 2) * (math.Pi / 2) * x * x * x * x / + float64((2*k+1)*(2*k+2)) * float64(4*k+1) / float64(4*k+5) + if math.Abs(next) < 1e-18*math.Abs(sum) { + break + } + term = next + } + return sum +} + +// fresnelSSeries sums S(x) = Σ (−1)^k (π/2)^{2k+1} x^{4k+3}/((2k+1)!(4k+3)) +// by the term-ratio recurrence, with the same validity range as the +// cosine branch. +func fresnelSSeries(x float64) float64 { + sum := 0.0 + term := (math.Pi / 2) * x * x * x / 3 // the k = 0 term + for k := range 200 { + sum += term + next := -term * (math.Pi / 2) * (math.Pi / 2) * x * x * x * x / + float64((2*k+2)*(2*k+3)) * float64(4*k+3) / float64(4*k+7) + if math.Abs(next) < 1e-18*math.Abs(sum) { + break + } + term = next + } + return sum +} + +// fresnelAsym evaluates C or S from the auxiliary functions above. +// x must be positive and at least fresnelAsymptoticFrom. +func fresnelAsym(x float64, sine bool) float64 { + f, g := fresnelAux(x) + u := math.Pi * x * x / 2 + cosu, sinu := math.Cos(u), math.Sin(u) + if sine { + return 0.5 - f*cosu - g*sinu + } + return 0.5 + f*sinu - g*cosu +} + +// fresnelAux returns the auxiliary functions f and g of the Fresnel +// asymptotics, +// +// S(x) = ½ − f·cos(πx²/2) − g·sin(πx²/2) +// C(x) = ½ + f·sin(πx²/2) − g·cos(πx²/2) +// +// as the classical expansions +// +// f(x) = (1/(πx))·Σ (−1)^k (4k−1)!!·t^{2k} +// g(x) = (1/(πx))·Σ (−1)^k (4k+1)!!·t^{2k+1}, t = 1/(πx²), +// +// whose coefficients follow from (4k+1)(4k+3). x must be positive and +// at least fresnelAsymptoticFrom: there t ≤ 1/314 and the terms fall +// fast enough for four of them to pass double precision. The +// coefficients and the reconstruction were checked against mpmath to +// 1e-20 at x = 8, 12, 20, 100, 1000. +func fresnelAux(x float64) (f, g float64) { + t := 1 / (math.Pi * x * x) + cf, cg := 1.0, t + sumF, sumG := 1.0, t + for k := range 20 { + cf = -cf * float64(4*k+1) * float64(4*k+3) * t * t + cg = -cg * float64(4*k+3) * float64(4*k+5) * t * t + sumF += cf + sumG += cg + if math.Abs(cf)+math.Abs(cg) < 1e-19 { + break + } + } + inv := 1 / (math.Pi * x) + return inv * sumF, inv * sumG +} + +// beyond the float64 cancellation ceiling. The working bit size grows +// with the largest intermediate term; the tail is cut once the terms +// fall below a few 1e-40 relative to the running sum, which bounds the +// final float64 rounding far below the 1e-15 relative level. +func fresnelBig(x float64, sine bool) float64 { + // The largest intermediate term of the series grows like e^{πx²/2}, + // so the working precision must grow by π/(2·ln2) ≈ 2.266 bits per + // unit of x²; a smaller coefficient silently loses the cancellation + // and the result turns to noise around x ≈ 22. + bits := uint(64 + int(2.2663*x*x) + 96) + xf := new(big.Float).SetPrec(bits).SetFloat64(x) + halfPi := new(big.Float).SetPrec(bits).SetFloat64(math.Pi / 2) + hh := new(big.Float).SetPrec(bits).Mul(halfPi, halfPi) + x4 := new(big.Float).SetPrec(bits).Mul(xf, xf) + x4.Mul(x4, xf) + x4.Mul(x4, xf) + + term := new(big.Float).SetPrec(bits) + if sine { + // S seed: (π/2)·x³/3. + term.Mul(halfPi, xf) + term.Mul(term, xf) + term.Mul(term, xf) + term.Quo(term, big.NewFloat(3)) + } else { + term.Set(xf) + } + sum := new(big.Float).SetPrec(bits) + + tiny := new(big.Float).SetPrec(bits).SetFloat64(1e-40) + tinyNeg := new(big.Float).SetPrec(bits).SetFloat64(-1e-40) + for k := range 200000 { + sum.Add(sum, term) + an := 4*k + 1 + d1, d2, d3 := 2*k+1, 2*k+2, 4*k+5 + if sine { + an = 4*k + 3 + d1, d2, d3 = 2*k+2, 2*k+3, 4*k+7 + } + next := new(big.Float).SetPrec(bits).Neg(term) + next.Mul(next, hh) + next.Mul(next, x4) + next.Mul(next, big.NewFloat(float64(an))) + next.Quo(next, big.NewFloat(float64(d1))) + next.Quo(next, big.NewFloat(float64(d2))) + next.Quo(next, big.NewFloat(float64(d3))) + if next.Cmp(tiny) < 0 && next.Cmp(tinyNeg) > 0 { + break + } + term = next + } + out, _ := sum.Float64() + return out +} diff --git a/internal/core/grid.go b/internal/core/grid.go new file mode 100644 index 0000000..8e77655 --- /dev/null +++ b/internal/core/grid.go @@ -0,0 +1,118 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/bits" +) + +// Multilinear interpolation over a regular grid of any rank: the +// rank-2 generalisation of Interpolate2D's bilinear read, one weight +// per axis, clamped at the boundary. + +// InterpolateGrid reads the grid at every query point by multilinear +// interpolation. grid has one axis per dimension (each at least two +// samples); origins[i] and steps[i] place axis i, with steps strictly +// positive; queries is a rank-2 (m × rank) matrix, one row per query, +// axes in the grid's own order. Queries outside the grid clamp to the +// boundary, matching Interpolate's convention; a position that comes +// out NaN is an error, since there is nothing sensible to clamp it to. +// A rank above 12 is an error: 2^12 corners per query is where this +// implementation stays honest about its cost. +// +// The grid and the queries are widened once and the walk reads the +// payload slices, the same values the accessors returned; the axes +// combine bottom up, deepest first, which is the order the recursive +// walk evaluated, so the bits are the recursive walk's. +func InterpolateGrid(grid *Array, origins, steps []float64, queries *Array) (*Array, error) { + const name = "InterpolateGrid" + dims := grid.NDim() + if dims < 1 { + return nil, errf("%s: the grid must have at least one axis", name) + } + if dims > 12 { + return nil, errf("%s: the grid has %d axes, above the 12 this multilinear read supports", name, dims) + } + if len(origins) != dims || len(steps) != dims { + return nil, errf("%s: origins and steps need %d entries each, got %d and %d", name, dims, len(origins), len(steps)) + } + shape := grid.Shape() + for i := range dims { + if !isFiniteStep(steps[i]) || steps[i] <= 0 { + return nil, errf("%s: axis %d needs a positive finite step, got %g", name, i, steps[i]) + } + if shape[i] < 2 { + return nil, errf("%s: axis %d needs at least two samples, got %d", name, i, shape[i]) + } + } + if queries.NDim() != 2 || queries.Shape()[1] != dims { + return nil, errf("%s: queries must be rank 2 with %d columns, got shape %s", name, dims, shapeText(queries.Shape())) + } + if grid.dt == Complex || queries.dt == Complex { + return nil, errf("%s: complex grids are not supported", name) + } + m := queries.Shape()[0] + out := &Array{shape: []int{m}, dt: Float} + out.alloc(m) + gv, qv := floatPayload(grid), floatPayload(queries) + // The axis strides and the corner flats are functions of the shape + // alone: precomputed once, the corner table stepped by the lowest + // cleared bit so no query re-walks the axes to find its corners. + strides := make([]int, dims) + strides[dims-1] = 1 + for i := dims - 2; i >= 0; i-- { + strides[i] = strides[i+1] * shape[i+1] + } + corners := 1 << dims + flat := make([]int, corners) + for c := 1; c < corners; c++ { + flat[c] = flat[c&(c-1)] + strides[bits.TrailingZeros32(uint32(c))] + } + buf := make([]float64, corners) + idx := make([]int, dims) + frac := make([]float64, dims) + for q := range m { + base := 0 + for i := range dims { + t := (qv[q*dims+i] - origins[i]) / steps[i] + // NaN compares false against both clamps below and converts + // to the platform's indefinite integer, which then indexes + // far outside the grid: the clamp contract only holds for + // the infinities, so an undefined position is a loud error. + if math.IsNaN(t) { + return nil, errf("%s: query %d axis %d is NaN, which cannot be clamped", name, q, i) + } + if t < 0 { + t = 0 + } + if t > float64(shape[i]-1) { + t = float64(shape[i] - 1) + } + idx[i] = min(int(t), shape[i]-2) + frac[i] = t - float64(idx[i]) + base += idx[i] * strides[i] + } + // The corners seed the deepest level and the axes combine from + // the last one up: buf[l] takes the lo corner first, the hi + // corner second, the operand order the recursive walk kept. + for c := range corners { + buf[c] = gv[base+flat[c]] + } + for i := dims - 1; i >= 0; i-- { + w := frac[i] + lo := 1 - w + for l := range 1 << i { + buf[l] = lo*buf[l] + w*buf[l+1< (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/rand/v2" + "strings" + "testing" + "time" +) + +// Guard pins: each test guards a repair. An invalid +// einsum label is an error instead of a silent contraction, the pInf +// norm reports NaN for an all-NaN line, circular padding folds in +// constant time, BesselKn owns its domain, the float32 sparse matmul +// accumulates in float64, a NaN interpolation query is an +// error, strided views read their own elements everywhere, unknown +// dtypes are refused at the constructors, Scalar.String prints one +// sign, and Mean survives a sum that would overflow. The TestPin +// guards at the end cover paths that already agreed with their naive +// references and only needed the coverage pinned. + +// ---------- Einsum rejects labels outside a-z and A-Z ---------- + +func TestEinsumInvalidLabels(t *testing.T) { + m := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2) + v := mustFloats(t, []float64{7, 8}, 2) + // Three specs the old label reader silently computed: the pound + // sign decoded as a label and returned a copy, the euro sign did + // the same on a vector, and the digit is refused for consistency + // with the general engine's parser. + for _, c := range []struct { + spec string + op *Array + bad string + }{ + {"£->£", m, "£"}, + {"€->€", v, "€"}, + {"i0,i0->i0", m, "0"}, + } { + out, err := Einsum(c.spec, c.op) + if err == nil { + t.Errorf("Einsum(%q) = %v, want an error", c.spec, out.Shape()) + continue + } + if !strings.Contains(err.Error(), c.bad) { + t.Errorf("Einsum(%q) error %q does not name the offending character %q", c.spec, err, c.bad) + } + } + // The valid surface is unchanged: the table path still answers the + // same results bit for bit, uppercase letters included. + a := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2) + b := mustFloats(t, []float64{5, 6, 7, 8}, 2, 2) + table, err := Einsum("ij,jk->ik", a, b) + if err != nil { + t.Fatal(err) + } + direct, derr := MatMul2D(a, b) + if derr != nil { + t.Fatal(derr) + } + if !samePayload(t, table, direct) { + t.Error("Einsum ij,jk->ik no longer matches MatMul2D bit for bit") + } + general, err := Einsum("ij,jk", a, b) + if err != nil { + t.Fatal(err) + } + if !samePayload(t, table, general) { + t.Error("Einsum table and general paths disagree") + } + upA := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2) + upB := mustFloats(t, []float64{5, 6, 7, 8}, 2, 2) + up, err := Einsum("AB,BC->AC", upA, upB) + if err != nil { + t.Fatalf("uppercase labels rejected: %v", err) + } + if !samePayload(t, up, table) { + t.Error("uppercase labels answered a different product") + } + x := mustFromInts(t, []int64{1, 2}, 2) + y := mustFromInts(t, []int64{3, 4}, 2) + dot, err := Einsum("i,i->", x, y) + if err != nil { + t.Fatal(err) + } + if dot.Dtype() != Int || dot.ints[0] != 11 { + t.Errorf("Einsum i,i-> = %s %v, want int [11]", dot.Dtype(), dot.ints) + } +} + +// samePayload reports whether two arrays hold identical values, dtype +// included, comparing int payloads as integers so the check stays bit +// exact for every element type. +func samePayload(t *testing.T, a, b *Array) bool { + t.Helper() + if a.Dtype() != b.Dtype() || !sameShape(a.Shape(), b.Shape()) { + return false + } + for i := range a.Len() { + switch a.Dtype() { + case Int: + if a.ints[i] != b.ints[i] { + return false + } + case Float32: + if math.Float32bits(a.floats32[i]) != math.Float32bits(b.floats32[i]) { + return false + } + case Float: + if math.Float64bits(a.floats[i]) != math.Float64bits(b.floats[i]) { + return false + } + default: + if a.complexes[i] != b.complexes[i] { + return false + } + } + } + return true +} + +// ---------- Norm with p = +Inf propagates NaN ---------- + +func TestNormInfNaN(t *testing.T) { + allNaN := mustFloats(t, []float64{math.NaN(), math.NaN()}, 2) + out, err := Norm(allNaN, math.Inf(1), 0, false) + if err != nil { + t.Fatal(err) + } + if v := out.FloatAt(0); !math.IsNaN(v) { + t.Errorf("Norm(all-NaN, Inf) = %v, want NaN like p = 1 and p = 2", v) + } + // NaN between finite values is skipped, as it always was: the + // magnitudes decide. + mixed := mustFloats(t, []float64{math.NaN(), 3, -9}, 3) + mixOut, err := Norm(mixed, math.Inf(1), 0, false) + if err != nil { + t.Fatal(err) + } + if v := mixOut.FloatAt(0); v != 9 { + t.Errorf("Norm([NaN, 3, -9], Inf) = %v, want 9", v) + } + // Lines without NaN keep their exact answers. + clean := mustFloats(t, []float64{1, -4, 2}, 3) + cleanOut, err := Norm(clean, math.Inf(1), 0, false) + if err != nil { + t.Fatal(err) + } + if v := cleanOut.FloatAt(0); v != 4 { + t.Errorf("Norm([1, -4, 2], Inf) = %v, want 4", v) + } + // The float32 path seeds the same way. + f32NaN := &Array{shape: []int{2}, dt: Float32, floats32: []float32{float32(math.NaN()), 5}} + f32Out, err := Norm(f32NaN, math.Inf(1), 0, false) + if err != nil { + t.Fatal(err) + } + if v := f32Out.FloatAt(0); v != 5 { + t.Errorf("Norm(float32 [NaN, 5], Inf) = %v, want 5", v) + } + f32All := &Array{shape: []int{2}, dt: Float32, floats32: []float32{float32(math.NaN()), float32(math.NaN())}} + f32AllOut, err := Norm(f32All, math.Inf(1), 0, false) + if err != nil { + t.Fatal(err) + } + if v := f32AllOut.FloatAt(0); !math.IsNaN(v) { + t.Errorf("Norm(float32 all-NaN, Inf) = %v, want NaN", v) + } + // Per-line behaviour on a 2-D array along both orientations. + g := mustFloats(t, []float64{math.NaN(), 2, 3, math.NaN()}, 2, 2) + byDim1, err := Norm(g, math.Inf(1), 1, false) + if err != nil { + t.Fatal(err) + } + if byDim1.FloatAt(0) != 2 || byDim1.FloatAt(1) != 3 { + t.Errorf("Norm dim 1 = [%v, %v], want [2, 3]", byDim1.FloatAt(0), byDim1.FloatAt(1)) + } + byDim0, err := Norm(g, math.Inf(1), 0, false) + if err != nil { + t.Fatal(err) + } + if byDim0.FloatAt(0) != 3 || byDim0.FloatAt(1) != 2 { + t.Errorf("Norm dim 0 = [%v, %v], want [3, 2]", byDim0.FloatAt(0), byDim0.FloatAt(1)) + } +} + +// ---------- Circular padding folds in constant time ---------- + +func TestPadCircularConstantTime(t *testing.T) { + // Small cases agree with a naive modulo fold, both directions. + src := mustFromInts(t, []int64{1, 2, 3}, 3) + got, err := Pad(src, []int{4, 5}, "circular", 0) + if err != nil { + t.Fatal(err) + } + for i := range 12 { + want := src.ints[((i-4)%3+3)%3] + if got.ints[i] != want { + t.Fatalf("circular pad [%d] = %d, want %d", i, got.ints[i], want) + } + } + // A pad far longer than the axis used to fold one step at a time, + // a quadratic walk: a pre-pad of two million on a one-element axis + // must land in constants, not hours. + one := mustFloats(t, []float64{7}, 1) + start := time.Now() + big, err := Pad(one, []int{2_000_000, 0}, "circular", 0) + if err != nil { + t.Fatal(err) + } + if elapsed := time.Since(start); elapsed > 30*time.Second { + t.Fatalf("circular pad of 2000000 on a length-1 axis took %s", elapsed) + } + if big.Len() != 2_000_001 { + t.Fatalf("padded length %d, want 2000001", big.Len()) + } + for _, i := range []int{0, 1, 999_999, 2_000_000} { + if v := big.FloatAt(i); v != 7 { + t.Fatalf("circular pad [%d] = %v, want 7", i, v) + } + } +} + +// ---------- BesselKn refuses every non-positive argument ---------- + +func TestBesselKnDomain(t *testing.T) { + neg := mustFloats(t, []float64{-1}, 1) + for _, n := range []int{0, 2} { + out, err := BesselKn(n, neg) + if err == nil { + t.Errorf("BesselKn(%d, -1) = %v, want an error", n, out.FloatAt(0)) + } else if !strings.Contains(err.Error(), "the argument must be positive") { + t.Errorf("BesselKn(%d, -1) error %q does not state the domain", n, err) + } else if !strings.HasPrefix(err.Error(), "tensor: BesselKn") { + t.Errorf("BesselKn(%d, -1) error %q lacks the prefixed name", n, err) + } + } + zero := mustFloats(t, []float64{0}, 1) + if _, err := BesselKn(0, zero); err == nil { + t.Error("BesselKn(0, 0) accepted the origin") + } + nan := mustFloats(t, []float64{math.NaN()}, 1) + if _, err := BesselKn(1, nan); err == nil { + t.Error("BesselKn(1, NaN) accepted NaN") + } + // The error names the offending element in a mixed array. + mixed := mustFloats(t, []float64{1, -2}, 2) + if _, err := BesselKn(3, mixed); err == nil || !strings.Contains(err.Error(), "element 1") { + t.Errorf("BesselKn(3, [1, -2]) error %v does not name element 1", err) + } + // Positive arguments are untouched: the recurrence still holds. + pos := mustFloats(t, []float64{1, 4}, 2) + for _, n := range []int{0, 1, 4} { + out, err := BesselKn(n, pos) + if err != nil { + t.Fatalf("BesselKn(%d, [1, 4]): %v", n, err) + } + for i, x := range []float64{1, 4} { + if v := out.FloatAt(i); !(v > 0) || math.IsInf(v, 1) { + t.Errorf("BesselKn(%d, %g) = %v, want a finite positive value", n, x, v) + } + } + } +} + +// ---------- SpMatMul accumulates float32 in float64 ---------- + +func TestSpMatMulFloat32Accumulation(t *testing.T) { + // One output cell fed by seven stored coordinates: the found case + // where narrowing every product before the float32 addition drifts + // ten ulps from the float64 reference, while widening once per + // coordinate stays within one. + vals := []float32{-0.72071517, 0.0010362695, 0.9238949, 0.0019417476, 205664.48, 0.8203619, 0.0019267823} + dens := []float32{-0.066897266, 0, 0.7552673, 0, 0, -0.8196033, 0} + indices := make([]int64, 0, 14) + for j := range vals { + indices = append(indices, 0, int64(j)) + } + idx, err := FromInts(indices, 7, 2) + if err != nil { + t.Fatal(err) + } + v32, err := FromFloat32s(vals, 7) + if err != nil { + t.Fatal(err) + } + s, err := NewSparseCOO(idx, v32, []int{1, 7}) + if err != nil { + t.Fatal(err) + } + d32, err := FromFloat32s(dens, 7, 1) + if err != nil { + t.Fatal(err) + } + got, err := SpMatMul(s, d32) + if err != nil { + t.Fatal(err) + } + if got.Dtype() != Float32 { + t.Fatalf("dtype %s, want float32", got.Dtype()) + } + // The widening walk: every coordinate widens, accumulates in float64 + // and narrows the cell on arrival. + exp := float32(0) + for j := range vals { + exp = float32(float64(exp) + float64(vals[j])*float64(dens[j])) + } + if bits := math.Float32bits(got.floats32[0]); bits != math.Float32bits(exp) { + t.Errorf("SpMatMul float32 = %v (%08x), want %v (%08x)", + got.floats32[0], bits, exp, math.Float32bits(exp)) + } + // Within two ulps of the float64 accumulation narrowed once. + acc := 0.0 + for j := range vals { + acc += float64(vals[j]) * float64(dens[j]) + } + ref := float32(acc) + if d := ulpDiff32(got.floats32[0], ref); d > 2 { + t.Errorf("SpMatMul float32 = %v sits %d ulps from the float64 reference %v, want at most 2", + got.floats32[0], d, ref) + } + // An int sparse times a float32 dense promotes and stays sane. + vi, err := FromInts([]int64{3, -1, 2}, 3) + if err != nil { + t.Fatal(err) + } + idxI, err := FromInts([]int64{0, 0, 0, 1, 0, 2}, 3, 2) + if err != nil { + t.Fatal(err) + } + si, err := NewSparseCOO(idxI, vi, []int{1, 3}) + if err != nil { + t.Fatal(err) + } + di, err := FromFloat32s([]float32{0.1, 0.2, 0.3}, 3, 1) + if err != nil { + t.Fatal(err) + } + gotI, err := SpMatMul(si, di) + if err != nil { + t.Fatal(err) + } + if gotI.Dtype() != Float32 { + t.Fatalf("mixed dtype %s, want float32", gotI.Dtype()) + } + if want := float32(3*0.1 - 0.2 + 2*0.3); math.Abs(float64(gotI.floats32[0]-want)) > 1e-6 { + t.Errorf("int×float32 sparse product = %v, want about %v", gotI.floats32[0], want) + } +} + +// ulpDiff32 counts the representable float32 steps between a and b. +func ulpDiff32(a, b float32) int { + return int(int64(math.Float32bits(a)) - int64(math.Float32bits(b))) +} + +// ---------- InterpolateMonotone refuses a NaN query ---------- + +func TestPCHIPNaNQuery(t *testing.T) { + xs := mustFloats(t, []float64{0, 1, 2}, 3) + ys := mustFloats(t, []float64{0, 1, 4}, 3) + q := mustFloats(t, []float64{0.5, math.NaN()}, 2) + out, err := InterpolateMonotone(xs, ys, q) + if err == nil { + t.Errorf("InterpolateMonotone with a NaN query = %v, want an error", out.Shape()) + } else if !strings.Contains(err.Error(), "NaN") { + t.Errorf("InterpolateMonotone NaN error %q does not name NaN", err) + } + fine, err := InterpolateMonotone(xs, ys, mustFloats(t, []float64{0.5, 1.9}, 2)) + if err != nil { + t.Fatal(err) + } + if v := fine.FloatAt(0); !(v > 0 && v < 1) { + t.Errorf("InterpolateMonotone(0.5) = %v, want a value inside (0, 1)", v) + } +} + +// ---------- Strided views read their own elements ---------- + +func TestStridedRawReads(t *testing.T) { + // Equal walks both arrays' own elements, not their payloads: two + // views carrying {4, 1} over different padding are equal, and one + // of them equals the dense array of the same elements. + v1 := &Array{shape: []int{2}, dt: Int, ints: []int64{4, 9, 1, 9}, strides: []int{2}} + v2 := &Array{shape: []int{2}, dt: Int, ints: []int64{4, 7, 1, 7}, strides: []int{2}} + dense, err := FromInts([]int64{4, 1}, 2) + if err != nil { + t.Fatal(err) + } + if !Equal(v1, v2) { + t.Error("Equal of two identical strided int views = false, want true") + } + if !Equal(v1, dense) { + t.Error("Equal of a strided int view and its dense twin = false, want true") + } + v3 := &Array{shape: []int{2}, dt: Int, ints: []int64{4, 9, 2, 9}, strides: []int{2}} + if Equal(v1, v3) { + t.Error("Equal of views holding {4, 1} and {4, 2} = true, want false") + } + // ArgMax/ArgMin compare the view's elements as ints. + got, err := ArgMax(v1) + if err != nil { + t.Fatal(err) + } + if got != 0 { + t.Errorf("ArgMax(strided int view of [4, 1]) = %d, want 0", got) + } + gotMin, err := ArgMin(v1) + if err != nil { + t.Fatal(err) + } + if gotMin != 1 { + t.Errorf("ArgMin(strided int view of [4, 1]) = %d, want 1", gotMin) + } + // Sort gathers the view's own float32 elements. + fv := &Array{shape: []int{3}, dt: Float32, floats32: []float32{9, 7, 1, 7, 2, 7}, strides: []int{2}} + sorted, err := Sort(fv) + if err != nil { + t.Fatal(err) + } + wantSorted, err := FromFloat32s([]float32{1, 2, 9}, 3) + if err != nil { + t.Fatal(err) + } + if !samePayload(t, sorted, wantSorted) { + t.Errorf("Sort(strided float32 view of [9, 1, 2]) = %v, want [1, 2, 9]", sorted.floats32) + } +} + +// ---------- The constructors refuse unknown dtypes ---------- + +func TestUnknownDtypeRejected(t *testing.T) { + if _, err := Zeros(Dtype(99), 2); err == nil { + t.Error("Zeros(Dtype(99), 2) created a payload") + } + if _, err := Ones(Dtype(200), 2); err == nil { + t.Error("Ones(Dtype(200), 2) created a payload") + } + if New(Dtype(99), 2) != nil { + t.Error("New(Dtype(99), 2) returned an array, want nil") + } + for _, dt := range []Dtype{Int, Float32, Float, Complex} { + if _, err := Zeros(dt, 2); err != nil { + t.Errorf("Zeros(%s, 2): %v", dt, err) + } + } +} + +// ---------- Scalar.String prints one sign ---------- + +func TestScalarString(t *testing.T) { + c := Scalar{isComplex: true, c: complex(4, -2)} + if got := c.String(); got != "complex (4-2i)" { + t.Errorf("Scalar.String() = %q, want %q", got, "complex (4-2i)") + } + pos := Scalar{isComplex: true, c: complex(0, 2.5)} + if got := pos.String(); got != "complex (0+2.5i)" { + t.Errorf("Scalar.String() = %q, want %q", got, "complex (0+2.5i)") + } + if strings.Contains(Scalar{isComplex: true, c: complex(1, -1)}.String(), "+-") { + t.Error("Scalar.String() still prints +-") + } + if got := (Scalar{i: 7}).String(); got != "int 7" { + t.Errorf("Scalar.String() = %q, want %q", got, "int 7") + } + if got := (Scalar{isFloat: true, f: 2.5}).String(); got != "float 2.5" { + t.Errorf("Scalar.String() = %q, want %q", got, "float 2.5") + } +} + +// ---------- Mean survives a sum that would overflow ---------- + +func TestMeanOverflow(t *testing.T) { + maxs := mustFloats(t, []float64{math.MaxFloat64, math.MaxFloat64}, 2) + m, err := Mean(maxs) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(m) != math.Float64bits(math.MaxFloat64) { + t.Errorf("Mean of two MaxFloat64s = %v, want %v", m, math.MaxFloat64) + } + threes := mustFloats(t, []float64{math.MaxFloat64, math.MaxFloat64, math.MaxFloat64}, 3) + m3, err := Mean(threes) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(m3) != math.Float64bits(math.MaxFloat64) { + t.Errorf("Mean of three MaxFloat64s = %v, want %v", m3, math.MaxFloat64) + } + neg := mustFloats(t, []float64{-math.MaxFloat64, -math.MaxFloat64}, 2) + mn, err := Mean(neg) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(mn) != math.Float64bits(-math.MaxFloat64) { + t.Errorf("Mean of two −MaxFloat64s = %v, want %v", mn, -math.MaxFloat64) + } + // Ordinary inputs keep the plain Sum rounding bit for bit. + plain := mustFloats(t, []float64{0.1, 0.2, 0.3, 700, -3.25}, 5) + mp, err := Mean(plain) + if err != nil { + t.Fatal(err) + } + if want := Sum(plain).Float() / 5; math.Float64bits(mp) != math.Float64bits(want) { + t.Errorf("Mean of ordinary input = %v, want the Sum rounding %v", mp, want) + } + // Int arrays keep their true-division mean. + mi, err := Mean(mustFromInts(t, []int64{1, 2}, 2)) + if err != nil { + t.Fatal(err) + } + if mi != 1.5 { + t.Errorf("Mean of int [1, 2] = %v, want 1.5", mi) + } + // A NaN payload still reaches the caller as NaN, and an infinite + // element keeps the plain path's infinite mean. + if mNaN, err := Mean(mustFloats(t, []float64{math.NaN(), 1}, 2)); err != nil || !math.IsNaN(mNaN) { + t.Errorf("Mean of [NaN, 1] = %v, %v, want NaN", mNaN, err) + } + if mInf, err := Mean(mustFloats(t, []float64{math.Inf(1), 1}, 2)); err != nil || !math.IsInf(mInf, 1) { + t.Errorf("Mean of [+Inf, 1] = %v, %v, want +Inf", mInf, err) + } +} + +// ---------- Coverage pins (no fix behind them) ---------- + +// TestPinEinsumBatchedVecMat walks the two batched +// matrix-vector patterns against a naive loop, for int and float32, +// with one batch and with several. +func TestPinEinsumBatchedVecMat(t *testing.T) { + rng := rand.New(rand.NewPCG(606, 6)) + for _, batches := range []int{1, 3} { + n, k, m := 2, 3, 4 + for _, dt := range []Dtype{Int, Float32} { + a := randWhole(t, rng, dt, batches, n, k) // (b, n, k) + bm := randWhole(t, rng, dt, batches, k, m) // (b, k, m) + v := randWhole(t, rng, dt, batches, k) // (b, k) + mv, err := Einsum("bij,bj->bi", a, v) + if err != nil { + t.Fatalf("bij,bj->bi (%d batches, %s): %v", batches, dt, err) + } + vm, err := Einsum("bi,bij->bj", v, bm) + if err != nil { + t.Fatalf("bi,bij->bj (%d batches, %s): %v", batches, dt, err) + } + for bt := range batches { + for i := range n { + var wantF float64 + var wantI int64 + for p := range k { + av := int64At(a, (bt*n+i)*k+p) + bv := int64At(v, bt*k+p) + if dt == Int { + wantI += av * bv + } else { + wantF += float64(av) * float64(bv) + } + } + if dt == Int { + if mv.ints[bt*n+i] != wantI { + t.Errorf("bij,bj->bi int batch %d row %d = %d, want %d", bt, i, mv.ints[bt*n+i], wantI) + } + } else if d := math.Abs(float64(mv.floats32[bt*n+i]) - wantF); d > 1e-5*judgingScale(wantF) { + t.Errorf("bij,bj->bi float32 batch %d row %d = %v, want %v", bt, i, mv.floats32[bt*n+i], wantF) + } + } + for j := range m { + var wantF float64 + var wantI int64 + for p := range k { + bv := int64At(v, bt*k+p) + mav := int64At(bm, (bt*k+p)*m+j) + if dt == Int { + wantI += bv * mav + } else { + wantF += float64(bv) * float64(mav) + } + } + if dt == Int { + if vm.ints[bt*m+j] != wantI { + t.Errorf("bi,bij->bj int batch %d col %d = %d, want %d", bt, j, vm.ints[bt*m+j], wantI) + } + } else if d := math.Abs(float64(vm.floats32[bt*m+j]) - wantF); d > 1e-5*judgingScale(wantF) { + t.Errorf("bi,bij->bj float32 batch %d col %d = %v, want %v", bt, j, vm.floats32[bt*m+j], wantF) + } + } + } + } + } +} + +// judgingScale grows the tolerance with the magnitude it guards. +func judgingScale(want float64) float64 { + s := math.Abs(want) + if s < 1 { + return 1 + } + return s +} + +// TestPinScanDimTypes walks CumSum and CumProd for int, +// float32 and complex against step-by-step references. +func TestPinScanDimTypes(t *testing.T) { + iv := []int64{3, -1, 2, 5, 0, 7} + ia := mustFromInts(t, iv, 2, 3) + cs, err := CumSum(ia, 1) + if err != nil { + t.Fatal(err) + } + // Row-major (2, 3): row 0 scans 3, -1, 2 and row 1 scans 5, 0, 7. + wantRow := [][]int64{{3, 2, 4}, {5, 5, 12}} + for i := range 2 { + for j := range 3 { + if cs.ints[i*3+j] != wantRow[i][j] { + t.Errorf("CumSum int [%d][%d] = %d, want %d", i, j, cs.ints[i*3+j], wantRow[i][j]) + } + } + } + cp, err := CumProd(ia, 0) + if err != nil { + t.Fatal(err) + } + for j := range 3 { + if cp.ints[j] != iv[j] { + t.Errorf("CumProd int first row [%d] = %d, want %d", j, cp.ints[j], iv[j]) + } + if cp.ints[3+j] != iv[j]*iv[3+j] { + t.Errorf("CumProd int second row [%d] = %d, want %d", j, cp.ints[3+j], iv[j]*iv[3+j]) + } + } + // Float32 narrows the carry once per step, exactly as the stored + // value feeds the next one. + fv := []float32{0.5, -1.25, 3.5, 2, -0.125, 4} + fa, err := FromFloat32s(fv, 2, 3) + if err != nil { + t.Fatal(err) + } + fcs, err := CumSum(fa, 0) + if err != nil { + t.Fatal(err) + } + for j := range 3 { + carry := float64(fv[j]) + if math.Float32bits(fcs.floats32[j]) != math.Float32bits(float32(carry)) { + t.Errorf("CumSum float32 [%d] = %v, want %v", j, fcs.floats32[j], float32(carry)) + } + for i := 1; i < 2; i++ { + carry = float64(float32(carry)) + float64(fv[i*3+j]) + if math.Float32bits(fcs.floats32[i*3+j]) != math.Float32bits(float32(carry)) { + t.Errorf("CumSum float32 [%d] = %v, want %v", i*3+j, fcs.floats32[i*3+j], float32(carry)) + } + } + } + fcp, err := CumProd(fa, 1) + if err != nil { + t.Fatal(err) + } + for i := range 2 { + carry := float32(1) + for j := range 3 { + carry = float32(float64(carry) * float64(fv[i*3+j])) + if math.Float32bits(fcp.floats32[i*3+j]) != math.Float32bits(carry) { + t.Errorf("CumProd float32 [%d] = %v, want %v", i*3+j, fcp.floats32[i*3+j], carry) + } + } + } + // Complex adds exactly. + cv := []complex128{1 + 1i, 2, -1i, 1} + ca := mustFromComplexes(t, cv, 2, 2) + ccs, err := CumSum(ca, 1) + if err != nil { + t.Fatal(err) + } + if ccs.complexes[0] != 1+1i || ccs.complexes[1] != 3+1i || + ccs.complexes[2] != -1i || ccs.complexes[3] != 1-1i { + t.Errorf("CumSum complex = %v", ccs.complexes) + } + // CumProd along dim 1 accumulates within each row. + ccp, err := CumProd(ca, 1) + if err != nil { + t.Fatal(err) + } + // Row 0: (1+1i)·2 = 2+2i; row 1: (−1i)·1 = −1i. + if ccp.complexes[0] != 1+1i || ccp.complexes[1] != 2+2i || + ccp.complexes[2] != -1i || ccp.complexes[3] != -1i { + t.Errorf("CumProd complex = %v", ccp.complexes) + } +} + +// TestPinMeanAxisEmptyDim pins the NaN fill a mean over an +// empty dimension reports. +func TestPinMeanAxisEmptyDim(t *testing.T) { + a, err := Zeros(Float, 3, 0, 4) + if err != nil { + t.Fatal(err) + } + m, err := MeanAxis(a, 1) + if err != nil { + t.Fatal(err) + } + if !sameShape(m.Shape(), []int{3, 4}) { + t.Fatalf("MeanAxis over an empty dim has shape %s, want (3, 4)", shapeText(m.Shape())) + } + for i := range m.Len() { + if !math.IsNaN(m.floats[i]) { + t.Errorf("MeanAxis over an empty dim [%d] = %v, want NaN", i, m.floats[i]) + } + } +} + +// TestPinEinsumSlotSumMultiOperand walks the three- and +// four-operand slot sums against naive loops, float64 and float32. +func TestPinEinsumSlotSumMultiOperand(t *testing.T) { + rng := rand.New(rand.NewPCG(607, 7)) + for _, dt := range []Dtype{Float, Float32} { + a := randWhole(t, rng, dt, 3, 4) + b := randWhole(t, rng, dt, 4, 5) + c := randWhole(t, rng, dt, 5, 6) + got3, err := Einsum("ik,kj,jl->il", a, b, c) + if err != nil { + t.Fatalf("3 operands %s: %v", dt, err) + } + for i := range 3 { + for l := range 6 { + var want float64 + for k := range 4 { + for j := range 5 { + want += float64At(a, i*4+k) * float64At(b, k*5+j) * float64At(c, j*6+l) + } + } + g := float64At(got3, i*6+l) + if math.Abs(g-want) > 1e-9*judgingScale(want) { + t.Errorf("3 operands %s [%d,%d] = %v, want %v", dt, i, l, g, want) + } + } + } + d := randWhole(t, rng, dt, 6, 2) + got4, err := Einsum("ik,kj,jl,lm->im", a, b, c, d) + if err != nil { + t.Fatalf("4 operands %s: %v", dt, err) + } + for i := range 3 { + for m := range 2 { + var want float64 + for k := range 4 { + for j := range 5 { + for l := range 6 { + want += float64At(a, i*4+k) * float64At(b, k*5+j) * + float64At(c, j*6+l) * float64At(d, l*2+m) + } + } + } + g := float64At(got4, i*2+m) + tol := 1e-9 + if dt == Float32 { + tol = 1e-5 + } + if math.Abs(g-want) > tol*judgingScale(want) { + t.Errorf("4 operands %s [%d,%d] = %v, want %v", dt, i, m, g, want) + } + } + } + } +} + +// randWhole builds a small random array of the given dtype, holding +// only whole numbers so the naive references stay exact. +func randWhole(t *testing.T, rng *rand.Rand, dt Dtype, shape ...int) *Array { + t.Helper() + n := 1 + for _, d := range shape { + n *= d + } + switch dt { + case Int: + vals := make([]int64, n) + for i := range vals { + vals[i] = int64(rng.IntN(7) - 3) + } + return mustFromInts(t, vals, shape...) + case Float32: + vals := make([]float32, n) + for i := range vals { + vals[i] = float32(rng.IntN(7) - 3) + } + a, err := FromFloat32s(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a + default: + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(rng.IntN(7) - 3) + } + return mustFloats(t, vals, shape...) + } +} + +// int64At reads element i as an int64 magnitudes-only value. +func int64At(a *Array, i int) int64 { + switch a.Dtype() { + case Int: + return a.ints[i] + case Float32: + return int64(a.floats32[i]) + default: + return int64(a.floats[i]) + } +} + +// float64At reads element i as a float64 value. +func float64At(a *Array, i int) float64 { + switch a.Dtype() { + case Float32: + return float64(a.floats32[i]) + default: + return a.floatAt(i) + } +} diff --git a/internal/core/index.go b/internal/core/index.go new file mode 100644 index 0000000..0bf5d0a --- /dev/null +++ b/internal/core/index.go @@ -0,0 +1,233 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +// Indexing and slicing. Every read and write goes through a +// validated multi-dimensional index; updates are functional: +// WithInt, WithFloat and WithComplex return a new array and leave the +// receiver untouched. Slice returns a read-only view when the selection +// is contiguous and a copy otherwise; either way the source is never +// written through. + +// IntAt returns the element at the given index as an int64; the array +// must be an integer-class dtype, bool included, and every widening is +// exact. A float or complex array is refused with the wording the tests +// pin. +func IntAt(a *Array, index ...int) (int64, error) { + if !intClass(a.dt) { + return 0, errf("IntAt: array is %s, not int", a.dt) + } + off, err := flatIndex(a.shape, index) + if err != nil { + return 0, err + } + return a.intAt(off), nil +} + +// FloatAt returns the element at the given index as a float64; every +// real dtype widens exactly, bool included. A complex array is refused +// with the wording the tests pin. +func FloatAt(a *Array, index ...int) (float64, error) { + if a.dt == Complex { + return 0, errf("FloatAt: array is %s, not float", a.dt) + } + off, err := flatIndex(a.shape, index) + if err != nil { + return 0, err + } + return a.floatAt(off), nil +} + +// ComplexAt returns the element at the given index; the array must be +// complex. +func ComplexAt(a *Array, index ...int) (complex128, error) { + if a.dt != Complex { + return 0, errf("ComplexAt: array is %s, not complex", a.dt) + } + off, err := flatIndex(a.shape, index) + if err != nil { + return 0, err + } + return a.complexes[a.physIndex(off)], nil +} + +// WithInt returns a new array with v stored at the given index in the +// receiver's own dtype; the receiver is unchanged. Every integer-class +// dtype is written, bool included: the store narrows with Go's +// conversion, the implicit-store cast the ladder has always carried, and +// bool stores v != 0. A float or complex receiver is refused with the +// wording the tests pin. +func WithInt(a *Array, v int64, index ...int) (*Array, error) { + if !intClass(a.dt) { + return nil, errf("WithInt: array is %s, not int", a.dt) + } + off, err := flatIndex(a.shape, index) + if err != nil { + return nil, err + } + out := a.cloneArray() + switch a.dt { + case Int: + out.ints[off] = v + case Bool: + out.bools[off] = v != 0 + case Int8: + out.i8s[off] = int8(v) + case Uint8: + out.u8s[off] = uint8(v) + case Int16: + out.i16s[off] = int16(v) + case Uint16: + out.u16s[off] = uint16(v) + case Int32: + out.i32s[off] = int32(v) + default: + out.u32s[off] = uint32(v) + } + return out, nil +} + +// WithFloat returns a new float array with v at the given index; the +// receiver is unchanged. It writes the float dtype only: the narrow +// float widths and the integer dtypes keep the pinned refusals. +func WithFloat(a *Array, v float64, index ...int) (*Array, error) { + if a.dt != Float { + return nil, errf("WithFloat: array is %s, not float", a.dt) + } + off, err := flatIndex(a.shape, index) + if err != nil { + return nil, err + } + out := a.cloneArray() + out.floats[off] = v + return out, nil +} + +// WithComplex returns a new complex array with v at the given index; the +// receiver is unchanged. It writes the complex dtype only, the same +// kind-gate rule WithFloat carries. +func WithComplex(a *Array, v complex128, index ...int) (*Array, error) { + if a.dt != Complex { + return nil, errf("WithComplex: array is %s, not complex", a.dt) + } + off, err := flatIndex(a.shape, index) + if err != nil { + return nil, err + } + out := a.cloneArray() + out.complexes[off] = v + return out, nil +} + +// Slice selects the half-open range [start, stop) along one dimension. +// +// When the selection is a contiguous region of the source, the slice +// spans whole rows of it (dim 0) or takes the dimension in full, the +// result is a read-only view sharing the source's storage: the payload +// is rebased to the first selected element and the shape narrowed, so +// payload[i] is still element i and every kernel treats the view as an +// ordinary dense array. Any other selection (an interior column range +// of a 2-D array, say) is copied, because giving it a strided view +// needs the kernel-boundary materialisation audit described in +// docs/ARCHITECTURE.md. Either way the receiver is never written +// through, and a view never aliases for writing: nothing in the +// library writes through an array it did not allocate, and optimiser +// updates route through materialised parameters. +func Slice(a *Array, dim, start, stop int) (*Array, error) { + if dim < 0 || dim >= len(a.shape) { + return nil, errf("Slice: dimension %d is out of range for shape %s", dim, shapeText(a.shape)) + } + if start < 0 || stop > a.shape[dim] || start > stop { + return nil, errf("Slice: range [%d:%d] is out of range for dimension %d of size %d", + start, stop, dim, a.shape[dim]) + } + newShape := a.Shape() + newShape[dim] = stop - start + // A strided source cannot be sliced by payload arithmetic, so + // reduce it to a dense array first; the copy path below and the + // view path both assume payload[i] is element i. + if !a.isContiguous() { + a = a.materialise() + } + // Physical offset of the first selected element, plus the block of + // elements that one index along dim stands for. + tail := 1 + for d := dim + 1; d < a.NDim(); d++ { + tail *= a.shape[d] + } + head := 1 + for d := range dim { + head *= a.shape[d] + } + if head == 1 || stop-start == a.shape[dim] { + view := &Array{shape: newShape, dt: a.dt} + view.rebase(a, start*tail) + return view, nil + } + total := 1 + for _, d := range newShape { + total *= d + } + out := &Array{shape: newShape, dt: a.dt} + out.alloc(total) + // The selection is one contiguous run of (stop-start) blocks per + // position along the leading dimension, and the source above is + // dense, so each row moves whole rather than as gathered offsets. + // copyRun dispatches every dtype the package stores, the narrow + // payloads included. + run := (stop - start) * tail + stride := a.shape[dim] * tail + for r := range head { + copyRun(out, a, r*run, r*stride+start*tail, run) + } + return out, nil +} + +// rebase points a's payload at element off of src's payload, sharing +// the storage. off must be a valid payload index of src; the elements +// beyond the view's own count stay invisible because Len is shape-based. +func (a *Array) rebase(src *Array, off int) { + switch a.dt { + case Int: + a.ints = src.ints[off:] + case Float16: + a.halves = src.halves[off:] + case Float32: + a.floats32 = src.floats32[off:] + case Float: + a.floats = src.floats[off:] + case Complex: + a.complexes = src.complexes[off:] + case Bool: + a.bools = src.bools[off:] + case Int8: + a.i8s = src.i8s[off:] + case Uint8: + a.u8s = src.u8s[off:] + case Int16: + a.i16s = src.i16s[off:] + case Uint16: + a.u16s = src.u16s[off:] + case Int32: + a.i32s = src.i32s[off:] + default: + a.u32s = src.u32s[off:] + } +} + +// flatIndex validates a multi-dimensional index against a shape and +// returns the row-major flat offset. +func flatIndex(shape []int, index []int) (int, error) { + if len(index) != len(shape) { + return 0, errf("index %v does not match the shape %s", index, shapeText(shape)) + } + off := 0 + for d, i := range index { + if i < 0 || i >= shape[d] { + return 0, errf("index %d is out of range for dimension %d of size %d", i, d, shape[d]) + } + off = off*shape[d] + i + } + return off, nil +} diff --git a/internal/core/index_test.go b/internal/core/index_test.go new file mode 100644 index 0000000..252f4e0 --- /dev/null +++ b/internal/core/index_test.go @@ -0,0 +1,317 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "slices" + "strings" + "testing" +) + +func TestIndexing(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + if v, err := IntAt(a, 0, 0); err != nil || v != 1 { + t.Fatalf("IntAt(0,0): %d %v", v, err) + } + if v, err := IntAt(a, 1, 2); err != nil || v != 6 { + t.Fatalf("IntAt(1,2): %d %v", v, err) + } + // Row-major storage: (0,2) is the third element. + if v, _ := IntAt(a, 0, 2); v != 3 { + t.Fatalf("IntAt(0,2): %d", v) + } + + f := mustFromFloats(t, []float64{1.5}, 1) + if _, err := IntAt(f, 0); err == nil || !strings.Contains(err.Error(), "not int") { + t.Fatalf("IntAt on float: %v", err) + } + // FloatAt widens every real dtype, int included: the widening-reader + // contract the dtype surface settled. + if v, err := FloatAt(a, 0, 0); err != nil || v != 1 { + t.Fatalf("FloatAt on int: %v %v", v, err) + } + if _, err := IntAt(a, 2, 0); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("index out of range: %v", err) + } + if _, err := IntAt(a, 0); err == nil || !strings.Contains(err.Error(), "does not match the shape") { + t.Fatalf("wrong arity: %v", err) + } +} + +func TestFunctionalUpdate(t *testing.T) { + a := mustFromInts(t, []int64{1, 2}, 2) + updated, err := WithInt(a, 9, 0) + if err != nil { + t.Fatalf("WithInt: %v", err) + } + if v, _ := IntAt(updated, 0); v != 9 { + t.Fatalf("updated: %d", v) + } + // The receiver stays untouched. + if v, _ := IntAt(a, 0); v != 1 { + t.Fatalf("receiver mutated: %d", v) + } + + f := mustFromFloats(t, []float64{1.5, 2.5}, 2) + uf, err := WithFloat(f, 9.5, 1) + if err != nil { + t.Fatalf("WithFloat: %v", err) + } + if v, _ := FloatAt(uf, 1); v != 9.5 { + t.Fatalf("WithFloat value: %v", v) + } + if v, _ := FloatAt(f, 1); v != 2.5 { + t.Fatalf("WithFloat receiver mutated: %v", v) + } + + if _, err := WithInt(f, 1, 0); err == nil || !strings.Contains(err.Error(), "not int") { + t.Fatalf("WithInt on float: %v", err) + } + if _, err := WithFloat(a, 1.0, 0); err == nil || !strings.Contains(err.Error(), "not float") { + t.Fatalf("WithFloat on int: %v", err) + } + if _, err := WithInt(a, 1, 5); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("WithInt range: %v", err) + } +} + +func TestSlice(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4}, 4) + + s, err := Slice(a, 0, 1, 3) + if err != nil { + t.Fatalf("Slice: %v", err) + } + if s.Len() != 2 { + t.Fatalf("Slice len: %d", s.Len()) + } + if v, _ := IntAt(s, 0); v != 2 { + t.Fatalf("Slice(0): %d", v) + } + if v, _ := IntAt(s, 1); v != 3 { + t.Fatalf("Slice(1): %d", v) + } + + // A functional update through the view leaves the source untouched: + // WithInt copies rather than writing the shared payload. + updated, _ := WithInt(s, 99, 0) + if v, _ := IntAt(a, 1); v != 2 { + t.Fatalf("WithInt through a slice view mutated the source: %d", v) + } + if v, _ := IntAt(updated, 0); v != 99 { + t.Fatalf("updated slice: %d", v) + } + + // Slicing a dimension of a 2-D array picks whole rows. + m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + rows, err := Slice(m, 0, 1, 2) + if err != nil { + t.Fatalf("Slice rows: %v", err) + } + want := mustFromInts(t, []int64{4, 5, 6}, 1, 3) + if !Equal(want, rows) { + t.Fatalf("rows: %s", rows) + } + + cols, err := Slice(m, 1, 0, 2) + if err != nil { + t.Fatalf("Slice cols: %v", err) + } + wantCols := mustFromInts(t, []int64{1, 2, 4, 5}, 2, 2) + if !Equal(wantCols, cols) { + t.Fatalf("cols: %s", cols) + } + + // An empty range yields an empty array of the right shape. + none, err := Slice(a, 0, 2, 2) + if err != nil || none.Len() != 0 || none.Shape()[0] != 0 { + t.Fatalf("empty slice: %s %v", none, err) + } + + if _, err := Slice(a, 1, 0, 1); err == nil || !strings.Contains(err.Error(), "dimension 1 is out of range") { + t.Fatalf("Slice dim: %v", err) + } + if _, err := Slice(a, 0, 3, 2); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("Slice reversed: %v", err) + } + if _, err := Slice(a, 0, 0, 5); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("Slice beyond: %v", err) + } +} + +// TestWideningReadersAndUpdates pins the widening-reader contract of +// the dtypes: IntAt reads the integer class and bool with exact widenings, +// FloatAt reads every real dtype, ComplexAt keeps its gate, WithInt +// writes every integer-class dtype with the implicit-store cast, and the +// kind gates the suite pins stay loud. +func TestWideningReadersAndUpdates(t *testing.T) { + i8, err := FromInt8s([]int8{-2, 3}, 2) + if err != nil { + t.Fatal(err) + } + bl, err := FromBools([]bool{true, false}, 2) + if err != nil { + t.Fatal(err) + } + if v, err := IntAt(i8, 0); err != nil || v != -2 { + t.Fatalf("IntAt int8: %d %v", v, err) + } + if v, err := IntAt(bl, 0); err != nil || v != 1 { + t.Fatalf("IntAt bool: %d %v", v, err) + } + if v, err := FloatAt(i8, 1); err != nil || v != 3 { + t.Fatalf("FloatAt int8: %v %v", v, err) + } + if v, err := FloatAt(bl, 0); err != nil || v != 1 { + t.Fatalf("FloatAt bool: %v %v", v, err) + } + f32, err := FromFloat32s([]float32{1.5}, 1) + if err != nil { + t.Fatal(err) + } + if v, err := FloatAt(f32, 0); err != nil || v != 1.5 { + t.Fatalf("FloatAt float32: %v %v", v, err) + } + if _, err := ComplexAt(i8, 0); err == nil || !strings.Contains(err.Error(), "not complex") { + t.Fatalf("ComplexAt int8: %v", err) + } + + // WithInt writes the integer class with the implicit-store cast. + w, err := WithInt(i8, 300, 1) + if err != nil { + t.Fatalf("WithInt int8: %v", err) + } + stored := int64(300) + if got := w.RawInt8s()[1]; got != int8(stored) { + t.Fatalf("WithInt int8 = %d, want the wrapped %d", got, int8(stored)) + } + if got := i8.RawInt8s()[1]; got != 3 { + t.Fatalf("WithInt int8 mutated its receiver: %d", got) + } + wb, err := WithInt(bl, 2, 0) + if err != nil || !wb.RawBools()[0] { + t.Fatalf("WithInt bool(2): %v %v", wb.RawBools(), err) + } + wb, err = WithInt(bl, 0, 1) + if err != nil || wb.RawBools()[1] { + t.Fatalf("WithInt bool(0): %v %v", wb.RawBools(), err) + } + // The kind gates the suite pins stay loud for the other receivers. + if _, err := WithFloat(i8, 1, 0); err == nil || !strings.Contains(err.Error(), "not float") { + t.Fatalf("WithFloat int8: %v", err) + } + if _, err := WithComplex(i8, 1, 0); err == nil || !strings.Contains(err.Error(), "not complex") { + t.Fatalf("WithComplex int8: %v", err) + } + + // Slice carries the narrow dtypes: the view path rebases the narrow + // payload and the copy path walks setFrom. + m8, err := FromInt8s([]int8{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + view, err := Slice(m8, 0, 0, 1) + if err != nil { + t.Fatalf("Slice int8 view: %v", err) + } + if v, err := IntAt(view, 0, 2); err != nil || v != 3 { + t.Fatalf("Slice int8 view element: %d %v", v, err) + } + cols, err := Slice(m8, 1, 0, 2) + if err != nil { + t.Fatalf("Slice int8 copy: %v", err) + } + if want := []int8{1, 2, 4, 5}; !slices.Equal(cols.RawInt8s(), want) { + t.Fatalf("Slice int8 copy = %v, want %v", cols.RawInt8s(), want) + } +} + +// TestAdvancedIndexingNarrowDtypes pins the advanced-indexing contract: +// Scatter builds its output through cloneArray so narrow and bool +// destinations carry payloads, Gather and Take read and write the narrow +// payloads, and Nonzero sees zeros in them. +func TestAdvancedIndexingNarrowDtypes(t *testing.T) { + idx, err := FromInts([]int64{0, 1}, 2, 1) + if err != nil { + t.Fatal(err) + } + selfB, err := FromBools([]bool{false, true, false, true}, 2, 2) + if err != nil { + t.Fatal(err) + } + srcB, err := FromBools([]bool{true, false}, 2, 1) + if err != nil { + t.Fatal(err) + } + out, err := Scatter(selfB, 1, idx, srcB) + if err != nil { + t.Fatalf("Scatter bool: %v", err) + } + if want := []bool{true, true, false, false}; !slices.Equal(out.RawBools(), want) { + t.Fatalf("Scatter bool = %v, want %v", out.RawBools(), want) + } + + // An int8 destination from an int source stores with the + // implicit-store cast, exactly as setConverted stores it. + self8, err := FromInt8s([]int8{1, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatal(err) + } + srcI, err := FromInts([]int64{300, 9}, 2, 1) + if err != nil { + t.Fatal(err) + } + out8, err := Scatter(self8, 1, idx, srcI) + if err != nil { + t.Fatalf("Scatter int8 from int: %v", err) + } + wrap := srcI.RawInts()[0] + if want := []int8{int8(wrap), 2, 3, 9}; !slices.Equal(out8.RawInt8s(), want) { + t.Fatalf("Scatter int8 from int = %v, want %v", out8.RawInt8s(), want) + } + + src8, err := FromInt8s([]int8{5, 6, 7, 8}, 2, 2) + if err != nil { + t.Fatal(err) + } + gidx, err := FromInts([]int64{1, 0}, 2, 1) + if err != nil { + t.Fatal(err) + } + g, err := Gather(src8, 1, gidx) + if err != nil { + t.Fatalf("Gather int8: %v", err) + } + if want := []int8{6, 7}; !slices.Equal(g.RawInt8s(), want) { + t.Fatalf("Gather int8 = %v, want %v", g.RawInt8s(), want) + } + + tb, err := FromBools([]bool{true, false, true}, 3) + if err != nil { + t.Fatal(err) + } + tidx, err := FromInts([]int64{2, 0}, 2) + if err != nil { + t.Fatal(err) + } + tk, err := Take(tb, tidx) + if err != nil { + t.Fatalf("Take bool: %v", err) + } + if want := []bool{true, true}; !slices.Equal(tk.RawBools(), want) { + t.Fatalf("Take bool = %v, want %v", tk.RawBools(), want) + } + + nz8, err := FromInt8s([]int8{0, 2, 0, 0}, 4) + if err != nil { + t.Fatal(err) + } + nz, err := Nonzero(nz8) + if err != nil { + t.Fatalf("Nonzero int8: %v", err) + } + if want := []int{1}; !slices.Equal(nz[0], want) { + t.Fatalf("Nonzero int8 = %v, want %v", nz[0], want) + } +} diff --git a/internal/core/indexing2.go b/internal/core/indexing2.go new file mode 100644 index 0000000..3800dcf --- /dev/null +++ b/internal/core/indexing2.go @@ -0,0 +1,395 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sync" + +// Advanced indexing: Gather, Scatter, Nonzero, Take. The first three +// operate along a chosen dimension; Take gathers along the flat axis. +// Indices are int64 arrays; out-of-range indices are an error naming +// the offending position. Gather and Scatter read and write through the +// standard accessors, so a source may sit anywhere on the promotion +// ladder: complex works the same way as the others, except for the one +// narrowing the library rejects (complex into int or float32, Astype's +// rule). + +// Gather returns a copy indexed by index along dim: out[i_0, …, i_dim, +// …, i_N] = self[i_0, …, index[i_…, j, …], …, i_N]. The output shape +// is the index shape (after the dim is dropped from self). +func Gather(src *Array, dim int, index *Array) (*Array, error) { + if index.dt != Int { + return nil, errf("Gather: index must be an int array, got %s", index.dt) + } + if dim < 0 || dim >= src.NDim() { + return nil, errf("Gather: dimension %d out of range for shape %s", dim, shapeText(src.shape)) + } + // All non-dim dimensions of src must match the corresponding index dim. + if !gatherCompatible(src.shape, index.shape, dim) { + return nil, errf("Gather: shape mismatch %s vs %s on dim %d", shapeText(src.shape), shapeText(index.shape), dim) + } + out := &Array{shape: index.Shape(), dt: src.dt} + n := index.Len() + out.alloc(n) + if n == 0 { + return out, nil + } + // The source is read through the same accessors the element walk + // used, so a strided source materialises to exactly the values those + // reads produced; a contiguous source is its own payload. + ss := src.materialise() + // The workers write disjoint output slots. The serial walk reported + // the first offending position, so the workers publish the smallest + // one under the same lock-free-on-success rule Interpolate2D uses. + var mu sync.Mutex + firstPos := n + 1 + var firstVal int + fail := func(pos, idx int) { + mu.Lock() + defer mu.Unlock() + if pos < firstPos { + firstPos, firstVal = pos, idx + } + } + ndim := index.NDim() + switch out.dt { + case Int: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.ints, ss.ints, index, ss.shape, s, e, dim, ndim, fail) + }) + case Bool: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.bools, ss.bools, index, ss.shape, s, e, dim, ndim, fail) + }) + case Int8: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.i8s, ss.i8s, index, ss.shape, s, e, dim, ndim, fail) + }) + case Uint8: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.u8s, ss.u8s, index, ss.shape, s, e, dim, ndim, fail) + }) + case Int16: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.i16s, ss.i16s, index, ss.shape, s, e, dim, ndim, fail) + }) + case Uint16: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.u16s, ss.u16s, index, ss.shape, s, e, dim, ndim, fail) + }) + case Int32: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.i32s, ss.i32s, index, ss.shape, s, e, dim, ndim, fail) + }) + case Uint32: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.u32s, ss.u32s, index, ss.shape, s, e, dim, ndim, fail) + }) + case Float16: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.halves, ss.halves, index, ss.shape, s, e, dim, ndim, fail) + }) + case Float32: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.floats32, ss.floats32, index, ss.shape, s, e, dim, ndim, fail) + }) + case Float: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.floats, ss.floats, index, ss.shape, s, e, dim, ndim, fail) + }) + default: + parallelMin(n, copyMinPerWorker, func(s, e int) { + gatherFill(out.complexes, ss.complexes, index, ss.shape, s, e, dim, ndim, fail) + }) + } + if firstPos <= n { + return nil, errf("Gather: index %d out of range for dimension %d of size %d at position %d", + firstVal, dim, src.shape[dim], firstPos) + } + return out, nil +} + +// gatherFill writes dst[i] = src[off] for the worker's range of the +// index walk, where off folds the output's own coordinates with the +// gathered axis replaced by the index. An out-of-range index reports +// through fail and skips the write; a call that failed anywhere +// discards its output, so the skipped slot never surfaces. The odometer +// seeds from the flat start, so the chunks cover disjoint slots in the +// same coordinates one serial walk produced. The loops are bounded by +// the index array's own extent, never by a payload length. +func gatherFill[T any](dst, src []T, index *Array, srcShape []int, s, e, dim, ndim int, fail func(pos, idx int)) { + shape := index.shape + coord := make([]int, ndim) + if s > 0 { + rest := s + for d := ndim - 1; d >= 0; d-- { + coord[d] = rest % shape[d] + rest /= shape[d] + } + } + contig := index.isContiguous() + ip := index.ints + for i := s; i < e; i++ { + idx := int(ip[i]) + if !contig { + idx = int(ip[index.physIndex(i)]) + } + if idx < 0 || idx >= srcShape[dim] { + fail(i, idx) + advanceOdometer(coord, shape) + continue + } + // The source offset needs one fold per element: the output's own + // coordinates with the gathered axis replaced by the index. + off := 0 + for d := range ndim { + c := coord[d] + if d == dim { + c = idx + } + off = off*srcShape[d] + c + } + dst[i] = src[off] + advanceOdometer(coord, shape) + } +} + +// gatherCompatible reports whether src and index can combine on dim: +// every non-dim dimension must match in length. +func gatherCompatible(srcShape, indexShape []int, dim int) bool { + if len(srcShape) != len(indexShape) { + return false + } + for d := range srcShape { + if d == dim { + continue + } + if srcShape[d] != indexShape[d] { + return false + } + } + return true +} + +// Scatter is the inverse of Gather: it writes src values into self at +// positions indexed by index along dim. It returns a new array; the +// receiver is unchanged. Unindexed positions in the result keep self's +// values. +func Scatter(self *Array, dim int, index *Array, src *Array) (*Array, error) { + if index.dt != Int { + return nil, errf("Scatter: index must be an int array, got %s", index.dt) + } + if dim < 0 || dim >= self.NDim() { + return nil, errf("Scatter: dimension %d out of range for shape %s", dim, shapeText(self.shape)) + } + if !sameShape(index.shape, src.shape) { + return nil, errf("Scatter: index and src shapes must agree, got %s and %s", shapeText(index.shape), shapeText(src.shape)) + } + if !gatherCompatible(self.shape, index.shape, dim) { + return nil, errf("Scatter: shape mismatch %s vs %s on dim %d", shapeText(self.shape), shapeText(index.shape), dim) + } + // The source is read through its own accessors, so it may sit above + // self on the promotion ladder; complex into int or float32 is the + // one narrowing the library rejects (Astype's rule), and it must be + // a loud error rather than a read of a payload the source lacks. + if !canStore(self.dt, src.dt) { + return nil, errf("Scatter: cannot store a %s source in a %s array", src.dt, self.dt) + } + // The output starts as a full clone of self in self's own dtype, + // narrow payloads included; the setConverted walk below is + // dtype-generic. + out := self.cloneArray() + dst := make([]int, index.NDim()) + for i := range index.Len() { + idx := int(index.ints[index.physIndex(i)]) + if idx < 0 || idx >= self.shape[dim] { + return nil, errf("Scatter: index %d out of range for dimension %d of size %d at position %d", idx, dim, self.shape[dim], i) + } + // The destination offset needs one fold per element: the source's + // own coordinates with the scatter axis replaced by the index. + off := 0 + for d := range dst { + c := dst[d] + if d == dim { + c = idx + } + off = off*self.shape[d] + c + } + out.setConverted(off, src, i) + advanceOdometer(dst, index.shape) + } + return out, nil +} + +// Nonzero returns the multi-dimensional indices of every non-zero +// element, grouped per dimension. The result has a.NDim() slices; the +// i-th entry in each slice is the coordinate of the i-th nonzero +// element in that dimension. A complex array is rejected (no notion of +// zero in the lattice; equality with the zero value is already +// supported by Eq). +func Nonzero(a *Array) ([][]int, error) { + if a.dt == Complex { + return nil, errf("Nonzero: complex arrays have no notion of zero") + } + perDim := make([][]int, a.NDim()) + coord := make([]int, a.NDim()) + for i := range a.Len() { + if !isZero(a, i) { + for d := range a.NDim() { + perDim[d] = append(perDim[d], coord[d]) + } + } + advanceOdometer(coord, a.shape) + } + return perDim, nil +} + +func isZero(a *Array, i int) bool { + // A strided array maps the logical index to a payload slot through + // its strides, so only a contiguous operand may read the payload at + // the logical position directly; the accessor path keeps Argwhere + // and Nonzero correct the day a public strided view exists. + if a.strides != nil { + switch a.dt { + case Int, Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32: + // intAt widens the integer class exactly and reads bool + // against zero. + return a.intAt(i) == 0 + case Float16, Float32, Float: + return a.floatAt(i) == 0 + } + return false + } + switch a.dt { + case Int: + return a.ints[i] == 0 + case Bool: + return !a.bools[i] + case Int8: + return a.i8s[i] == 0 + case Uint8: + return a.u8s[i] == 0 + case Int16: + return a.i16s[i] == 0 + case Uint16: + return a.u16s[i] == 0 + case Int32: + return a.i32s[i] == 0 + case Uint32: + return a.u32s[i] == 0 + case Float16: + // The value test in bit space: clearing the sign bit folds -0.0 + // onto +0.0 exactly as the float comparisons below do. + return a.halves[i]&0x7FFF == 0 + case Float32: + return a.floats32[i] == 0 + case Float: + return a.floats[i] == 0 + } + return false +} + +// Take returns a 1-D copy whose elements are self[indices[i]] for each +// i, flat-indexed. Negative indices are an error. The walk is bounded by +// the index array's element count, not by the length of the payload it +// shares: a Slice view's payload can be longer. +func Take(a *Array, indices *Array) (*Array, error) { + if indices.dt != Int { + return nil, errf("Take: indices must be an int array, got %s", indices.dt) + } + if indices.NDim() != 1 { + return nil, errf("Take: indices must be 1-D, got shape %s", shapeText(indices.shape)) + } + n := indices.Len() + out := &Array{shape: []int{n}, dt: a.dt} + out.alloc(n) + if n == 0 { + return out, nil + } + src := a + if !src.isContiguous() { + src = src.materialise() + } + k := intPayload(indices) + // The workers write disjoint output slots and report the smallest + // offending position, the one the serial validation walk named. + bound := int64(a.Len()) + var mu sync.Mutex + firstPos := n + 1 + var firstVal int64 + fail := func(pos int, idx int64) { + mu.Lock() + defer mu.Unlock() + if pos < firstPos { + firstPos, firstVal = pos, idx + } + } + switch out.dt { + case Int: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.ints, src.ints, k, s, e, bound, fail) + }) + case Bool: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.bools, src.bools, k, s, e, bound, fail) + }) + case Int8: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.i8s, src.i8s, k, s, e, bound, fail) + }) + case Uint8: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.u8s, src.u8s, k, s, e, bound, fail) + }) + case Int16: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.i16s, src.i16s, k, s, e, bound, fail) + }) + case Uint16: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.u16s, src.u16s, k, s, e, bound, fail) + }) + case Int32: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.i32s, src.i32s, k, s, e, bound, fail) + }) + case Uint32: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.u32s, src.u32s, k, s, e, bound, fail) + }) + case Float16: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.halves, src.halves, k, s, e, bound, fail) + }) + case Float32: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.floats32, src.floats32, k, s, e, bound, fail) + }) + case Float: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.floats, src.floats, k, s, e, bound, fail) + }) + default: + parallelMin(n, copyMinPerWorker, func(s, e int) { + takeFill(out.complexes, src.complexes, k, s, e, bound, fail) + }) + } + if firstPos <= n { + return nil, errf("Take: index %d out of range for flat size %d at position %d", firstVal, a.Len(), firstPos) + } + return out, nil +} + +// takeFill writes dst[i] = src[k[i]] over the worker's range. An +// out-of-range index reports through fail and skips the write; a call +// that failed anywhere discards its output. +func takeFill[T any](dst, src []T, k []int64, s, e int, bound int64, fail func(pos int, idx int64)) { + for i := s; i < e; i++ { + v := k[i] + if v < 0 || v >= bound { + fail(i, v) + continue + } + dst[i] = src[v] + } +} diff --git a/internal/core/init.go b/internal/core/init.go new file mode 100644 index 0000000..77fe29f --- /dev/null +++ b/internal/core/init.go @@ -0,0 +1,121 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +// TruncatedNormal returns an array of the given shape with values +// drawn from a normal distribution truncated to ±2σ. Rejection sampling +// redraws until a value lands inside the window. All draws share the +// Generator for reproducibility; the dtype is float32. An invalid +// shape yields a nil array, like the other no-error constructors, and +// a std that is zero, negative or NaN is degenerate: the answer is the +// all-zero array, the escape from a retry loop that negative and NaN +// values could never leave. +func TruncatedNormal(g *Generator, shape []int, mean, std float64) *Array { + n, sh, err := checkedDims(shape) + if err != nil { + return nil + } + if !(std > 0) { + // std <= 0 or NaN: the window is a point, inverted or undefined, + // and the negative and NaN cases could never leave the retry loop + // below, so the degenerate fallback returns the all-zero array. + out := &Array{shape: sh, dt: Float32} + out.alloc(n) + return out + } + out := make([]float32, n) + twoStd := 2.0 * std + for i := range n { + // The scalar draw consumes the generator exactly as one + // Normal(g, 1, 0, std) call would (std is validated above, so + // that call cannot error), without its per-element allocation. + for { + v := std * g.normalUnit() + if v >= -twoStd && v <= twoStd { + out[i] = float32(v + mean) + break + } + } + } + return &Array{shape: sh, dt: Float32, floats32: out} +} + +// ZerosLike and OnesLike return a zero- or one-filled array with the +// same shape and dtype as a. The result is a copy: mutations on it +// do not affect a. +func ZerosLike(a *Array) *Array { + return fullLike(a, 0) +} + +func OnesLike(a *Array) *Array { + return fullLike(a, 1) +} + +// fullLike creates a new array with a's shape and dtype, filled with v. +// The narrow integer widths and bool take the same implicit-store cast +// every other fill path carries: Go's conversion through int64, and +// v != 0 for bool. +func fullLike(a *Array, v float64) *Array { + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(a.Len()) + switch a.dt { + case Int: + for i := range out.ints { + out.ints[i] = int64(v) + } + case Bool: + bv := v != 0 + for i := range out.bools { + out.bools[i] = bv + } + case Int8: + w := int8(int64(v)) + for i := range out.i8s { + out.i8s[i] = w + } + case Uint8: + w := uint8(int64(v)) + for i := range out.u8s { + out.u8s[i] = w + } + case Int16: + w := int16(int64(v)) + for i := range out.i16s { + out.i16s[i] = w + } + case Uint16: + w := uint16(int64(v)) + for i := range out.u16s { + out.u16s[i] = w + } + case Int32: + w := int32(int64(v)) + for i := range out.i32s { + out.i32s[i] = w + } + case Uint32: + w := uint32(int64(v)) + for i := range out.u32s { + out.u32s[i] = w + } + case Float16: + hv := HalfFromFloat64(v) + for i := range out.halves { + out.halves[i] = hv + } + case Float32: + for i := range out.floats32 { + out.floats32[i] = float32(v) + } + case Float: + for i := range out.floats { + out.floats[i] = v + } + default: + for i := range out.complexes { + out.complexes[i] = complex(v, 0) + } + } + return out +} diff --git a/internal/core/int_exactness_pins_test.go b/internal/core/int_exactness_pins_test.go new file mode 100644 index 0000000..c3eccd6 --- /dev/null +++ b/internal/core/int_exactness_pins_test.go @@ -0,0 +1,677 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "slices" + "testing" +) + +// Regression pins for int64 exactness and the dtype ladder. Wherever a +// view is possible the test slices first: a Slice view's payload runs +// past its element count, and that is where this class hides. + +// failOnPanic turns an unexpected panic into an ordinary test failure, +// so a panic-class regression is reported as one failing test instead +// of aborting the whole test binary. +func failOnPanic(t *testing.T) { + if r := recover(); r != nil { + t.Fatalf("unexpected panic: %v", r) + } +} + +// denseInts reads a dense int array the kernel just allocated. +func denseInts(a *Array) []int64 { return a.ints[:a.Len()] } + +// denseBools reads a dense bool array the kernel just allocated. +func denseBools(a *Array) []bool { return a.bools[:a.Len()] } + +// TestTopKIntExact pins: TopK on an int array +// ranked the values and wrote them back through float64, which rounds +// neighbours together above 2^53 and reports the wrong value and the +// wrong index. +func TestTopKIntExact(t *testing.T) { + big1 := int64(1)<<53 + 1 // 2^53 + 1 + big2 := int64(1)<<53 + 3 // 2^53 + 3 + a := mustFromInts(t, []int64{big1, big2, 5}, 3) + vals, idxs, err := TopK(a, 2, 0) + if err != nil { + t.Fatalf("TopK: %v", err) + } + if got, want := denseInts(vals), []int64{big2, big1}; !slices.Equal(got, want) { + t.Errorf("TopK values = %v, want %v", got, want) + } + if got, want := denseInts(idxs), []int64{1, 0}; !slices.Equal(got, want) { + t.Errorf("TopK indices = %v, want %v", got, want) + } + + // A single huge value comes back unchanged, not rounded. + single := int64(1)<<62 - 12345 + sv, si, err := TopK(mustFromInts(t, []int64{single}, 1), 1, 0) + if err != nil { + t.Fatalf("TopK single: %v", err) + } + if got := sv.ints[0]; got != single { + t.Errorf("TopK single value = %d, want %d", got, single) + } + if got := si.ints[0]; got != 0 { + t.Errorf("TopK single index = %d, want 0", got) + } + + // A rebased view: payload 4, Len 3. + parent := mustFromInts(t, []int64{0, big1, big2, 7, 9}, 5) + view := mustSlice(t, parent, 0, 1, 4) + vv, vi, err := TopK(view, 2, 0) + if err != nil { + t.Fatalf("TopK view: %v", err) + } + if got, want := denseInts(vv), []int64{big2, big1}; !slices.Equal(got, want) { + t.Errorf("TopK view values = %v, want %v", got, want) + } + if got, want := denseInts(vi), []int64{1, 0}; !slices.Equal(got, want) { + t.Errorf("TopK view indices = %v, want %v", got, want) + } + + // 2-D int along dim 0: the layout stays (k, suffix) row-major and + // the values stay exact. + m := mustFromInts(t, []int64{big1, 2, big2, 1}, 2, 2) + mv, mi, err := TopK(m, 1, 0) + if err != nil { + t.Fatalf("TopK 2-D: %v", err) + } + if got, want := denseInts(mv), []int64{big2, 2}; !slices.Equal(got, want) { + t.Errorf("TopK 2-D values = %v, want %v", got, want) + } + if got, want := denseInts(mi), []int64{1, 0}; !slices.Equal(got, want) { + t.Errorf("TopK 2-D indices = %v, want %v", got, want) + } +} + +// TestTopKEmptyTrailingDim pins: shape (3, 0) +// with dim 0 leaves the per-line stride zero, so the line count divides +// 0 by 0. +func TestTopKEmptyTrailingDim(t *testing.T) { + defer failOnPanic(t) + a := mustFloats(t, nil, 3, 0) + if _, _, err := TopK(a, 2, 0); err == nil { + t.Error("TopK((3, 0), 2, 0): expected an error, got none") + } + b := mustFromInts(t, nil, 2, 0) + if _, _, err := TopK(b, 1, 0); err == nil { + t.Error("TopK int (2, 0), 1, 0: expected an error, got none") + } + // An empty chosen dimension keeps its own diagnostic. + c := mustFloats(t, nil, 0) + if _, _, err := TopK(c, 0, 0); err == nil { + t.Error("TopK((0,), 0, 0): expected an error, got none") + } +} + +// TestSearchSortedIntExact pins: SearchSorted +// widened both operands through float64, so 2^60+1 and 2^60+2 both +// rounded onto 2^60 and the needle landed past its own neighbour. +func TestSearchSortedIntExact(t *testing.T) { + base := int64(1) << 60 + hay := mustFromInts(t, []int64{base, base + 2}, 2) + got, err := SearchSorted(hay, mustFromInts(t, []int64{base + 1}, 1)) + if err != nil { + t.Fatalf("SearchSorted: %v", err) + } + if got.ints[0] != 1 { + t.Errorf("SearchSorted(2^60+1) = %d, want 1", got.ints[0]) + } + + // The rightmost rule still counts an equal element and the last one. + exact, err := SearchSorted(hay, mustFromInts(t, []int64{base}, 1)) + if err != nil { + t.Fatalf("SearchSorted exact: %v", err) + } + if exact.ints[0] != 1 { + t.Errorf("SearchSorted(2^60) = %d, want 1", exact.ints[0]) + } + top, err := SearchSorted(hay, mustFromInts(t, []int64{base + 2}, 1)) + if err != nil { + t.Fatalf("SearchSorted top: %v", err) + } + if top.ints[0] != 2 { + t.Errorf("SearchSorted(2^60+2) = %d, want 2", top.ints[0]) + } + + // A rebased needle view: payload 5, Len 2. + parent := mustFromInts(t, []int64{7, base - 1, base + 1, base + 5, 8}, 5) + view := mustSlice(t, parent, 0, 1, 3) + gotView, err := SearchSorted(hay, view) + if err != nil { + t.Fatalf("SearchSorted view: %v", err) + } + if want := []int64{0, 1}; !slices.Equal(denseInts(gotView), want) { + t.Errorf("SearchSorted view = %v, want %v", denseInts(gotView), want) + } + + // The mixed and float paths keep the documented float64 widening. + mixed, err := SearchSorted(mustFloats(t, []float64{1, 2, 3}, 3), mustFromInts(t, []int64{2}, 1)) + if err != nil { + t.Fatalf("SearchSorted mixed: %v", err) + } + if mixed.ints[0] != 2 { + t.Errorf("SearchSorted mixed = %d, want 2", mixed.ints[0]) + } +} + +// TestAssignBinsIntExact pins: AssignBins chose +// the bin through the same float64 widening, so a value just above a +// bin edge rounded onto the edge and fell into the wrong bin. +func TestAssignBinsIntExact(t *testing.T) { + base := int64(1) << 60 + edges := mustFromInts(t, []int64{0, base + 100, 2 * base}, 3) + got, err := AssignBins(mustFromInts(t, []int64{base + 64}, 1), edges) + if err != nil { + t.Fatalf("AssignBins: %v", err) + } + if got.ints[0] != 0 { + t.Errorf("AssignBins(2^60+64) = %d, want 0", got.ints[0]) + } + + // Edge cases: exactly on the second edge, above the last edge, below + // the first. + atEdge, err := AssignBins(mustFromInts(t, []int64{base + 100}, 1), edges) + if err != nil { + t.Fatalf("AssignBins edge: %v", err) + } + if atEdge.ints[0] != 1 { + t.Errorf("AssignBins(2^60+100) = %d, want 1", atEdge.ints[0]) + } + over, err := AssignBins(mustFromInts(t, []int64{2*base + 5}, 1), edges) + if err != nil { + t.Fatalf("AssignBins over: %v", err) + } + if over.ints[0] != 1 { + t.Errorf("AssignBins(2^61+5) = %d, want 1 (the outer bin)", over.ints[0]) + } + under, err := AssignBins(mustFromInts(t, []int64{-1}, 1), edges) + if err != nil { + t.Fatalf("AssignBins under: %v", err) + } + if under.ints[0] != 0 { + t.Errorf("AssignBins(-1) = %d, want 0 (the outer bin)", under.ints[0]) + } + + // A rebased value view: payload 5, Len 2. + parent := mustFromInts(t, []int64{9, base + 64, base + 100, 2*base + 7, 9}, 5) + view := mustSlice(t, parent, 0, 1, 3) + gotView, err := AssignBins(view, edges) + if err != nil { + t.Fatalf("AssignBins view: %v", err) + } + if want := []int64{0, 1}; !slices.Equal(denseInts(gotView), want) { + t.Errorf("AssignBins view = %v, want %v", denseInts(gotView), want) + } + + // The float path is unchanged. + fl, err := AssignBins(mustFloats(t, []float64{2.5}, 1), mustFloats(t, []float64{0, 1, 2}, 3)) + if err != nil { + t.Fatalf("AssignBins float: %v", err) + } + if fl.ints[0] != 1 { + t.Errorf("AssignBins(2.5) = %d, want 1", fl.ints[0]) + } +} + +// TestMaskMixedIntFloatExact pins: mixing an int +// and a float operand compared the int side through float64, so values +// above 2^53 collapsed onto their rounded neighbours and every +// comparison reported the wrong answer. +func TestMaskMixedIntFloatExact(t *testing.T) { + base := int64(1) << 60 + rounded := mustFromFloats(t, []float64{float64(base)}, 1) + big := mustFromInts(t, []int64{base + 1}, 1) + bigger := mustFromInts(t, []int64{base + 64}, 1) + + eq, err := Eq(big, rounded) + if err != nil { + t.Fatalf("Eq: %v", err) + } + if eq.bools[0] { + t.Errorf("Eq(2^60+1, float 2^60) answered true, want false") + } + ne, err := Ne(big, rounded) + if err != nil { + t.Fatalf("Ne: %v", err) + } + if !ne.bools[0] { + t.Errorf("Ne(2^60+1, float 2^60) answered false, want true") + } + le, err := Le(bigger, rounded) + if err != nil { + t.Fatalf("Le: %v", err) + } + if le.bools[0] { + t.Errorf("Le(2^60+64, float 2^60) answered true, want false") + } + gt, err := Gt(bigger, rounded) + if err != nil { + t.Fatalf("Gt: %v", err) + } + if !gt.bools[0] { + t.Errorf("Gt(2^60+64, float 2^60) answered false, want true") + } + lt, err := Lt(rounded, bigger) + if err != nil { + t.Fatalf("Lt: %v", err) + } + if !lt.bools[0] { + t.Errorf("Lt(float 2^60, 2^60+64) answered false, want true") + } + ge, err := Ge(mustFromInts(t, []int64{base}, 1), rounded) + if err != nil { + t.Fatalf("Ge: %v", err) + } + if !ge.bools[0] { + t.Errorf("Ge(2^60, float 2^60) answered false, want true") + } + + // The scalar twins: an int scalar against a float array and a float + // scalar against an int array. + eqi, err := EqI(rounded, base+64) + if err != nil { + t.Fatalf("EqI: %v", err) + } + if eqi.bools[0] { + t.Errorf("EqI(float 2^60, 2^60+64) answered true, want false") + } + eqf, err := EqF(big, float64(base)) + if err != nil { + t.Fatalf("EqF: %v", err) + } + if eqf.bools[0] { + t.Errorf("EqF(2^60+1, float 2^60) answered true, want false") + } + lei, err := LeI(rounded, base-64) + if err != nil { + t.Fatalf("LeI: %v", err) + } + if lei.bools[0] { + t.Errorf("LeI(float 2^60, 2^60-64) answered true, want false") + } + gei, err := GeI(rounded, base+64) + if err != nil { + t.Fatalf("GeI: %v", err) + } + if gei.bools[0] { + t.Errorf("GeI(float 2^60, 2^60+64) answered true, want false") + } + + // A complex array against an int scalar compares the real part + // exactly too. + eqc, err := EqI(mustFromComplexes(t, []complex128{complex(float64(base), 0)}, 1), base+64) + if err != nil { + t.Fatalf("EqI complex: %v", err) + } + if eqc.bools[0] { + t.Errorf("EqI(complex 2^60, 2^60+64) answered true, want false") + } + + // Fractional floats behave as IEEE plus exact integer arithmetic. + frac := mustFromFloats(t, []float64{3.5}, 1) + if m, _ := Eq(mustFromInts(t, []int64{3}, 1), frac); m.bools[0] { + t.Errorf("Eq(3, 3.5) answered true, want false") + } + if m, _ := Le(mustFromInts(t, []int64{3}, 1), frac); !m.bools[0] { + t.Errorf("Le(3, 3.5) answered false, want true") + } + if m, _ := Gt(mustFromInts(t, []int64{4}, 1), frac); !m.bools[0] { + t.Errorf("Gt(4, 3.5) answered false, want true") + } + if m, _ := Ge(mustFromInts(t, []int64{4}, 1), frac); !m.bools[0] { + t.Errorf("Ge(4, 3.5) answered false, want true") + } + neg := mustFromFloats(t, []float64{-3.5}, 1) + if m, _ := Gt(mustFromInts(t, []int64{-3}, 1), neg); !m.bools[0] { + t.Errorf("Gt(-3, -3.5) answered false, want true") + } + if m, _ := Lt(mustFromInts(t, []int64{-4}, 1), neg); !m.bools[0] { + t.Errorf("Lt(-4, -3.5) answered false, want true") + } + if m, _ := Le(mustFromInts(t, []int64{-3}, 1), neg); m.bools[0] { + t.Errorf("Le(-3, -3.5) answered true, want false") + } + + // NaN and the infinities keep their IEEE answers. + nan := mustFromFloats(t, []float64{math.NaN()}, 1) + if m, _ := Eq(mustFromInts(t, []int64{5}, 1), nan); m.bools[0] { + t.Errorf("Eq(5, NaN) answered true, want false") + } + if m, _ := Ne(mustFromInts(t, []int64{5}, 1), nan); !m.bools[0] { + t.Errorf("Ne(5, NaN) answered false, want true") + } + if m, _ := Gt(mustFromInts(t, []int64{math.MaxInt64}, 1), mustFromFloats(t, []float64{math.Inf(1)}, 1)); m.bools[0] { + t.Errorf("Gt(MaxInt64, +Inf) answered true, want false") + } + if m, _ := Lt(mustFromInts(t, []int64{math.MinInt64}, 1), mustFromFloats(t, []float64{math.Inf(-1)}, 1)); m.bools[0] { + t.Errorf("Lt(MinInt64, -Inf) answered true, want false") + } + if m, _ := Gt(mustFromInts(t, []int64{math.MinInt64}, 1), mustFromFloats(t, []float64{math.Inf(-1)}, 1)); !m.bools[0] { + t.Errorf("Gt(MinInt64, -Inf) answered false, want true") + } + + // Both boundaries of the int64 range compare against the floats that + // sit exactly on them. + if m, _ := Eq(mustFromInts(t, []int64{math.MinInt64}, 1), mustFromFloats(t, []float64{math.Ldexp(-1, 63)}, 1)); !m.bools[0] { + t.Errorf("Eq(MinInt64, -2^63) answered false, want true") + } + if m, _ := Lt(mustFromInts(t, []int64{math.MaxInt64}, 1), mustFromFloats(t, []float64{math.Ldexp(1, 63)}, 1)); !m.bools[0] { + t.Errorf("Lt(MaxInt64, 2^63) answered false, want true") + } + + // A rebased int view against a float array: payload 5, Len 2. + parent := mustFromInts(t, []int64{9, base, base + 64, base + 100, 9}, 5) + view := mustSlice(t, parent, 0, 1, 3) + pair := mustFromFloats(t, []float64{float64(base), float64(base)}, 2) + gotView, err := Le(view, pair) + if err != nil { + t.Fatalf("Le view: %v", err) + } + if want := []bool{true, false}; !slices.Equal(denseBools(gotView), want) { + t.Errorf("Le(view, floats) = %v, want %v", denseBools(gotView), want) + } +} + +// TestTakeViewIndices pins: Take walked the +// index array's payload rather than its element count, so a view index +// wrote past the output it had allocated. +func TestTakeViewIndices(t *testing.T) { + defer failOnPanic(t) + src := mustFromFloats(t, []float64{10, 20, 30, 40, 50, 60, 70, 80, 90, 100}, 10) + parent := mustFromInts(t, []int64{0, 1, 2, 3, 4, 5, 6, 7, 8, 9}, 10) + + // The report's shape: Len 3 over a payload of 9. + view := mustSlice(t, parent, 0, 1, 4) + out, err := Take(src, view) + if err != nil { + t.Fatalf("Take: %v", err) + } + if out.Len() != 3 { + t.Fatalf("Take length = %d, want 3", out.Len()) + } + if got, want := denseFloats(out), []float64{20, 30, 40}; !slices.Equal(got, want) { + t.Errorf("Take = %v, want %v", got, want) + } + + // A view rebased past zero: the indices are 3, 4 and 5, not the + // payload's first three slots. + off := mustSlice(t, parent, 0, 3, 6) + outOff, err := Take(src, off) + if err != nil { + t.Fatalf("Take offset view: %v", err) + } + if got, want := denseFloats(outOff), []float64{40, 50, 60}; !slices.Equal(got, want) { + t.Errorf("Take offset view = %v, want %v", got, want) + } + + // An index beyond the source is still the documented error. + short := mustFromFloats(t, []float64{1, 2, 3}, 3) + if _, err := Take(short, view); err == nil { + t.Error("Take with an out-of-range index: expected an error, got none") + } +} + +// denseFloats reads a dense float array the kernel just allocated. +func denseFloats(a *Array) []float64 { return a.floats[:a.Len()] } + +// TestScatterMixedDtypes pins: Scatter read the +// source payload by the destination dtype, so any source above the +// destination on the ladder read a nil payload and panicked. +func TestScatterMixedDtypes(t *testing.T) { + defer failOnPanic(t) + self := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + idx := mustFromInts(t, []int64{0, 1}, 2, 1) + zvals := mustFromComplexes(t, []complex128{1 + 2i, 3 + 4i}, 2, 1) + + // Int destination, float source: Go's cast, as Astype documents. + out, err := Scatter(self, 1, idx, mustFromFloats(t, []float64{9.5, 8.5}, 2, 1)) + if err != nil { + t.Fatalf("Scatter int dest from a float source: %v", err) + } + if got, want := denseInts(out), []int64{9, 2, 3, 8}; !slices.Equal(got, want) { + t.Errorf("Scatter int dest = %v, want %v", got, want) + } + if got, want := denseInts(self), []int64{1, 2, 3, 4}; !slices.Equal(got, want) { + t.Errorf("Scatter wrote through its receiver: %v", got) + } + + // Float32 destination, int source. + self32 := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 2, 2) + out32, err := Scatter(self32, 1, idx, mustFromInts(t, []int64{7, 9}, 2, 1)) + if err != nil { + t.Fatalf("Scatter float32 dest from an int source: %v", err) + } + if got, want := out32.floats32[:4], []float32{7, 2, 3, 9}; !slices.Equal(got, want) { + t.Errorf("Scatter float32 dest = %v, want %v", got, want) + } + + // Float destination, complex source: the real part, the same rule + // Astype applies to complex to float. + selfF := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + outF, err := Scatter(selfF, 1, idx, zvals) + if err != nil { + t.Fatalf("Scatter float dest from a complex source: %v", err) + } + if got, want := denseFloats(outF), []float64{1, 2, 3, 3}; !slices.Equal(got, want) { + t.Errorf("Scatter float dest = %v, want %v", got, want) + } + + // Complex destination from a real source keeps working. + selfC := mustFromComplexes(t, []complex128{0, 0, 0, 0}, 2, 2) + outC, err := Scatter(selfC, 1, idx, mustFromFloats(t, []float64{5, 6}, 2, 1)) + if err != nil { + t.Fatalf("Scatter complex dest from a float source: %v", err) + } + if got, want := outC.complexes[:4], []complex128{5, 0, 0, 6}; !slices.Equal(got, want) { + t.Errorf("Scatter complex dest = %v, want %v", got, want) + } + + // Narrowing a complex source into int or float32 is the one + // direction the library rejects (Astype's rule): a loud error, not a + // panic on the nil payload. + if _, err := Scatter(self, 1, idx, zvals); err == nil { + t.Error("Scatter int dest from a complex source: expected an error, got none") + } + if _, err := Scatter(self32, 1, idx, zvals); err == nil { + t.Error("Scatter float32 dest from a complex source: expected an error, got none") + } + + // A rebased source view: payload 4, Len 2. + parentF := mustFromFloats(t, []float64{0, 9.5, 8.5, 99}, 4, 1) + viewSrc := mustSlice(t, parentF, 0, 1, 3) + outV, err := Scatter(self, 1, idx, viewSrc) + if err != nil { + t.Fatalf("Scatter view source: %v", err) + } + if got, want := denseInts(outV), []int64{9, 2, 3, 8}; !slices.Equal(got, want) { + t.Errorf("Scatter view source = %v, want %v", got, want) + } + + // A rebased index view decides the same positions as a dense one. + parentI := mustFromInts(t, []int64{9, 0, 1, 9}, 4, 1) + viewIdx := mustSlice(t, parentI, 0, 1, 3) + outVI, err := Scatter(self, 1, viewIdx, mustFromFloats(t, []float64{9.5, 8.5}, 2, 1)) + if err != nil { + t.Fatalf("Scatter view index: %v", err) + } + if got, want := denseInts(outVI), []int64{9, 2, 3, 8}; !slices.Equal(got, want) { + t.Errorf("Scatter view index = %v, want %v", got, want) + } +} + +// TestComplexValuesViewPrefix pins: the complex +// path handed out the whole payload of a rebased view, so every +// consumer saw the invisible tail elements. +func TestComplexValuesViewPrefix(t *testing.T) { + c := mustFromComplexes(t, []complex128{0, 1, 2, 3, 4}, 5) + view := mustSlice(t, c, 0, 1, 4) // Len 3, payload 4 + vals, err := view.ComplexValues("probe") + if err != nil { + t.Fatalf("ComplexValues: %v", err) + } + if len(vals) != 3 { + t.Fatalf("ComplexValues returned %d values, want 3", len(vals)) + } + if want := []complex128{1, 2, 3}; !slices.Equal(vals, want) { + t.Errorf("ComplexValues = %v, want %v", vals, want) + } + + // A 2-D view: the flat payload is 12 long, the view holds 4. + m := mustFromComplexes(t, []complex128{0, 0, 0, 0, 1, 2, 3, 4, 0, 0, 0, 0}, 3, 4) + row := mustSlice(t, m, 0, 1, 2) // Len 4, payload 8 + rowVals, err := row.ComplexValues("probe") + if err != nil { + t.Fatalf("ComplexValues row: %v", err) + } + if want := []complex128{1, 2, 3, 4}; !slices.Equal(rowVals, want) { + t.Errorf("ComplexValues row = %v, want %v", rowVals, want) + } + + // A strided view is not a contiguous prefix: the elements come from + // the accessor. + strided := &Array{shape: []int{2}, dt: Complex, complexes: []complex128{1, 2, 3, 4}, strides: []int{2}} + sv, err := strided.ComplexValues("probe") + if err != nil { + t.Fatalf("ComplexValues strided: %v", err) + } + if want := []complex128{1, 3}; !slices.Equal(sv, want) { + t.Errorf("ComplexValues strided = %v, want %v", sv, want) + } +} + +// TestDiffIntExact pins: Diff forced its output +// to float and subtracted through float64, so int64 neighbours above +// 2^53 lost their difference. +func TestDiffIntExact(t *testing.T) { + defer failOnPanic(t) + base := int64(1) << 60 + d, err := Diff(mustFromInts(t, []int64{base, base + 300}, 2), 1, 0) + if err != nil { + t.Fatalf("Diff: %v", err) + } + if d.Dtype() != Int { + t.Fatalf("Diff of an int array is %s, want int", d.Dtype()) + } + if d.Len() != 1 { + t.Fatalf("Diff length = %d, want 1", d.Len()) + } + if got := d.ints[0]; got != 300 { + t.Errorf("Diff = %d, want 300", got) + } + + // Order 2 stays exact as well. + d2, err := Diff(mustFromInts(t, []int64{0, base, 2*base + 7}, 3), 2, 0) + if err != nil { + t.Fatalf("Diff order 2: %v", err) + } + if got := d2.ints[0]; got != 7 { + t.Errorf("Diff order 2 = %d, want 7", got) + } + + // A rebased view input: payload 5, Len 3. + parent := mustFromInts(t, []int64{9, base, base + 300, 2 * base, 9}, 5) + view := mustSlice(t, parent, 0, 1, 4) + dv, err := Diff(view, 1, 0) + if err != nil { + t.Fatalf("Diff view: %v", err) + } + if got, want := denseInts(dv), []int64{300, base - 300}; !slices.Equal(got, want) { + t.Errorf("Diff view = %v, want %v", got, want) + } + + // The float, float32 and complex routes keep their dtypes. + df, err := Diff(mustFromFloats(t, []float64{1, 4, 9}, 3), 1, 0) + if err != nil { + t.Fatalf("Diff float: %v", err) + } + if df.Dtype() != Float { + t.Errorf("Diff of a float array is %s, want float", df.Dtype()) + } + if got, want := denseFloats(df), []float64{3, 5}; !slices.Equal(got, want) { + t.Errorf("Diff float = %v, want %v", got, want) + } + df32, err := Diff(mustFromFloat32s(t, []float32{1, 4, 9}, 3), 1, 0) + if err != nil { + t.Fatalf("Diff float32: %v", err) + } + if df32.Dtype() != Float { + t.Errorf("Diff of a float32 array is %s, want float (unchanged)", df32.Dtype()) + } + dc, err := Diff(mustFromComplexes(t, []complex128{1 + 1i, 3 + 4i}, 2), 1, 0) + if err != nil { + t.Fatalf("Diff complex: %v", err) + } + if dc.Dtype() != Complex { + t.Errorf("Diff of a complex array is %s, want complex", dc.Dtype()) + } + if got := dc.complexes[0]; got != 2+3i { + t.Errorf("Diff complex = %v, want 2+3i", got) + } +} + +// TestAstypeIntIdentity pins: Astype routed +// every conversion through float64, so converting an int array to int +// rounded its own values. +func TestAstypeIntIdentity(t *testing.T) { + base := int64(1) << 60 + big := int64(1)<<53 + 1 + a := mustFromInts(t, []int64{big, 1, 2}, 3) + out, err := Astype(a, Int) + if err != nil { + t.Fatalf("Astype int to int: %v", err) + } + if got, want := denseInts(out), []int64{big, 1, 2}; !slices.Equal(got, want) { + t.Errorf("Astype int to int = %v, want %v", got, want) + } + // It is a copy: writing the result leaves the source alone. + out.ints[0] = 42 + if a.ints[0] != big { + t.Errorf("Astype aliased its input: source is now %d", a.ints[0]) + } + + // A rebased view input: payload 4, Len 2. + parent := mustFromInts(t, []int64{9, big, big - 2, 9}, 4) + view := mustSlice(t, parent, 0, 1, 3) + vout, err := Astype(view, Int) + if err != nil { + t.Fatalf("Astype view: %v", err) + } + if got, want := denseInts(vout), []int64{big, big - 2}; !slices.Equal(got, want) { + t.Errorf("Astype view = %v, want %v", got, want) + } + + // The other identity conversions are copies too. + f := mustFromFloats(t, []float64{1.5, 2.5}, 2) + fout, err := Astype(f, Float) + if err != nil { + t.Fatalf("Astype float to float: %v", err) + } + fout.floats[0] = 9 + if f.floats[0] != 1.5 { + t.Errorf("Astype float to float aliased its input") + } + c := mustFromComplexes(t, []complex128{1 + 2i}, 1) + cout, err := Astype(c, Complex) + if err != nil { + t.Fatalf("Astype complex to complex: %v", err) + } + cout.complexes[0] = 0 + if c.complexes[0] != 1+2i { + t.Errorf("Astype complex to complex aliased its input") + } + + // The documented narrowing survives: float to int is Go's cast. + narrow, err := Astype(mustFromFloats(t, []float64{9.7, -9.7, float64(base)}, 3), Int) + if err != nil { + t.Fatalf("Astype float to int: %v", err) + } + if got, want := denseInts(narrow), []int64{9, -9, base}; !slices.Equal(got, want) { + t.Errorf("Astype float to int = %v, want %v", got, want) + } +} diff --git a/internal/core/interpolate2d.go b/internal/core/interpolate2d.go new file mode 100644 index 0000000..bc68b91 --- /dev/null +++ b/internal/core/interpolate2d.go @@ -0,0 +1,199 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "sync" +) + +// Bilinear interpolation over a regular grid, the two-dimensional +// companion of the 1-D Interpolate in arrayutil.go. Where the 1-D +// entry point clamps queries to the boundary, the grid version +// refuses them: in two dimensions a silently clamped value is far +// harder to spot, and the honest answer to an out-of-domain query is +// an error. + +// interp2dMinPerWorker is the per-worker chunk floor for the bilinear +// queries. A query costs a couple of Floor calls and a dozen flops +// over four grid reads, roughly a dozen nanoseconds, so a chunk in +// the hundreds already outweighs a worker's start-up; the sweep around +// the value pinned 256 ahead of both 1024 and 4096 at the +// million-query size. +const interp2dMinPerWorker = 256 + +// Interpolate2D evaluates the bilinear interpolation of a regular +// grid at a set of query points. grid is a rank-2, rows × cols array +// whose element (i, j) samples the interpolated field at +// (x0 + j·dx, y0 + i·dy); xs and ys hold the query coordinates, must +// share their length, and the result takes xs's shape. Bilinear +// interpolation is exact on every function that is bilinear within a +// cell, so linear fields come back unchanged and grid nodes are +// reproduced exactly. +// +// Every query must lie inside the closed rectangle the grid spans, +// boundaries included; a query outside it is an error, never a +// silent clamp. Either sign of dx and dy works, so descending axes +// need no preprocessing. +// +// Errors: grid not rank 2 or an extent below 2, complex inputs, a +// grid dtype the surface refuses (bool and the narrow integer +// widths, refused by name; convert with Astype first), xs and ys of +// different lengths, a zero or non-finite spacing or origin +// coordinate, a non-finite query coordinate, or a query outside the +// domain. +func Interpolate2D(grid, xs, ys *Array, x0, y0, dx, dy float64) (*Array, error) { + if grid.NDim() != 2 { + return nil, errf("Interpolate2D: grid must be rank 2, got shape %s", shapeText(grid.Shape())) + } + rows, cols := grid.Shape()[0], grid.Shape()[1] + if rows < 2 || cols < 2 { + return nil, errf("Interpolate2D: grid needs at least 2×2 samples, got %d×%d", rows, cols) + } + if grid.dt == Complex { + return nil, errf("Interpolate2D: complex grids are not supported") + } + if xs.Len() != ys.Len() { + return nil, errf("Interpolate2D: xs and ys must share their length, got %d and %d", xs.Len(), ys.Len()) + } + if xs.dt == Complex || ys.dt == Complex { + return nil, errf("Interpolate2D: complex query coordinates are not supported") + } + if narrowRefused(grid.dt) { + // Bool and the narrow integer grids carry no bilinear kernel, + // and the default grid reads below would reach a nil float64 + // payload for them: the narrow-dtype refusal names the grid dtype + // and the conversion instead, after the rank and complex gates + // whose texts the existing pins carry. + return nil, errf("Interpolate2D: dtype %s is not supported; convert with Astype", grid.dt) + } + if dx == 0 || dy == 0 || math.IsNaN(dx) || math.IsNaN(dy) || + math.IsInf(dx, 0) || math.IsInf(dy, 0) || + math.IsNaN(x0) || math.IsNaN(y0) || math.IsInf(x0, 0) || math.IsInf(y0, 0) { + return nil, errf("Interpolate2D: the origin and spacings must be finite with dx, dy non-zero, got x0=%g y0=%g dx=%g dy=%g", x0, y0, dx, dy) + } + // The domain runs between the first and last sample along each + // axis, which for a negative spacing means a reversed interval. + xLo, xHi := min(x0, x0+float64(cols-1)*dx), max(x0, x0+float64(cols-1)*dx) + yLo, yHi := min(y0, y0+float64(rows-1)*dy), max(y0, y0+float64(rows-1)*dy) + // The coordinates and the grid are read through dense payload + // windows: materialise once, widen the coordinates exactly as + // floatAt does (float32 and int convert exactly), and bind the + // grid slices so the loop pays no per-element accessor dispatch. + xf, yf := widened(xs), widened(ys) + g := grid.materialise() + n := xs.Len() + out := &Array{shape: append([]int{}, xs.shape...), dt: Float} + out.alloc(n) + // The workers write disjoint output slots and report a query + // error through a mutex touched only on the failing path; the + // sequential contract is that the first offending query wins, so + // the reported error is the one with the smallest query index. + var mu sync.Mutex + var firstErr error + firstIdx := 0 + fail := func(i int, err error) { + mu.Lock() + defer mu.Unlock() + if firstErr == nil || i < firstIdx { + firstErr, firstIdx = err, i + } + } + parallelMin(n, interp2dMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + x, y := xf[i], yf[i] + if math.IsNaN(x) || math.IsInf(x, 0) || math.IsNaN(y) || math.IsInf(y, 0) { + fail(i, errf("Interpolate2D: query %d is not finite (x=%g, y=%g)", i, x, y)) + return + } + if x < xLo || x > xHi || y < yLo || y > yHi { + fail(i, errf("Interpolate2D: query (%g, %g) lies outside the grid domain x ∈ [%g, %g], y ∈ [%g, %g]", x, y, xLo, xHi, yLo, yHi)) + return + } + // Cell indices from the fractional grid position, folded onto + // the final cell at the exact upper boundary; the lower clamps + // only absorb rounding at the exact lower boundary. + gx := (x - x0) / dx + gy := (y - y0) / dy + j := min(max(int(math.Floor(gx)), 0), cols-2) + r := min(max(int(math.Floor(gy)), 0), rows-2) + tx := gx - float64(j) + ty := gy - float64(r) + var v00, v01, v10, v11 float64 + switch g.dt { + case Float16: + // The widening of a half grid value is exact, so the + // fold runs on the same bits the float32 path runs on. + b := r*cols + j + v00, v01 = HalfToFloat64(g.halves[b]), HalfToFloat64(g.halves[b+1]) + v10, v11 = HalfToFloat64(g.halves[b+cols]), HalfToFloat64(g.halves[b+cols+1]) + case Float32: + b := r*cols + j + v00, v01 = float64(g.floats32[b]), float64(g.floats32[b+1]) + v10, v11 = float64(g.floats32[b+cols]), float64(g.floats32[b+cols+1]) + case Int: + b := r*cols + j + v00, v01 = float64(g.ints[b]), float64(g.ints[b+1]) + v10, v11 = float64(g.ints[b+cols]), float64(g.ints[b+cols+1]) + default: + b := r*cols + j + v00, v01 = g.floats[b], g.floats[b+1] + v10, v11 = g.floats[b+cols], g.floats[b+cols+1] + } + // The explicit conversions are the FMA fence, the one the + // window catalogue uses: a contracting build would fuse the + // bare products into the sums and the answer would drift a + // ulp off the portable bits. + top := float64((1-tx)*v00) + float64(tx*v01) + bot := float64((1-tx)*v10) + float64(tx*v11) + out.floats[i] = float64((1-ty)*top) + float64(ty*bot) + } + }) + if firstErr != nil { + return nil, firstErr + } + return out, nil +} + +// widened returns the array's elements as a float64 payload, widened +// exactly as floatAt widens: float16, float32, int, bool and the narrow +// integer widths all convert exactly, and float aliases the payload +// itself. The loops are bounded by the logical length: a rebased view +// carries a payload longer than its extent. +func widened(a *Array) []float64 { + m := a.materialise() + switch m.dt { + case Float16: + out := make([]float64, m.Len()) + for i := range m.Len() { + out[i] = HalfToFloat64(m.halves[i]) + } + return out + case Float32: + out := make([]float64, m.Len()) + for i := range m.Len() { + out[i] = float64(m.floats32[i]) + } + return out + case Int: + out := make([]float64, m.Len()) + for i := range m.Len() { + out[i] = float64(m.ints[i]) + } + return out + case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32: + // The narrow real class widens through the same accessor the + // query gates let through: bool reads 0/1 and every narrow + // integer widens exactly, the values floatAt hands over. Without + // this arm a bool query array reached the default and read a nil + // float64 payload, panicking on the first slot. + out := make([]float64, m.Len()) + for i := range m.Len() { + out[i] = m.floatAt(i) + } + return out + default: + return m.floats + } +} diff --git a/internal/core/interpolate2d_test.go b/internal/core/interpolate2d_test.go new file mode 100644 index 0000000..20f5a7d --- /dev/null +++ b/internal/core/interpolate2d_test.go @@ -0,0 +1,233 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +// TestInterpolate2DLinearField checks that a linear field comes back +// unchanged: bilinear interpolation is exact on it, so the queries +// pin hand-computed values of 2 + 3x − 1.5y on a 5×4 grid that does +// not start at the origin and has unequal spacings. +func TestInterpolate2DLinearField(t *testing.T) { + const x0, dx, y0, dy = 0.0, 0.5, 10.0, 2.0 + gridVals := make([]float64, 0, 20) + for i := range 5 { + for j := range 4 { + x := x0 + float64(j)*dx + y := y0 + float64(i)*dy + gridVals = append(gridVals, 2+3*x-1.5*y) + } + } + grid := mustFloats(t, gridVals, 5, 4) + qx := []float64{1.3, 0, 1.5, 0.75, 0.25} + qy := []float64{12.7, 10, 18, 16, 11} + want := []float64{-13.15, -13, -20.5, -19.75, 2 + 0.75 - 16.5} + got, err := Interpolate2D(grid, mustFloats(t, qx), mustFloats(t, qy), x0, y0, dx, dy) + if err != nil { + t.Fatalf("Interpolate2D: %v", err) + } + for i := range want { + if v := got.FloatAt(i); math.Abs(v-want[i]) > 1e-12*math.Abs(want[i]) { + t.Errorf("query (%g, %g) = %v, want %v", qx[i], qy[i], v, want[i]) + } + } +} + +// TestInterpolate2DPatchMidpoints pins exact values inside a single +// bilinear patch: the cell corners (1, 2; 3, 7) average to 13/4 at +// the centre, and the off-centre point lands on 53/16. Both targets +// are dyadic, so the comparison is exact. +func TestInterpolate2DPatchMidpoints(t *testing.T) { + grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2) + got, err := Interpolate2D(grid, mustFloats(t, []float64{0.5, 0.25}), mustFloats(t, []float64{0.5, 0.75}), 0, 0, 1, 1) + if err != nil { + t.Fatalf("Interpolate2D: %v", err) + } + if got.FloatAt(0) != 3.25 { + t.Errorf("centre = %v, want 3.25", got.FloatAt(0)) + } + if got.FloatAt(1) != 3.3125 { + t.Errorf("off-centre = %v, want 3.3125", got.FloatAt(1)) + } +} + +// TestInterpolate2DNodes checks that every grid node is reproduced +// exactly, and that a grid with a descending x axis (dx < 0) gives +// the mirrored result. +func TestInterpolate2DNodes(t *testing.T) { + grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2) + for i := range 2 { + for j := range 2 { + got, err := Interpolate2D(grid, mustFloats(t, []float64{float64(j)}), mustFloats(t, []float64{float64(i)}), 0, 0, 1, 1) + if err != nil { + t.Fatalf("Interpolate2D(%d, %d): %v", i, j, err) + } + if want := grid.FloatAt(2*i + j); got.FloatAt(0) != want { + t.Errorf("node (%d, %d) = %v, want %v", i, j, got.FloatAt(0), want) + } + } + } + got, err := Interpolate2D(grid, mustFloats(t, []float64{0.5}), mustFloats(t, []float64{0.5}), 1, 0, -1, 1) + if err != nil { + t.Fatalf("Interpolate2D descending: %v", err) + } + if got.FloatAt(0) != 3.25 { + t.Errorf("descending axis: %v, want 3.25", got.FloatAt(0)) + } +} + +// TestInterpolate2DOutsideRejects pins the no-silent-clamping +// contract: queries beyond any edge of the domain, and non-finite +// coordinates, are errors. +func TestInterpolate2DOutsideRejects(t *testing.T) { + grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2) + cases := []struct { + name string + xs, ys []float64 + errSubstr string + }{ + {"below x", []float64{-0.1}, []float64{0.5}, "outside the grid domain"}, + {"above x", []float64{1.0001}, []float64{0.5}, "outside the grid domain"}, + {"below y", []float64{0.5}, []float64{-0.5}, "outside the grid domain"}, + {"above y", []float64{0.5}, []float64{1.5}, "outside the grid domain"}, + {"not a number", []float64{0.5}, []float64{math.NaN()}, "not finite"}, + } + for _, c := range cases { + _, err := Interpolate2D(grid, mustFloats(t, c.xs), mustFloats(t, c.ys), 0, 0, 1, 1) + if err == nil { + t.Errorf("%s: expected an error", c.name) + } else if !strings.Contains(err.Error(), c.errSubstr) { + t.Errorf("%s: error %q lacks %q", c.name, err, c.errSubstr) + } + } +} + +// TestInterpolate2DValidation pins the input contracts: rank-2 grids +// of extent ≥ 2, matched coordinate lengths, finite non-zero +// spacings, and real-valued arrays. +func TestInterpolate2DValidation(t *testing.T) { + grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2) + xs := mustFloats(t, []float64{0.5}) + ys := mustFloats(t, []float64{0.5}) + if _, err := Interpolate2D(mustFloats(t, []float64{1, 2, 3, 4}, 4), xs, ys, 0, 0, 1, 1); err == nil { + t.Error("expected an error for a rank-1 grid") + } + if _, err := Interpolate2D(mustFloats(t, []float64{1, 2, 3, 4}, 1, 4), xs, ys, 0, 0, 1, 1); err == nil { + t.Error("expected an error for a 1×4 grid") + } + if _, err := Interpolate2D(grid, mustFloats(t, []float64{0.5, 0.6}), ys, 0, 0, 1, 1); err == nil { + t.Error("expected an error for mismatched coordinate lengths") + } + if _, err := Interpolate2D(grid, xs, ys, 0, 0, 0, 1); err == nil { + t.Error("expected an error for dx = 0") + } + if _, err := Interpolate2D(grid, xs, ys, math.Inf(1), 0, 1, 1); err == nil { + t.Error("expected an error for a non-finite origin") + } + complexGrid, _ := FromComplexes([]complex128{1, 2, 3, 4}, 2, 2) + if _, err := Interpolate2D(complexGrid, xs, ys, 0, 0, 1, 1); err == nil { + t.Error("expected an error for a complex grid") + } + if _, err := Interpolate2D(grid, mustFromComplexes(t, []complex128{0.5, 1}, 2), ys, 0, 0, 1, 1); err == nil { + t.Error("expected an error for complex query coordinates") + } +} + +// TestInterpolate2DShape checks that the result takes the query +// array's shape, not just its length. +func TestInterpolate2DShape(t *testing.T) { + grid := mustFloats(t, []float64{1, 2, 3, 7}, 2, 2) + xs := mustFloats(t, []float64{0.5, 0.5, 0.5, 0.5}, 2, 2) + ys := mustFloats(t, []float64{0.5, 0.5, 0.5, 0.5}) + got, err := Interpolate2D(grid, xs, ys, 0, 0, 1, 1) + if err != nil { + t.Fatalf("Interpolate2D: %v", err) + } + if shape := got.Shape(); len(shape) != 2 || shape[0] != 2 || shape[1] != 2 { + t.Errorf("result shape %v, want (2, 2)", shape) + } +} + +// TestInterpolate2DNarrowGridRefused pins the narrow-dtype refusal at +// the entry: bool and the narrow integer grids carry no bilinear +// kernel, and the default grid reads would reach a nil float64 payload +// for them. The refusal names the grid dtype and the conversion, after +// the rank and complex gates whose texts the pins above carry. +func TestInterpolate2DNarrowGridRefused(t *testing.T) { + must := func(a *Array, err error) *Array { + if err != nil { + t.Fatal(err) + } + return a + } + xs := mustFloats(t, []float64{0.5}) + ys := mustFloats(t, []float64{0.5}) + cases := []struct { + name string + grid *Array + }{ + {"bool", must(FromBools([]bool{true, false, true, true}, 2, 2))}, + {"int8", must(FromInt8s([]int8{1, 2, 3, 4}, 2, 2))}, + {"uint8", must(FromUint8s([]uint8{1, 2, 3, 4}, 2, 2))}, + {"int16", must(FromInt16s([]int16{1, 2, 3, 4}, 2, 2))}, + {"uint16", must(FromUint16s([]uint16{1, 2, 3, 4}, 2, 2))}, + {"int32", must(FromInt32s([]int32{1, 2, 3, 4}, 2, 2))}, + {"uint32", must(FromUint32s([]uint32{1, 2, 3, 4}, 2, 2))}, + } + for _, c := range cases { + _, err := Interpolate2D(c.grid, xs, ys, 0, 0, 1, 1) + if err == nil || !strings.Contains(err.Error(), "Interpolate2D") || + !strings.Contains(err.Error(), c.name) || + !strings.Contains(err.Error(), "convert with Astype") { + t.Errorf("Interpolate2D %s grid: err = %v, want the narrow-dtype refusal naming %s and the conversion", + c.name, err, c.name) + } + } +} + +// TestInterpolate2DNarrowQueries pins the narrow query dtypes: a bool or +// narrow integer coordinate array widens exactly the way floatAt widens +// it, answering the values the equivalent float queries answer. The +// widened helper used to read the nil float64 payload for these dtypes, +// so the first query slot panicked instead of interpolating. +func TestInterpolate2DNarrowQueries(t *testing.T) { + g, err := FromFloats([]float64{0, 1, 2, 3}, 2, 2) + if err != nil { + t.Fatal(err) + } + // The field samples 2y + x at (x0+j·dx, y0+i·dy), so the query + // (x, y) answers 2y + x whatever dtype carries the coordinates. + mk := func(a *Array, err error) *Array { + if err != nil { + t.Fatal(err) + } + return a + } + cases := []struct { + name string + xs *Array + }{ + {"bool", mk(FromBools([]bool{false, true}, 2))}, + {"int8", mk(FromInt8s([]int8{0, 1}, 2))}, + {"uint16", mk(FromUint16s([]uint16{0, 1}, 2))}, + {"int32", mk(FromInt32s([]int32{0, 1}, 2))}, + } + ys := mustFloats(t, []float64{0, 0.5}) + for _, c := range cases { + got, err := Interpolate2D(g, c.xs, ys, 0, 0, 1, 1) + if err != nil { + t.Fatalf("Interpolate2D %s queries: %v", c.name, err) + } + want := []float64{0, 2} + for i, w := range want { + if v := got.FloatAt(i); v != w { + t.Errorf("Interpolate2D %s queries[%d] = %v, want %v", c.name, i, v, w) + } + } + } +} diff --git a/internal/core/interpolate_knot_test.go b/internal/core/interpolate_knot_test.go new file mode 100644 index 0000000..dfcc1e3 --- /dev/null +++ b/internal/core/interpolate_knot_test.go @@ -0,0 +1,38 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// TestInterpolateQueryOnKnotTakesTheLeftSegment pins the segment choice +// for a query landing exactly on a knot: the bisection takes the first +// knot at or above the query, so the left segment is used and the query +// sits at its upper end with t exactly 1. With a lower y of zero the +// interpolation is exactly the knot's own y. A repeated knot makes the +// choice observable: the query returns the first of the repeated ys, +// where the right segment would return the second. +func TestInterpolateQueryOnKnotTakesTheLeftSegment(t *testing.T) { + xs := mustFloats(t, []float64{0, 1, 2, 3}) + ys := mustFloats(t, []float64{0, 10, 20, 30}) + q := mustFloats(t, []float64{1, 2}) + out, err := Interpolate(xs, ys, q) + if err != nil { + t.Fatalf("Interpolate: %v", err) + } + for i, w := range []float64{10, 20} { + if got := out.FloatAt(i); got != w { + t.Fatalf("query on knot %v answered %v, want exactly %v", q.FloatAt(i), got, w) + } + } + dup, err := Interpolate( + mustFloats(t, []float64{0, 1, 1, 2}), + mustFloats(t, []float64{10, 20, 30, 40}), + mustFloats(t, []float64{1})) + if err != nil { + t.Fatalf("Interpolate: %v", err) + } + if got := dup.FloatAt(0); got != 20 { + t.Fatalf("query on a repeated knot answered %v, want the left segment's 20", got) + } +} diff --git a/internal/core/jacobian.go b/internal/core/jacobian.go new file mode 100644 index 0000000..f998756 --- /dev/null +++ b/internal/core/jacobian.go @@ -0,0 +1,100 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// Numerical Jacobians for vector-valued functions, the standalone +// companion of the central differences LevenbergMarquardt builds +// internally: one column per input coordinate, two evaluations per +// column, and a per-column step that scales with the coordinate's +// magnitude so every column carries the same relative resolution. + +// JacobianOptions tunes Jacobian. Step is the absolute difference +// step applied to every coordinate; zero or negative selects the +// default per-column step sqrt(ε)·max(1, |x_j|), the largest step +// whose central-difference truncation error still sits below the +// rounding floor. +type JacobianOptions struct { + Step float64 +} + +// Jacobian returns the Jacobian of f at x as an (m × n) float array, +// entry (i, j) holding ∂f_i/∂x_j by central differences with the +// step opts.Step, or sqrt(ε)·max(1, |x_j|) per column when unset. x +// may hold any real dtype and shape; its n elements are perturbed one +// at a time. f must map the point and every probe to a real rank-1 +// array of one fixed length m: a non-vector output, an output that +// changes length between columns, an empty output, or an error from +// f is reported. The probe arrays handed to f are reused between +// columns, so f must not retain them. +func Jacobian(f func(x *Array) (*Array, error), x *Array, opts JacobianOptions) (*Array, error) { + if x.dt == Complex { + return nil, errf("Jacobian: complex points are not supported") + } + n := x.Len() + if n == 0 { + return nil, errf("Jacobian: the point must not be empty") + } + // The unperturbed evaluation fixes m up front and holds f to it + // for every column that follows. + base0, err := f(x) + if err != nil { + return nil, base.WrapErr("Jacobian", err) + } + if base0.dt == Complex { + return nil, errf("Jacobian: complex outputs are not supported") + } + if base0.NDim() != 1 { + return nil, errf("Jacobian: f must return a vector, got shape %s", shapeText(base0.Shape())) + } + m := base0.Len() + if m == 0 { + return nil, errf("Jacobian: f returns an empty vector") + } + out := &Array{shape: []int{m, n}, dt: Float} + out.alloc(m * n) + // The point widened once, then two probe copies reused across the + // coordinate sweep: f receives an array it may keep for the + // duration of the call, but each column rewrites both from + // scratch. + p := make([]float64, n) + for k := range n { + p[k] = x.floatAt(k) + } + pp := make([]float64, n) + pm := make([]float64, n) + for j := range n { + h := opts.Step + if h <= 0 { + h = math.Sqrt(base.EpsF) * math.Max(1, math.Abs(p[j])) + } + copy(pp, p) + copy(pm, p) + pp[j] += h + pm[j] -= h + fp, errP := f(&Array{shape: append([]int{}, x.shape...), dt: Float, floats: pp}) + fm, errM := f(&Array{shape: append([]int{}, x.shape...), dt: Float, floats: pm}) + if errP != nil { + return nil, base.WrapErr("Jacobian", errP) + } + if errM != nil { + return nil, base.WrapErr("Jacobian", errM) + } + if fp.NDim() != 1 || fp.Len() != m || fm.NDim() != 1 || fm.Len() != m { + return nil, errf("Jacobian: f must return %d values at column %d, got %d and %d", m, j, fp.Len(), fm.Len()) + } + if fp.dt == Complex || fm.dt == Complex { + return nil, errf("Jacobian: complex outputs are not supported") + } + for i := range m { + out.floats[i*n+j] = (fp.floatAt(i) - fm.floatAt(i)) / (2 * h) + } + } + return out, nil +} diff --git a/internal/core/jacobian_test.go b/internal/core/jacobian_test.go new file mode 100644 index 0000000..df088d1 --- /dev/null +++ b/internal/core/jacobian_test.go @@ -0,0 +1,156 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "errors" + "math" + "strings" + "testing" +) + +// TestJacobianAnalytic pins the Jacobian of f(x, y) = (x·y, x+y), +// whose exact derivative [[y, x], [1, 1]] the central differences +// must reproduce to 1e-8, at points that exercise the step scaling. +func TestJacobianAnalytic(t *testing.T) { + f := func(x *Array) (*Array, error) { + xv, yv := x.FloatAt(0), x.FloatAt(1) + return FromFloats([]float64{xv * yv, xv + yv}, 2) + } + cases := []struct { + name string + x, y float64 + want []float64 + }{ + {"first quadrant", 2, 3, []float64{3, 2, 1, 1}}, + {"negative coordinate", -1.5, 4, []float64{4, -1.5, 1, 1}}, + {"origin", 0, 0, []float64{0, 0, 1, 1}}, + } + for _, c := range cases { + j, err := Jacobian(f, mustFloats(t, []float64{c.x, c.y}), JacobianOptions{}) + if err != nil { + t.Fatalf("%s: %v", c.name, err) + } + if shape := j.Shape(); len(shape) != 2 || shape[0] != 2 || shape[1] != 2 { + t.Fatalf("%s: shape %v, want (2, 2)", c.name, shape) + } + for i, want := range c.want { + if got := j.FloatAt(i); math.Abs(got-want) > 1e-8 { + t.Errorf("%s: entry %d = %v, want %v", c.name, i, got, want) + } + } + } +} + +// TestJacobianNonlinear pins a nonlinear scalar-input case, +// f(x) = (x², x³) at x = 1.7, whose Jacobian is the column +// (2x, 3x²) = (3.4, 8.67). +func TestJacobianNonlinear(t *testing.T) { + f := func(x *Array) (*Array, error) { + xv := x.FloatAt(0) + return FromFloats([]float64{xv * xv, xv * xv * xv}, 2) + } + j, err := Jacobian(f, mustFloats(t, []float64{1.7}), JacobianOptions{}) + if err != nil { + t.Fatalf("Jacobian: %v", err) + } + // The x³ column sits at the default step's rounding floor + // (~2·10⁻⁸ of cancellation noise), so the pin is 2e-8 there. + tols := []float64{1e-8, 2e-8} + for i, want := range []float64{3.4, 8.67} { + if got := j.FloatAt(i); math.Abs(got-want) > tols[i] { + t.Errorf("entry %d = %v, want %v", i, got, want) + } + } +} + +// TestJacobianCustomStep checks that a caller-supplied step is +// honoured: on a linear field any step is exact, so the result pins +// the matrix rather than the step heuristic. +func TestJacobianCustomStep(t *testing.T) { + f := func(x *Array) (*Array, error) { + return FromFloats([]float64{2*x.FloatAt(0) - x.FloatAt(1)}, 1) + } + j, err := Jacobian(f, mustFloats(t, []float64{5, -7}), JacobianOptions{Step: 1e-6}) + if err != nil { + t.Fatalf("Jacobian: %v", err) + } + if shape := j.Shape(); len(shape) != 2 || shape[0] != 1 || shape[1] != 2 { + t.Fatalf("shape %v, want (1, 2)", shape) + } + for i, want := range []float64{2, -1} { + // The wide 1e-6 step leaves ~2e-9 of cancellation noise on + // values of size ~17; the pin still says which matrix it is. + if got := j.FloatAt(i); math.Abs(got-want) > 1e-8 { + t.Errorf("entry %d = %v, want %v", i, got, want) + } + } +} + +// TestJacobianRejects pins the contracts: complex points, empty +// points, non-vector outputs, outputs whose length changes between +// columns, and errors from f all come back as errors. +func TestJacobianRejects(t *testing.T) { + complexPoint, _ := FromComplexes([]complex128{1, 2}, 2) + if _, err := Jacobian(vectorIdentity, complexPoint, JacobianOptions{}); err == nil { + t.Error("expected an error for a complex point") + } else if !strings.Contains(err.Error(), "complex") { + t.Errorf("complex point: error %q lacks \"complex\"", err) + } + if _, err := Jacobian(vectorIdentity, mustFloats(t, nil), JacobianOptions{}); err == nil { + t.Error("expected an error for an empty point") + } + + matrixOut := func(x *Array) (*Array, error) { + return FromFloats([]float64{1, 0, 0, 1}, 2, 2) + } + if _, err := Jacobian(matrixOut, mustFloats(t, []float64{1}), JacobianOptions{}); err == nil { + t.Error("expected an error for a matrix-shaped output") + } else if !strings.Contains(err.Error(), "vector") { + t.Errorf("matrix output: error %q lacks \"vector\"", err) + } + + emptyOut := func(x *Array) (*Array, error) { + return FromFloats([]float64{}, 0) + } + if _, err := Jacobian(emptyOut, mustFloats(t, []float64{1}), JacobianOptions{}); err == nil { + t.Error("expected an error for an empty output") + } + + calls := 0 + changingLength := func(x *Array) (*Array, error) { + calls++ + if calls == 1 { + return FromFloats([]float64{1, 2}, 2) + } + return FromFloats([]float64{1, 2, 3}, 3) + } + if _, err := Jacobian(changingLength, mustFloats(t, []float64{1, 2}), JacobianOptions{}); err == nil { + t.Error("expected an error for an output length that changes between columns") + } + + sentinel := errf("f blew up") + failing := func(x *Array) (*Array, error) { + if x.FloatAt(0) != 1 { + return nil, sentinel + } + return FromFloats([]float64{1}, 1) + } + _, err := Jacobian(failing, mustFloats(t, []float64{1}), JacobianOptions{}) + if err == nil { + t.Error("expected the probe failure to propagate") + } else if !strings.Contains(err.Error(), "blew up") { + t.Errorf("probe failure: error %q lacks the cause", err) + } else if !errors.Is(err, sentinel) { + // The wrap keeps the chain open, so the sentinel stays reachable + // through the entry point's context: a revert to a plain %v + // formatting fails exactly here. + t.Errorf("probe failure: error %q does not unwrap to the cause", err) + } +} + +// vectorIdentity is a well-behaved f for the input-side rejections. +func vectorIdentity(x *Array) (*Array, error) { + return FromFloats([]float64{x.FloatAt(0), x.FloatAt(1)}, 2) +} diff --git a/internal/core/knot_ellipsis_pins_test.go b/internal/core/knot_ellipsis_pins_test.go new file mode 100644 index 0000000..33fec09 --- /dev/null +++ b/internal/core/knot_ellipsis_pins_test.go @@ -0,0 +1,292 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// Regression pins: interpolation knot and query validation, the einsum +// ellipsis axis order, the sparse COO length contract, the norm of an +// empty dimension, the TopK sort path, the Kron widen-once rule and +// the truncated-normal stream. + +// TestInterpolateRepeatedEdgeKnot pins the zero-width segment +// handling at the range edges: a query equal to a duplicated edge knot +// returns the first of the repeated ys, and queries past it take the +// upper end, where the old 0/0 division returned NaN. +func TestInterpolateRepeatedEdgeKnot(t *testing.T) { + ys := mustFloats(t, []float64{10, 20, 30}) + cases := []struct { + xs, q []float64 + want float64 + name string + }{ + {[]float64{0, 1, 1}, []float64{1}, 20, "right duplicate, query on it"}, + {[]float64{0, 1, 1}, []float64{1.5}, 30, "right duplicate, query above"}, + {[]float64{0, 0, 1}, []float64{0}, 10, "left duplicate, query on it"}, + {[]float64{0, 0, 1}, []float64{-1}, 10, "left duplicate, query below"}, + } + for _, tc := range cases { + got, err := Interpolate(mustFloats(t, tc.xs), ys, mustFloats(t, tc.q)) + if err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + if math.IsNaN(got.FloatAt(0)) || got.FloatAt(0) != tc.want { + t.Fatalf("%s: got %v, want %v", tc.name, got.FloatAt(0), tc.want) + } + } +} + +// TestInterpolateRejectsBadKnotsAndQueries pins that a +// non-finite knot or a NaN query is a loud error rather than a NaN +// flowing through the evaluation. +func TestInterpolateRejectsBadKnotsAndQueries(t *testing.T) { + ys := mustFloats(t, []float64{1, 2}) + if _, err := Interpolate(mustFloats(t, []float64{0, math.NaN()}), ys, mustFloats(t, []float64{0.5})); err == nil { + t.Fatal("expected an error for a NaN knot") + } + if _, err := Interpolate(mustFloats(t, []float64{0, math.Inf(1)}), ys, mustFloats(t, []float64{0.5})); err == nil { + t.Fatal("expected an error for an Inf knot") + } + if _, err := Interpolate(mustFloats(t, []float64{0, 1}), ys, mustFloats(t, []float64{math.NaN()})); err == nil { + t.Fatal("expected an error for a NaN query") + } +} + +// TestInterpolateMonotoneRejectsNaNKnots pins that the +// strictly-increasing gate cannot be bypassed by NaN, for which every +// comparison is false. +func TestInterpolateMonotoneRejectsNaNKnots(t *testing.T) { + xs := mustFloats(t, []float64{0, math.NaN()}) + ys := mustFloats(t, []float64{1, 2}) + if _, err := InterpolateMonotone(xs, ys, mustFloats(t, []float64{0.5})); err == nil { + t.Fatal("expected an error for a NaN knot") + } + if _, err := InterpolateMonotone(mustFloats(t, []float64{0, math.Inf(1)}), ys, mustFloats(t, []float64{0.5})); err == nil { + t.Fatal("expected an error for an Inf knot") + } +} + +// TestEinsumEllipsisAxisOrder pins that with two or more +// ellipsis axes per operand the output keeps the operand's own +// left-to-right axis order. The synthetic labels count from the right +// edge, and sorting them ascending transposed the batch axes. +func TestEinsumEllipsisAxisOrder(t *testing.T) { + a := mustFloats(t, []float64{ + 0.5, 1, 2, 3, + 4, 5, 6, 7, + 8, 9, 10, 11, + 12, 13, 14, 15, + 16, 17, 18, 19, + }, 5, 2, 2) + b := mustFloats(t, []float64{ + 2, 0, 0, 2, + 2, 0, 0, 2, + 2, 0, 0, 2, + 2, 0, 0, 2, + 20, 0, 0, 2, + }, 5, 2, 2) + want := make([]float64, 5*2*2) + for i := range 5 { + for j := range 2 { + for l := range 2 { + sum := 0.0 + for k := range 2 { + sum += a.FloatAt((i*2+j)*2+k) * b.FloatAt((i*2+k)*2+l) + } + want[(i*2+j)*2+l] = sum + } + } + } + for _, spec := range []string{"...jk,...kl->...jl", "...jk,...kl"} { + got, err := Einsum(spec, a, b) + if err != nil { + t.Fatalf("%s: %v", spec, err) + } + if got.Shape()[0] != 5 || got.Shape()[1] != 2 || got.Shape()[2] != 2 { + t.Fatalf("%s: shape %v, want [5 2 2]", spec, got.Shape()) + } + for i := range got.Len() { + if got.FloatAt(i) != want[i] { + t.Fatalf("%s: [%d] = %v, want %v", spec, i, got.FloatAt(i), want[i]) + } + } + } +} + +// TestSparseCOOLengthMismatch pins that a hand-built SparseCOO +// literal whose Values run shorter than its index rows is refused by +// every entry point that walks the coordinates, rather than panicking. +func TestSparseCOOLengthMismatch(t *testing.T) { + idx := mustFromInts(t, []int64{0, 0, 0, 1, 1, 0, 1, 1}, 4, 2) + vals := mustFloats(t, []float64{1, 2}) + s := &SparseCOO{Indices: idx, Values: vals, Shape: []int{2, 2}} + if _, err := s.Dense(); err == nil { + t.Fatal("Dense: expected a length error") + } + dense := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2) + if _, err := SpMul(s, dense); err == nil { + t.Fatal("SpMul: expected a length error") + } + if _, err := SpMatMul(s, dense); err == nil { + t.Fatal("SpMatMul: expected a length error") + } +} + +// TestNormEmptyDim pins the infinity norm of an empty +// reduced dimension: NaN like MinAxis and MaxAxis, while the p = 1 +// sum of nothing stays zero. +func TestNormEmptyDim(t *testing.T) { + a := mustFloats(t, []float64{}, 2, 0) + inf, err := Norm(a, math.Inf(1), 1, false) + if err != nil { + t.Fatalf("Norm p=Inf: %v", err) + } + if !math.IsNaN(inf.FloatAt(0)) || !math.IsNaN(inf.FloatAt(1)) { + t.Fatalf("Norm p=Inf over an empty line = %v, want NaN", inf.FloatAt(0)) + } + one, err := Norm(a, 1, 1, false) + if err != nil { + t.Fatalf("Norm p=1: %v", err) + } + if one.FloatAt(0) != 0 { + t.Fatalf("Norm p=1 over an empty line = %v, want 0", one.FloatAt(0)) + } +} + +// TestTopKSortPathMatchesElections runs the sort-based wide-k +// path against a reference repeated argmax on lines with duplicates +// and NaN: the two paths must agree element for element, values and +// indices, across dtypes. +func TestTopKSortPathMatchesElections(t *testing.T) { + refTopK := func(vals []float64, isInt []int64, dt Dtype, k int) ([]float64, []int64) { + n := len(vals) + used := make([]bool, n) + var outV []float64 + var outI []int64 + for range k { + best := -1 + for i := range n { + if used[i] { + continue + } + if dt == Int { + if best < 0 || isInt[i] > isInt[best] { + best = i + } + continue + } + if v := vals[i]; v == v && (best < 0 || v > vals[best]) { + best = i + } + } + if best < 0 { + if dt == Int { + outV = append(outV, 0) + } else { + outV = append(outV, math.NaN()) + } + outI = append(outI, 0) + continue + } + if dt == Int { + outV = append(outV, float64(isInt[best])) + } else { + outV = append(outV, vals[best]) + } + outI = append(outI, int64(best)) + used[best] = true + } + return outV, outI + } + vals := []float64{3, -1, 3, math.NaN(), 0.5, 3, -1, 2} + ints := make([]int64, len(vals)) + for i, v := range vals { + ints[i] = int64(v * 4) + } + for _, dt := range []Dtype{Float, Float32, Int} { + var a *Array + switch dt { + case Int: + a = mustFromInts(t, ints, len(ints)) + case Float32: + a = &Array{shape: []int{len(vals)}, dt: Float32, floats32: make([]float32, len(vals))} + for i, v := range vals { + a.floats32[i] = float32(v) + } + default: + a = mustFloats(t, vals, len(vals)) + } + for _, k := range []int{len(vals), len(vals) - 2, 5, 2} { + gotV, gotI, err := TopK(a, k, 0) + if err != nil { + t.Fatalf("%s k=%d: %v", dt, k, err) + } + refV, refI := refTopK(vals, ints, dt, k) + for i := range k { + gv := gotV.FloatAt(i) + if math.IsNaN(refV[i]) { + if !math.IsNaN(gv) { + t.Fatalf("%s k=%d [%d]: got %v, want NaN", dt, k, i, gv) + } + } else if dt == Int { + if gotV.ints[i] != int64(refV[i]) { + t.Fatalf("%s k=%d [%d]: got %v, want %v", dt, k, i, gotV.ints[i], refV[i]) + } + } else if gv != refV[i] { + t.Fatalf("%s k=%d [%d]: got %v, want %v", dt, k, i, gv, refV[i]) + } + if gotI.ints[i] != refI[i] { + t.Fatalf("%s k=%d [%d]: index %d, want %d", dt, k, i, gotI.ints[i], refI[i]) + } + } + } + } +} + +// TestKronNarrowsOnce pins the widen-once rule on Kron: an int operand +// widens to float64, the product narrows once, so an int above 2^24 +// keeps the single rounding of the sibling kernels. +func TestKronNarrowsOnce(t *testing.T) { + a := mustFromInts(t, []int64{1<<24 + 1}, 1, 1) + b := &Array{shape: []int{1, 1}, dt: Float32, floats32: []float32{3}} + got, err := Kron(a, b) + if err != nil { + t.Fatalf("Kron: %v", err) + } + av := float64(1<<24 + 1) + want := float32(av * 3) + if got.floats32[0] != want { + t.Fatalf("Kron int×float32 = %v, want %v (one narrowing)", got.floats32[0], want) + } +} + +// TestTruncatedNormalStream pins that the allocation-free draw +// consumes the generator exactly as the per-element Normal call did. +func TestTruncatedNormalStream(t *testing.T) { + g1 := NewGenerator(42) + g2 := NewGenerator(42) + direct := TruncatedNormal(g1, []int{512}, 0.5, 2) + manual := make([]float32, 512) + twoStd := 4.0 + for i := range 512 { + for { + arr, err := Normal(g2, 1, 0, 2) + if err != nil { + t.Fatalf("Normal: %v", err) + } + if v := arr.floats[0]; v >= -twoStd && v <= twoStd { + manual[i] = float32(v + 0.5) + break + } + } + } + for i := range 512 { + if direct.floats32[i] != manual[i] { + t.Fatalf("[%d] = %v, want %v", i, direct.floats32[i], manual[i]) + } + } +} diff --git a/internal/core/mask.go b/internal/core/mask.go new file mode 100644 index 0000000..0a04e46 --- /dev/null +++ b/internal/core/mask.go @@ -0,0 +1,1283 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "sync" +) + +// Masks. Element-wise comparisons answer a bool mask: the +// payload is the predicate itself, one byte per element against the +// eight the 0/1 int mask of earlier releases wrote. Masks compose +// through the bool logic operations And, Or, Xor and Not at the foot of +// this file, and bool carries no arithmetic. Select picks the elements +// where a mask is set and Where chooses element-wise between two +// arrays; both take a bool mask, and both still take the int mask of +// 0/1 the comparisons used to answer, a nonzero element standing for +// true. NaN follows IEEE: every ordering and equality comparison is +// false, except Ne, which is true. Complex arrays take part in Eq and +// Ne only: ordering has no complex meaning. + +// cmpOp selects the relation a comparison reports. +type cmpOp uint8 + +const ( + opEq cmpOp = iota + opNe + opLt + opLe + opGt + opGe +) + +// isOrdering reports whether op needs an order, which complex operands +// do not have. +func (o cmpOp) isOrdering() bool { return o == opLt || o == opLe || o == opGt || o == opGe } + +// mirror returns op with its operands swapped: x < y is y > x. +func (o cmpOp) mirror() cmpOp { + switch o { + case opLt: + return opGt + case opGt: + return opLt + case opLe: + return opGe + case opGe: + return opLe + default: + return o + } +} + +// applyInt applies op to two int64 values. +func applyInt(op cmpOp, x, y int64) bool { + switch op { + case opEq: + return x == y + case opNe: + return x != y + case opLt: + return x < y + case opLe: + return x <= y + case opGt: + return x > y + default: + return x >= y + } +} + +// applyFloat applies op to two float64 values; NaN follows IEEE, so +// every ordering and equality comparison is false and Ne is true. +func applyFloat(op cmpOp, x, y float64) bool { + switch op { + case opEq: + return x == y + case opNe: + return x != y + case opLt: + return x < y + case opLe: + return x <= y + case opGt: + return x > y + default: + return x >= y + } +} + +// applyComplex applies op to two complex values; complex takes part in +// Eq and Ne only, so an ordering op never reaches here. +func applyComplex(op cmpOp, x, y complex128) bool { + if op == opNe { + return x != y + } + return x == y +} + +// twoPow63 is 2^63 as a float64: the first value above every int64. +const twoPow63 = 9223372036854775808.0 + +// applyIntFloat reports op(x, y) for an int64 x and a float64 y, exactly. +// Widening x to float64 rounds two neighbouring ints above 2^53 onto the +// same value, which is how Eq once answered 1 for 2^60+1 against the +// float 2^60. An integral y in int64 range compares natively instead; a +// fractional y, which no int can equal, is decided by its floor and +// ceiling, because an integer is below a fractional y exactly when it is +// at most floor(y) and above it at least ceil(y). +func applyIntFloat(op cmpOp, x int64, y float64) bool { + switch { + case math.IsNaN(y): + return op == opNe + case math.IsInf(y, 1): + return op == opNe || op == opLt || op == opLe + case math.IsInf(y, -1): + return op == opNe || op == opGt || op == opGe + case y >= twoPow63: + // Above every int64. + return op == opNe || op == opLt || op == opLe + case y < -twoPow63: + // Below every int64. + return op == opNe || op == opGt || op == opGe + } + trunc := int64(y) + if float64(trunc) == y { + return applyInt(op, x, trunc) + } + floor, ceil := trunc, trunc+1 + if y < 0 { + floor, ceil = trunc-1, trunc + } + switch op { + case opLt, opLe: + return x <= floor + case opGt, opGe: + return x >= ceil + case opNe: + return true + default: // opEq + return false + } +} + +// cmpArray runs an element-wise comparison, producing a bool mask. An +// ordering op rejects complex operands; mixed int and float operands +// compare exactly through applyIntFloat, while the int, float and +// complex branches widen exactly. Dense same-width payloads take the +// raw walk above, a mixed narrow-width pair takes its int64 widening +// walk; a strided view and a pair that puts a bool beside a narrow +// width keep the accessor walk, which reads the same values in the +// same order. +func (a *Array) cmpArray(b *Array, name string, op cmpOp) (*Array, error) { + if !sameShape(a.shape, b.shape) { + return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape)) + } + if op.isOrdering() && (a.dt == Complex || b.dt == Complex) { + return nil, errf("%s: complex arrays have no ordering", name) + } + n := a.Len() + out := &Array{shape: a.Shape(), dt: Bool, bools: make([]bool, n)} + switch { + case a.dt == Int && b.dt == Int: + if a.isContiguous() && b.isContiguous() { + intCmpRun(op, a.ints[:n], b.ints[:n], out.bools) + return out, nil + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyInt(op, a.ints[a.physIndex(i)], b.ints[b.physIndex(i)]) + } + }) + case promote(a.dt, b.dt) == Complex: + if a.dt == Complex && b.dt == Complex && a.isContiguous() && b.isContiguous() { + complexCmpRun(op, a.complexes[:n], b.complexes[:n], out.bools) + return out, nil + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyComplex(op, a.complexAt(i), b.complexAt(i)) + } + }) + case a.dt != Int && b.dt != Int: + if a.dt == b.dt && a.isContiguous() && b.isContiguous() { + switch a.dt { + case Float: + floatCmpRun(op, a.floats[:n], b.floats[:n], out.bools) + return out, nil + case Float32: + floatCmpRun(op, a.floats32[:n], b.floats32[:n], out.bools) + return out, nil + case Bool: + boolCmpRun(op, a.bools[:n], b.bools[:n], out.bools) + return out, nil + case Int8: + narrowCmpRun(op, a.i8s[:n], b.i8s[:n], out.bools) + return out, nil + case Uint8: + narrowCmpRun(op, a.u8s[:n], b.u8s[:n], out.bools) + return out, nil + case Int16: + narrowCmpRun(op, a.i16s[:n], b.i16s[:n], out.bools) + return out, nil + case Uint16: + narrowCmpRun(op, a.u16s[:n], b.u16s[:n], out.bools) + return out, nil + case Int32: + narrowCmpRun(op, a.i32s[:n], b.i32s[:n], out.bools) + return out, nil + case Uint32: + narrowCmpRun(op, a.u32s[:n], b.u32s[:n], out.bools) + return out, nil + } + } + if narrowIntClass(a.dt) && narrowIntClass(b.dt) && a.dt != b.dt && a.isContiguous() && b.isContiguous() { + // Two narrow widths: the mixed walk below widens both sides + // exactly, so the comparison runs over the int64 widenings + // in the raw kernel. Bool is an integer-class dtype but not + // a narrow payload, so it keeps the accessor walk. + narrowCmpMixed(op, a, b, n, out.bools) + return out, nil + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyFloat(op, a.floatAt(i), b.floatAt(i)) + } + }) + default: + // One side int, the other float: the int side stays exact, and + // the float side takes the comparison in its own order. + ints, floats, o := a, b, op + if b.dt == Int { + ints, floats, o = b, a, op.mirror() + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyIntFloat(o, ints.ints[ints.physIndex(i)], floats.floatAt(i)) + } + }) + } + return out, nil +} + +// cmpScalarI compares against an int scalar. Int arrays compare +// natively, float arrays compare exactly against the int through +// applyIntFloat rather than through the rounded float64(v), and complex +// arrays compare a zero imaginary part plus the exact real part. +func (a *Array) cmpScalarI(v int64, name string, op cmpOp) (*Array, error) { + if op.isOrdering() && a.dt == Complex { + return nil, errf("%s: complex arrays have no ordering", name) + } + n := a.Len() + out := &Array{shape: a.Shape(), dt: Bool, bools: make([]bool, n)} + switch a.dt { + case Int: + if a.isContiguous() { + intCmpScalarRun(op, a.ints[:n], v, out.bools) + return out, nil + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyInt(op, a.ints[a.physIndex(i)], v) + } + }) + case Complex: + // The scalar is an int, so the comparison is exact against the + // real part and asks whether the imaginary part is zero: the + // equality test's answer is inverted for every other relation. + want := op == opEq + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + z := a.complexes[a.physIndex(i)] + eq := imag(z) == 0 && applyIntFloat(opEq, v, real(z)) + out.bools[i] = eq == want + } + }) + default: + if a.isContiguous() { + // The scalar relation is a property of the call, so the + // mirror leaves the loop and the walk reads each payload + // slice directly, widening exactly as the accessor would. + scalarCmpIRun(op, a, v, out.bools) + return out, nil + } + // The mirror is a property of the call, so it leaves the loop. + mo := op.mirror() + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyIntFloat(mo, v, a.floatAt(i)) + } + }) + } + return out, nil +} + +// scalarCmpIRun compares a dense real payload against an int scalar. The +// float sides keep the exact int-versus-float relation applyIntFloat +// gives; the integer-class sides widen exactly, so the int comparison is +// the accessor walk's answer, its float detour included. +func scalarCmpIRun(op cmpOp, a *Array, v int64, dst []bool) { + mo := op.mirror() + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + zs := dst[s:e] + switch a.dt { + case Float: + xs := a.floats[s:e] + for i := range zs { + zs[i] = applyIntFloat(mo, v, xs[i]) + } + case Float32: + xs := a.floats32[s:e] + for i := range zs { + zs[i] = applyIntFloat(mo, v, float64(xs[i])) + } + case Float16: + xs := a.halves[s:e] + for i := range zs { + zs[i] = applyIntFloat(mo, v, HalfToFloat64(xs[i])) + } + case Bool: + xs := a.bools[s:e] + for i := range zs { + w := int64(0) + if xs[i] { + w = 1 + } + zs[i] = applyInt(mo, v, w) + } + case Int8: + xs := a.i8s[s:e] + for i := range zs { + zs[i] = applyInt(mo, v, int64(xs[i])) + } + case Uint8: + xs := a.u8s[s:e] + for i := range zs { + zs[i] = applyInt(mo, v, int64(xs[i])) + } + case Int16: + xs := a.i16s[s:e] + for i := range zs { + zs[i] = applyInt(mo, v, int64(xs[i])) + } + case Uint16: + xs := a.u16s[s:e] + for i := range zs { + zs[i] = applyInt(mo, v, int64(xs[i])) + } + case Int32: + xs := a.i32s[s:e] + for i := range zs { + zs[i] = applyInt(mo, v, int64(xs[i])) + } + default: // Uint32 + xs := a.u32s[s:e] + for i := range zs { + zs[i] = applyInt(mo, v, int64(xs[i])) + } + } + }) +} + +// cmpScalarF compares against a float scalar. Int arrays compare +// exactly through applyIntFloat, real arrays compare in float64, and +// complex arrays compare a zero imaginary part plus the real part. +func (a *Array) cmpScalarF(v float64, name string, op cmpOp) (*Array, error) { + if op.isOrdering() && a.dt == Complex { + return nil, errf("%s: complex arrays have no ordering", name) + } + n := a.Len() + out := &Array{shape: a.Shape(), dt: Bool, bools: make([]bool, n)} + switch a.dt { + case Int: + if a.isContiguous() { + intCmpScalarFRun(op, a.ints[:n], v, out.bools) + return out, nil + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyIntFloat(op, a.ints[a.physIndex(i)], v) + } + }) + case Complex: + z := complex(v, 0) + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyComplex(op, a.complexes[a.physIndex(i)], z) + } + }) + default: + if a.isContiguous() { + switch a.dt { + case Float: + floatCmpScalarRun(op, a.floats[:n], v, out.bools) + return out, nil + case Float32: + floatCmpScalarRun(op, a.floats32[:n], v, out.bools) + return out, nil + default: + scalarCmpFRun(op, a, v, out.bools) + return out, nil + } + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = applyFloat(op, a.floatAt(i), v) + } + }) + } + return out, nil +} + +// intCmpScalarFRun compares a dense int payload against a float scalar +// through the exact int-versus-float relation. +func intCmpScalarFRun(op cmpOp, x []int64, v float64, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, zs := x[s:e], dst[s:e] + for i := range zs { + zs[i] = applyIntFloat(op, xs[i], v) + } + }) +} + +// scalarCmpFRun compares a dense half, bool or narrow-integer payload +// against a float scalar, widening each element exactly as the accessor +// read would. The float64 and float32 widths carry their own op-hoisted +// kernels and never reach here. +func scalarCmpFRun(op cmpOp, a *Array, v float64, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + zs := dst[s:e] + switch a.dt { + case Float16: + xs := a.halves[s:e] + for i := range zs { + zs[i] = applyFloat(op, HalfToFloat64(xs[i]), v) + } + case Bool: + xs := a.bools[s:e] + for i := range zs { + w := 0.0 + if xs[i] { + w = 1 + } + zs[i] = applyFloat(op, w, v) + } + case Int8: + xs := a.i8s[s:e] + for i := range zs { + zs[i] = applyFloat(op, float64(xs[i]), v) + } + case Uint8: + xs := a.u8s[s:e] + for i := range zs { + zs[i] = applyFloat(op, float64(xs[i]), v) + } + case Int16: + xs := a.i16s[s:e] + for i := range zs { + zs[i] = applyFloat(op, float64(xs[i]), v) + } + case Uint16: + xs := a.u16s[s:e] + for i := range zs { + zs[i] = applyFloat(op, float64(xs[i]), v) + } + case Int32: + xs := a.i32s[s:e] + for i := range zs { + zs[i] = applyFloat(op, float64(xs[i]), v) + } + default: // Uint32 + xs := a.u32s[s:e] + for i := range zs { + zs[i] = applyFloat(op, float64(xs[i]), v) + } + } + }) +} + +// The comparison kernels below hold one loop per relation. The relation +// is a property of the call, not of the element, so the switch leaves +// the loop and the body holds the comparison itself; the destination of +// every element is its own slot, so the walk splits across workers like +// the arithmetic maps. Each kernel writes the same predicate the +// accessor walk computed, from the same operands. + +// intCmpRun compares two dense int payloads, writing true where op +// holds. +func intCmpRun(op cmpOp, x, y []int64, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, ys, zs := x[s:e], y[s:e], dst[s:e] + switch op { + case opEq: + for i := range zs { + zs[i] = xs[i] == ys[i] + } + case opNe: + for i := range zs { + zs[i] = xs[i] != ys[i] + } + case opLt: + for i := range zs { + zs[i] = xs[i] < ys[i] + } + case opLe: + for i := range zs { + zs[i] = xs[i] <= ys[i] + } + case opGt: + for i := range zs { + zs[i] = xs[i] > ys[i] + } + default: // opGe + for i := range zs { + zs[i] = xs[i] >= ys[i] + } + } + }) +} + +// floatCmpRun compares two dense real payloads of one width. Both +// widenings are exact for the comparison, matching the accessor walk's +// values exactly, and NaN follows IEEE because the comparison is the +// same one. +func floatCmpRun[T float32 | float64](op cmpOp, x, y []T, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, ys, zs := x[s:e], y[s:e], dst[s:e] + switch op { + case opEq: + for i := range zs { + zs[i] = float64(xs[i]) == float64(ys[i]) + } + case opNe: + for i := range zs { + zs[i] = float64(xs[i]) != float64(ys[i]) + } + case opLt: + for i := range zs { + zs[i] = float64(xs[i]) < float64(ys[i]) + } + case opLe: + for i := range zs { + zs[i] = float64(xs[i]) <= float64(ys[i]) + } + case opGt: + for i := range zs { + zs[i] = float64(xs[i]) > float64(ys[i]) + } + default: // opGe + for i := range zs { + zs[i] = float64(xs[i]) >= float64(ys[i]) + } + } + }) +} + +// complexCmpRun compares two dense complex payloads; an ordering op never +// reaches here, so the walk holds either equality or its negation. +func complexCmpRun(op cmpOp, x, y []complex128, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, ys, zs := x[s:e], y[s:e], dst[s:e] + if op == opNe { + for i := range zs { + zs[i] = xs[i] != ys[i] + } + return + } + for i := range zs { + zs[i] = xs[i] == ys[i] + } + }) +} + +// intCmpScalarRun compares a dense int payload against an int scalar. +func intCmpScalarRun(op cmpOp, x []int64, v int64, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, zs := x[s:e], dst[s:e] + switch op { + case opEq: + for i := range zs { + zs[i] = xs[i] == v + } + case opNe: + for i := range zs { + zs[i] = xs[i] != v + } + case opLt: + for i := range zs { + zs[i] = xs[i] < v + } + case opLe: + for i := range zs { + zs[i] = xs[i] <= v + } + case opGt: + for i := range zs { + zs[i] = xs[i] > v + } + default: // opGe + for i := range zs { + zs[i] = xs[i] >= v + } + } + }) +} + +// floatCmpScalarRun compares a dense real payload of one width against a +// float scalar, widening exactly. +func floatCmpScalarRun[T float32 | float64](op cmpOp, x []T, v float64, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, zs := x[s:e], dst[s:e] + switch op { + case opEq: + for i := range zs { + zs[i] = float64(xs[i]) == v + } + case opNe: + for i := range zs { + zs[i] = float64(xs[i]) != v + } + case opLt: + for i := range zs { + zs[i] = float64(xs[i]) < v + } + case opLe: + for i := range zs { + zs[i] = float64(xs[i]) <= v + } + case opGt: + for i := range zs { + zs[i] = float64(xs[i]) > v + } + default: // opGe + for i := range zs { + zs[i] = float64(xs[i]) >= v + } + } + }) +} + +// narrowCmpRun compares two dense same-width narrow integer payloads, +// writing true where op holds. Every widening the accessor walk takes is +// exact, so the comparison in the payload's own type sees the accessor +// walk's values and answers with its predicate. +func narrowCmpRun[T int8 | uint8 | int16 | uint16 | int32 | uint32](op cmpOp, x, y []T, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, ys, zs := x[s:e], y[s:e], dst[s:e] + switch op { + case opEq: + for i := range zs { + zs[i] = xs[i] == ys[i] + } + case opNe: + for i := range zs { + zs[i] = xs[i] != ys[i] + } + case opLt: + for i := range zs { + zs[i] = xs[i] < ys[i] + } + case opLe: + for i := range zs { + zs[i] = xs[i] <= ys[i] + } + case opGt: + for i := range zs { + zs[i] = xs[i] > ys[i] + } + default: // opGe + for i := range zs { + zs[i] = xs[i] >= ys[i] + } + } + }) +} + +// boolCmpRun compares two dense bool payloads: false orders below true, +// the order the accessor walk's 0/1 widening carried. +func boolCmpRun(op cmpOp, x, y, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, ys, zs := x[s:e], y[s:e], dst[s:e] + switch op { + case opEq: + for i := range zs { + zs[i] = xs[i] == ys[i] + } + case opNe: + for i := range zs { + zs[i] = xs[i] != ys[i] + } + case opLt: + for i := range zs { + zs[i] = !xs[i] && ys[i] + } + case opLe: + for i := range zs { + zs[i] = !xs[i] || ys[i] + } + case opGt: + for i := range zs { + zs[i] = xs[i] && !ys[i] + } + default: // opGe + for i := range zs { + zs[i] = xs[i] || !ys[i] + } + } + }) +} + +// narrowCmpMixed compares two dense narrow integer payloads of +// different widths. Both widenings are exact, so the int64 comparison +// is the accessor walk's float64 comparison unchanged: no narrow value +// ever rounds. +func narrowCmpMixed(op cmpOp, a, b *Array, n int, dst []bool) { + switch a.dt { + case Int8: + narrowCmpMixedRun(op, a.i8s[:n], b, dst) + case Uint8: + narrowCmpMixedRun(op, a.u8s[:n], b, dst) + case Int16: + narrowCmpMixedRun(op, a.i16s[:n], b, dst) + case Uint16: + narrowCmpMixedRun(op, a.u16s[:n], b, dst) + case Int32: + narrowCmpMixedRun(op, a.i32s[:n], b, dst) + default: // Uint32 + narrowCmpMixedRun(op, a.u32s[:n], b, dst) + } +} + +// narrowCmpMixedRun is narrowCmpMixed's kernel: the left payload's width +// is fixed by the caller, the right side dispatches per dtype inside. +func narrowCmpMixedRun[A int8 | uint8 | int16 | uint16 | int32 | uint32](op cmpOp, xs []A, b *Array, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + as := xs[s:e] + switch b.dt { + case Int8: + narrowCmpMixedLoop(op, as, b.i8s[s:e], dst[s:e]) + case Uint8: + narrowCmpMixedLoop(op, as, b.u8s[s:e], dst[s:e]) + case Int16: + narrowCmpMixedLoop(op, as, b.i16s[s:e], dst[s:e]) + case Uint16: + narrowCmpMixedLoop(op, as, b.u16s[s:e], dst[s:e]) + case Int32: + narrowCmpMixedLoop(op, as, b.i32s[s:e], dst[s:e]) + default: // Uint32 + narrowCmpMixedLoop(op, as, b.u32s[s:e], dst[s:e]) + } + }) +} + +// narrowCmpMixedLoop holds one comparison loop per relation over two +// differently typed narrow payloads, both widened exactly into int64. +func narrowCmpMixedLoop[A, B int8 | uint8 | int16 | uint16 | int32 | uint32](op cmpOp, xs []A, ys []B, zs []bool) { + switch op { + case opEq: + for i := range zs { + zs[i] = int64(xs[i]) == int64(ys[i]) + } + case opNe: + for i := range zs { + zs[i] = int64(xs[i]) != int64(ys[i]) + } + case opLt: + for i := range zs { + zs[i] = int64(xs[i]) < int64(ys[i]) + } + case opLe: + for i := range zs { + zs[i] = int64(xs[i]) <= int64(ys[i]) + } + case opGt: + for i := range zs { + zs[i] = int64(xs[i]) > int64(ys[i]) + } + default: // opGe + for i := range zs { + zs[i] = int64(xs[i]) >= int64(ys[i]) + } + } +} + +func whereRun[T any](cond []int64, x, y, dst []T) { + parallelMin(len(cond), elementwiseMinPerWorker, func(s, e int) { + cs, xs, ys, zs := cond[s:e], x[s:e], y[s:e], dst[s:e] + for i := range zs { + if cs[i] != 0 { + zs[i] = xs[i] + } else { + zs[i] = ys[i] + } + } + }) +} + +// whereHalfRun is whereRun for float16, which narrows through +// HalfFromFloat64 as the accessor walk does: a half NaN canonicalises on +// the way out, so the payload cannot be copied across directly. +func whereHalfRun(cond []int64, x, y, dst []uint16) { + parallelMin(len(cond), elementwiseMinPerWorker, func(s, e int) { + cs, xs, ys, zs := cond[s:e], x[s:e], y[s:e], dst[s:e] + for i := range zs { + v := ys[i] + if cs[i] != 0 { + v = xs[i] + } + zs[i] = HalfFromFloat64(HalfToFloat64(v)) + } + }) +} + +// Eq returns a bool mask that is true where a equals b element-wise. +func Eq(a, b *Array) (*Array, error) { return a.cmpArray(b, "Eq", opEq) } + +// Ne returns a bool mask that is true where a differs from b +// element-wise. +func Ne(a, b *Array) (*Array, error) { return a.cmpArray(b, "Ne", opNe) } + +// Lt returns a bool mask that is true where a is below b element-wise. +func Lt(a, b *Array) (*Array, error) { return a.cmpArray(b, "Lt", opLt) } + +// Le returns a bool mask that is true where a is at most b element-wise. +func Le(a, b *Array) (*Array, error) { return a.cmpArray(b, "Le", opLe) } + +// Gt returns a bool mask that is true where a is above b element-wise. +func Gt(a, b *Array) (*Array, error) { return a.cmpArray(b, "Gt", opGt) } + +// Ge returns a bool mask that is true where a is at least b element-wise. +func Ge(a, b *Array) (*Array, error) { return a.cmpArray(b, "Ge", opGe) } + +// EqI is Eq against an int scalar. +func EqI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Eq", opEq) } + +// NeI is Ne against an int scalar. +func NeI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Ne", opNe) } + +// LtI is Lt against an int scalar. +func LtI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Lt", opLt) } + +// LeI is Le against an int scalar. +func LeI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Le", opLe) } + +// GtI is Gt against an int scalar. +func GtI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Gt", opGt) } + +// GeI is Ge against an int scalar. +func GeI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Ge", opGe) } + +// EqF is Eq against a float scalar. +func EqF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Eq", opEq) } + +// NeF is Ne against a float scalar. +func NeF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Ne", opNe) } + +// LtF is Lt against a float scalar. +func LtF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Lt", opLt) } + +// LeF is Le against a float scalar. +func LeF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Le", opLe) } + +// GtF is Gt against a float scalar. +func GtF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Gt", opGt) } + +// GeF is Ge against a float scalar. +func GeF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Ge", opGe) } + +// Select returns a 1-D copy of the elements where the mask m, of the +// same shape, is set: a bool mask selects where it reads true, an int +// mask where it reads nonzero. Renamed from `Mask` because the +// operation is element selection, not bitmasking; the mask is just a +// selector array, bool natively or the int 0/1 the comparisons answered +// before the bool dtype carried them. +func Select(a, m *Array) (*Array, error) { + if m.dt != Bool && m.dt != Int { + return nil, errf("Select: the mask must be a bool or int array, got %s", m.dt) + } + if !sameShape(a.shape, m.shape) { + return nil, errf("Select: shape mismatch %s vs %s", shapeText(a.shape), shapeText(m.shape)) + } + n := m.Len() + hits := 0 + if m.dt == Bool { + for i := range n { + if m.bools[i] { + hits++ + } + } + } else { + for i := range n { + if m.ints[i] != 0 { + hits++ + } + } + } + out := &Array{shape: []int{hits}, dt: a.dt} + out.alloc(hits) + if hits == 0 { + return out, nil + } + if !a.isContiguous() { + // A view has no aligned payload, so the gather keeps the + // accessor per hit. + k := 0 + if m.dt == Bool { + for i := range n { + if m.bools[i] { + out.setFrom(k, a, i) + k++ + } + } + } else { + for i := range n { + if m.ints[i] != 0 { + out.setFrom(k, a, i) + k++ + } + } + } + return out, nil + } + switch a.dt { + case Int: + gatherMasked(out.ints, a.ints[:n], m, n) + case Bool: + gatherMasked(out.bools, a.bools[:n], m, n) + case Int8: + gatherMasked(out.i8s, a.i8s[:n], m, n) + case Uint8: + gatherMasked(out.u8s, a.u8s[:n], m, n) + case Int16: + gatherMasked(out.i16s, a.i16s[:n], m, n) + case Uint16: + gatherMasked(out.u16s, a.u16s[:n], m, n) + case Int32: + gatherMasked(out.i32s, a.i32s[:n], m, n) + case Uint32: + gatherMasked(out.u32s, a.u32s[:n], m, n) + case Float16: + gatherMasked(out.halves, a.halves[:n], m, n) + case Float32: + gatherMasked(out.floats32, a.floats32[:n], m, n) + case Float: + gatherMasked(out.floats, a.floats[:n], m, n) + default: + gatherMasked(out.complexes, a.complexes[:n], m, n) + } + return out, nil +} + +// gatherMasked runs the compact gather for one payload width, reading +// the mask as bool when it is one and as a nonzero int otherwise. +func gatherMasked[T any](dst, src []T, m *Array, n int) { + if m.dt == Bool { + selectCompactBool(dst, src, m.bools[:n]) + return + } + selectCompact(dst, src, m.ints[:n]) +} + +// selectPlan splits n elements into per-worker chunks and folds each +// chunk's hit count into the running offsets the gathers write from, +// the schedule both compact walks share. The split answer is false when +// the walk is too small to split, and the caller gathers serially. +func selectPlan(n int, hits func(lo, hi int) int) (chunk int, base []int, split bool) { + w := workersFor(n) + chunk = (n + w - 1) / w + if w == 1 || chunk < elementwiseMinPerWorker { + return chunk, nil, false + } + nch := (n + chunk - 1) / chunk + base = make([]int, nch) + var wg sync.WaitGroup + for c := range nch { + wg.Go(func() { + base[c] = hits(c*chunk, min((c+1)*chunk, n)) + }) + } + wg.Wait() + off := 0 + for c := range nch { + k := base[c] + base[c] = off + off += k + } + return chunk, base, true +} + +// selectCompact gathers the elements of src whose int mask slot is +// nonzero into dst in ascending order. Each chunk counts its hits first, +// so the prefix sum hands every worker its own destination offset and +// the writes stay output-disjoint: the same count-then-gather shape the +// serial walk had, with the scan split across workers. +func selectCompact[T any](dst, src []T, mask []int64) { + chunk, base, split := selectPlan(len(mask), func(lo, hi int) int { + k := 0 + for i := lo; i < hi; i++ { + if mask[i] != 0 { + k++ + } + } + return k + }) + if !split { + k := 0 + for i := range mask { + if mask[i] != 0 { + dst[k] = src[i] + k++ + } + } + return + } + n := len(mask) + var wg sync.WaitGroup + for c := range base { + wg.Go(func() { + lo, hi := c*chunk, min((c+1)*chunk, n) + k := base[c] + for i := lo; i < hi; i++ { + if mask[i] != 0 { + dst[k] = src[i] + k++ + } + } + }) + } + wg.Wait() +} + +// selectCompactBool gathers the elements of src whose bool mask slot is +// true, with selectCompact's contract and schedule. +func selectCompactBool[T any](dst, src []T, mask []bool) { + chunk, base, split := selectPlan(len(mask), func(lo, hi int) int { + k := 0 + for i := lo; i < hi; i++ { + if mask[i] { + k++ + } + } + return k + }) + if !split { + k := 0 + for i := range mask { + if mask[i] { + dst[k] = src[i] + k++ + } + } + return + } + n := len(mask) + var wg sync.WaitGroup + for c := range base { + wg.Go(func() { + lo, hi := c*chunk, min((c+1)*chunk, n) + k := base[c] + for i := lo; i < hi; i++ { + if mask[i] { + dst[k] = src[i] + k++ + } + } + }) + } + wg.Wait() +} + +// Where selects element-wise between x and y: a set element of cond +// picks x, an unset one picks y. cond is a bool mask or an int array, +// whose nonzero elements stand for true; every other dtype is refused +// with the wording the tests pin. All three arrays must share a shape; +// the result dtype is promote(x, y) over the full ladder. +func Where(cond, x, y *Array) (*Array, error) { + if cond.dt != Int && cond.dt != Bool { + return nil, errf("Where: the condition must be a bool or int array, got %s", cond.dt) + } + if !sameShape(cond.shape, x.shape) || !sameShape(cond.shape, y.shape) { + return nil, errf("Where: shapes %s, %s and %s must agree", + shapeText(cond.shape), shapeText(x.shape), shapeText(y.shape)) + } + n := cond.Len() + out := &Array{shape: x.Shape(), dt: promote(x.dt, y.dt)} + out.alloc(n) + // A dense int-masked result whose dtype is each operand's own dtype + // picks payload slots directly; a bool condition, a mixed pair, a + // narrow float16 result and a view keep the accessor walk below, + // which reads the same values. + if cond.dt == Int && cond.isContiguous() && x.isContiguous() && y.isContiguous() && x.dt == out.dt && y.dt == out.dt { + cs := cond.ints[:n] + switch out.dt { + case Int: + whereRun(cs, x.ints[:n], y.ints[:n], out.ints) + return out, nil + case Bool: + whereRun(cs, x.bools[:n], y.bools[:n], out.bools) + return out, nil + case Int8: + whereRun(cs, x.i8s[:n], y.i8s[:n], out.i8s) + return out, nil + case Uint8: + whereRun(cs, x.u8s[:n], y.u8s[:n], out.u8s) + return out, nil + case Int16: + whereRun(cs, x.i16s[:n], y.i16s[:n], out.i16s) + return out, nil + case Uint16: + whereRun(cs, x.u16s[:n], y.u16s[:n], out.u16s) + return out, nil + case Int32: + whereRun(cs, x.i32s[:n], y.i32s[:n], out.i32s) + return out, nil + case Uint32: + whereRun(cs, x.u32s[:n], y.u32s[:n], out.u32s) + return out, nil + case Float16: + // The half result narrows through HalfFromFloat64 as the + // accessor walk does, so a payload copy would not do. + whereHalfRun(cs, x.halves[:n], y.halves[:n], out.halves) + return out, nil + case Float32: + whereRun(cs, x.floats32[:n], y.floats32[:n], out.floats32) + return out, nil + case Float: + whereRun(cs, x.floats[:n], y.floats[:n], out.floats) + return out, nil + default: + whereRun(cs, x.complexes[:n], y.complexes[:n], out.complexes) + return out, nil + } + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + var nz bool + if cond.dt == Bool { + nz = cond.boolAt(i) + } else { + nz = cond.ints[i] != 0 + } + src := y + if nz { + src = x + } + switch out.dt { + case Int: + // intAt widens a narrow or bool operand exactly: a mixed + // integer-class pair promotes to int and both sides + // convert on the way in. + out.ints[i] = src.intAt(i) + case Bool: + out.bools[i] = src.boolAt(i) + case Int8: + out.i8s[i] = int8(src.intAt(i)) + case Uint8: + out.u8s[i] = uint8(src.intAt(i)) + case Int16: + out.i16s[i] = int16(src.intAt(i)) + case Uint16: + out.u16s[i] = uint16(src.intAt(i)) + case Int32: + out.i32s[i] = int32(src.intAt(i)) + case Uint32: + out.u32s[i] = uint32(src.intAt(i)) + case Float16: + // A half result reads either operand through the exact + // widening and narrows once. + out.halves[i] = HalfFromFloat64(src.floatAt(i)) + case Float32: + out.floats32[i] = src.float32At(i) + case Float: + out.floats[i] = src.floatAt(i) + default: + out.complexes[i] = src.complexAt(i) + } + } + }) + return out, nil +} + +// logicOp selects the boolean connective a logic operation applies. +type logicOp uint8 + +const ( + logicAnd logicOp = iota + logicOr + logicXor +) + +// And returns the element-wise conjunction of two bool arrays; both +// operands must be bool. +func And(a, b *Array) (*Array, error) { return logicArray(a, b, "And", logicAnd) } + +// Or returns the element-wise disjunction of two bool arrays; both +// operands must be bool. +func Or(a, b *Array) (*Array, error) { return logicArray(a, b, "Or", logicOr) } + +// Xor returns the element-wise exclusive disjunction of two bool +// arrays; both operands must be bool. +func Xor(a, b *Array) (*Array, error) { return logicArray(a, b, "Xor", logicXor) } + +// Not returns the element-wise negation of a bool array; the operand +// must be bool. +func Not(a *Array) (*Array, error) { + if a.dt != Bool { + return nil, errf("Not: operands must be bool arrays, got %s", a.dt) + } + n := a.Len() + out := &Array{shape: a.Shape(), dt: Bool} + out.alloc(n) + if a.isContiguous() { + src := a.bools + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + xs, os := src[s:e], out.bools[s:e] + for i := range os { + os[i] = !xs[i] + } + }) + return out, nil + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.bools[i] = !a.boolAt(i) + } + }) + return out, nil +} + +// logicArray runs the binary logic connectives: the only arithmetic-like +// surface the bool dtype carries, answered as a bool array. +func logicArray(a, b *Array, name string, op logicOp) (*Array, error) { + if a.dt != Bool || b.dt != Bool { + return nil, errf("%s: operands must be bool arrays, got %s and %s", name, a.dt, b.dt) + } + if !sameShape(a.shape, b.shape) { + return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape)) + } + n := a.Len() + out := &Array{shape: a.Shape(), dt: Bool} + out.alloc(n) + if a.isContiguous() && b.isContiguous() { + boolLogicRun(op, a.bools, b.bools, out.bools) + return out, nil + } + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + x, y := a.boolAt(i), b.boolAt(i) + switch op { + case logicAnd: + out.bools[i] = x && y + case logicOr: + out.bools[i] = x || y + default: // logicXor + out.bools[i] = x != y + } + } + }) + return out, nil +} + +// boolLogicRun applies one connective over two dense bool payloads, the +// relation a property of the call so the switch leaves the element loop. +func boolLogicRun(op logicOp, x, y, dst []bool) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + xs, ys, os := x[s:e], y[s:e], dst[s:e] + switch op { + case logicAnd: + for i := range os { + os[i] = xs[i] && ys[i] + } + case logicOr: + for i := range os { + os[i] = xs[i] || ys[i] + } + default: // logicXor + for i := range os { + os[i] = xs[i] != ys[i] + } + } + }) +} diff --git a/internal/core/mask_test.go b/internal/core/mask_test.go new file mode 100644 index 0000000..ac1541f --- /dev/null +++ b/internal/core/mask_test.go @@ -0,0 +1,414 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "slices" + "strings" + "testing" +) + +func maskOf(t *testing.T, a *Array) []int64 { + t.Helper() + out := make([]int64, a.Len()) + for i := range out { + out[i], _ = IntAt(a, i) + } + return out +} + +func TestComparisonsArray(t *testing.T) { + a := mustFromInts(t, []int64{1, 5, 5}, 3) + b := mustFromInts(t, []int64{5, 5, 9}, 3) + + m, err := Lt(a, b) + if err != nil { + t.Fatalf("Lt: %v", err) + } + if got := maskOf(t, m); got[0] != 1 || got[1] != 0 || got[2] != 1 { + t.Fatalf("Lt: %v", got) + } + + m, _ = Eq(a, b) + if got := maskOf(t, m); got[0] != 0 || got[1] != 1 || got[2] != 0 { + t.Fatalf("Eq: %v", got) + } + m, _ = Ne(a, b) + if got := maskOf(t, m); got[0] != 1 || got[1] != 0 || got[2] != 1 { + t.Fatalf("Ne: %v", got) + } + m, _ = Ge(a, b) + if got := maskOf(t, m); got[0] != 0 || got[1] != 1 || got[2] != 0 { + t.Fatalf("Ge: %v", got) + } + + // Mixed dtypes compare exactly: int against float widens per element + // without rounding the integer side through float64 first. + f := mustFromFloats(t, []float64{0.5, 5.0, 5.5}, 3) + m, _ = Gt(a, f) + if got := maskOf(t, m); got[0] != 1 || got[1] != 0 || got[2] != 0 { + t.Fatalf("Gt mixed: %v", got) + } + + if _, err := Lt(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "shape mismatch") { + t.Fatalf("comparison shape: %v", err) + } +} + +func TestComparisonsScalar(t *testing.T) { + a := mustFromInts(t, []int64{1, 5, 9}, 3) + + gt, err := GtI(a, 4) + if err != nil { + t.Fatalf("GtI: %v", err) + } + if got := maskOf(t, gt); got[0] != 0 || got[1] != 1 || got[2] != 1 { + t.Fatalf("GtI: %v", got) + } + le, _ := LeI(a, 5) + if got := maskOf(t, le); got[0] != 1 || got[1] != 1 || got[2] != 0 { + t.Fatalf("LeI: %v", got) + } + eq, _ := EqI(a, 5) + if got := maskOf(t, eq); got[1] != 1 || got[0] != 0 { + t.Fatalf("EqI: %v", got) + } + + f := mustFromFloats(t, []float64{0.5, 2.5}, 2) + lt, _ := LtF(f, 2.0) + if got := maskOf(t, lt); got[0] != 1 || got[1] != 0 { + t.Fatalf("LtF: %v", got) + } + gtf, _ := GtI(f, 0) + if got := maskOf(t, gtf); got[0] != 1 || got[1] != 1 { + t.Fatalf("GtI on float: %v", got) + } + + // IEEE NaN semantics: everything false except Ne. + nan := mustFromFloats(t, []float64{math.NaN()}, 1) + nanEq, _ := Eq(nan, nan) + if got := maskOf(t, nanEq); got[0] != 0 { + t.Fatalf("NaN Eq: %v", got) + } + nanLt, _ := Lt(nan, nan) + if got := maskOf(t, nanLt); got[0] != 0 { + t.Fatalf("NaN Lt: %v", got) + } + nanNe, _ := Ne(nan, nan) + if got := maskOf(t, nanNe); got[0] != 1 { + t.Fatalf("NaN Ne: %v", got) + } + nanEqF, _ := EqF(nan, math.NaN()) + if got := maskOf(t, nanEqF); got[0] != 0 { + t.Fatalf("NaN EqF: %v", got) + } +} + +func TestMaskSelect(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4}, 4) + mask, err := GtI(a, 2) + if err != nil { + t.Fatalf("GtI: %v", err) + } + + selected, err := Select(a, mask) + if err != nil { + t.Fatalf("Mask: %v", err) + } + want := mustFromInts(t, []int64{3, 4}, 2) + if !Equal(want, selected) { + t.Fatalf("Mask: %s", selected) + } + + // A float array keeps its dtype through masking. + f := mustFromFloats(t, []float64{0.5, 1.5, 2.5}, 3) + fm, err := GeF(f, 1.0) + if err != nil { + t.Fatalf("GeF: %v", err) + } + fs, err := Select(f, fm) + if err != nil || fs.Dtype() != Float || fs.Len() != 2 { + t.Fatalf("Mask float: %s %v", fs, err) + } + + // Masks compose through the bool logic operations: And is logical + // and, Or logical or. + m1, err := GtI(a, 2) // [false, false, true, true] + if err != nil { + t.Fatalf("GtI: %v", err) + } + m2, err := GeI(a, 2) // [false, true, true, true] + if err != nil { + t.Fatalf("GeI: %v", err) + } + if m1.Dtype() != Bool || m2.Dtype() != Bool { + t.Fatalf("comparisons answered %s and %s, want bool", m1.Dtype(), m2.Dtype()) + } + both, err := And(m1, m2) + if err != nil { + t.Fatalf("Mask compose: %v", err) + } + andSelected, _ := Select(a, both) + if !Equal(mustFromInts(t, []int64{3, 4}, 2), andSelected) { + t.Fatalf("Mask and: %s", andSelected) + } + either, err := Or(m1, m2) + if err != nil { + t.Fatalf("Mask compose or: %v", err) + } + orSelected, _ := Select(a, either) + if !Equal(mustFromInts(t, []int64{2, 3, 4}, 3), orSelected) { + t.Fatalf("Mask or: %s", orSelected) + } + + if _, err := Select(a, f); err == nil || !strings.Contains(err.Error(), "must be a bool or int array") { + t.Fatalf("Mask dtype: %v", err) + } + if _, err := Select(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "shape mismatch") { + t.Fatalf("Mask shape: %v", err) + } +} + +func TestWhere(t *testing.T) { + cond := mustFromInts(t, []int64{1, 0, 1}, 3) + x := mustFromInts(t, []int64{10, 20, 30}, 3) + y := mustFromInts(t, []int64{-1, -2, -3}, 3) + + out, err := Where(cond, x, y) + if err != nil { + t.Fatalf("Where: %v", err) + } + want := mustFromInts(t, []int64{10, -2, 30}, 3) + if !Equal(want, out) { + t.Fatalf("Where: %s", out) + } + + // Mixed dtypes promote to float. + fy := mustFromFloats(t, []float64{-0.5, -0.5, -0.5}, 3) + out, err = Where(cond, x, fy) + if err != nil || out.Dtype() != Float { + t.Fatalf("Where promote: %s %v", out, err) + } + if v, _ := FloatAt(out, 1); v != -0.5 { + t.Fatalf("Where promote value: %v", v) + } + + if _, err := Where(fy, x, y); err == nil || !strings.Contains(err.Error(), "condition must be a bool or int array") { + t.Fatalf("Where cond dtype: %v", err) + } + if _, err := Where(mustFromInts(t, []int64{1, 0}, 2), x, y); err == nil || !strings.Contains(err.Error(), "must agree") { + t.Fatalf("Where shape: %v", err) + } +} + +// TestFloatScalarComparisons pins all six relations of the scalar-float +// kernels against a hand-computed mask. The sample holds an element +// equal to the scalar, so both the <= and the >= boundary are read, and +// the float32 payload takes the same kernel one width down. +func TestFloatScalarComparisons(t *testing.T) { + f := mustFromFloats(t, []float64{-1, 0, 1, 2, 2.5, 3}, 2, 3) + cases := []struct { + name string + got func() (*Array, error) + want []bool + }{ + {"LtF", func() (*Array, error) { return LtF(f, 2) }, []bool{true, true, true, false, false, false}}, + {"LeF", func() (*Array, error) { return LeF(f, 2) }, []bool{true, true, true, true, false, false}}, + {"GtF", func() (*Array, error) { return GtF(f, 2) }, []bool{false, false, false, false, true, true}}, + {"GeF", func() (*Array, error) { return GeF(f, 2) }, []bool{false, false, false, true, true, true}}, + {"EqF", func() (*Array, error) { return EqF(f, 2) }, []bool{false, false, false, true, false, false}}, + {"NeF", func() (*Array, error) { return NeF(f, 2) }, []bool{true, true, true, false, true, true}}, + } + for _, tc := range cases { + m, err := tc.got() + if err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + if m.Dtype() != Bool { + t.Fatalf("%s answered dtype %s, want bool", tc.name, m.Dtype()) + } + got := m.RawBools() + if len(got) != len(tc.want) { + t.Fatalf("%s: %d elements, want %d", tc.name, len(got), len(tc.want)) + } + for i := range tc.want { + if got[i] != tc.want[i] { + t.Fatalf("%s over [-1 0 1 2 2.5 3] against 2: mask %v, want %v", + tc.name, got, tc.want) + } + } + } + // The float32 payload one width down, boundary included. + g := mustFromFloat32s(t, []float32{-1, 2, 2.5}, 3) + le, err := LeF(g, 2) + if err != nil { + t.Fatalf("LeF float32: %v", err) + } + if got := le.RawBools(); got[0] != true || got[1] != true || got[2] != false { + t.Fatalf("LeF float32 over [-1 2 2.5] against 2: %v, want [true true false]", got) + } + gt, err := GtF(g, 2) + if err != nil { + t.Fatalf("GtF float32: %v", err) + } + if got := gt.RawBools(); got[0] != false || got[1] != false || got[2] != true { + t.Fatalf("GtF float32 over [-1 2 2.5] against 2: %v, want [false false true]", got) + } +} + +// TestWhereDenseFloatOperands pins the operand order of the dense pick +// over two float payloads of one width: a nonzero condition takes the +// first operand, a zero the second, on both the float64 and the float32 +// kernel. +func TestWhereDenseFloatOperands(t *testing.T) { + cond := mustFromInts(t, []int64{1, 0, 1, 0, 1}, 5) + x := mustFromFloats(t, []float64{10, 20, 30, 40, 50}, 5) + y := mustFromFloats(t, []float64{-1, -2, -3, -4, -5}, 5) + out, err := Where(cond, x, y) + if err != nil { + t.Fatalf("Where float: %v", err) + } + if out.Dtype() != Float { + t.Fatalf("Where float dtype: %s", out.Dtype()) + } + for i, want := range []float64{10, -2, 30, -4, 50} { + if got := out.RawFloats()[i]; got != want { + t.Fatalf("Where float [%d] = %v, want %v", i, got, want) + } + } + // An all-zero and an all-one condition pin each arm in full. + zeros := mustFromInts(t, []int64{0, 0, 0, 0, 0}, 5) + low, err := Where(zeros, x, y) + if err != nil { + t.Fatalf("Where float zeros: %v", err) + } + for i := range 5 { + if got := low.RawFloats()[i]; got != y.RawFloats()[i] { + t.Fatalf("Where float zero condition [%d] = %v, want the second operand %v", + i, got, y.RawFloats()[i]) + } + } + ones := mustFromInts(t, []int64{1, 1, 1, 1, 1}, 5) + high, err := Where(ones, x, y) + if err != nil { + t.Fatalf("Where float ones: %v", err) + } + for i := range 5 { + if got := high.RawFloats()[i]; got != x.RawFloats()[i] { + t.Fatalf("Where float one condition [%d] = %v, want the first operand %v", + i, got, x.RawFloats()[i]) + } + } + // The float32 pair takes its own kernel. + c32 := mustFromInts(t, []int64{0, 1, 1}, 3) + x32 := mustFromFloat32s(t, []float32{1.5, -2.5, 7.25}, 3) + y32 := mustFromFloat32s(t, []float32{9, 9, 9}, 3) + out32, err := Where(c32, x32, y32) + if err != nil { + t.Fatalf("Where float32: %v", err) + } + for i, want := range []float32{9, -2.5, 7.25} { + if got := out32.RawFloat32s()[i]; got != want { + t.Fatalf("Where float32 [%d] = %v, want %v", i, got, want) + } + } +} + +// TestLogicOpsBoolContract pins And, Or, Xor and Not: bool operands +// only, element-wise boolean semantics, bool results, and a loud named +// refusal for every other dtype. +func TestLogicOpsBoolContract(t *testing.T) { + x := narrowBools(t, []bool{true, true, false, false}, 4) + y := narrowBools(t, []bool{true, false, true, false}, 4) + + check := func(name string, got *Array, err error, want []bool) { + t.Helper() + if err != nil { + t.Fatalf("%s: %v", name, err) + } + if got.Dtype() != Bool { + t.Fatalf("%s answered dtype %s, want bool", name, got.Dtype()) + } + if !slices.Equal(got.RawBools(), want) { + t.Fatalf("%s = %v, want %v", name, got.RawBools(), want) + } + } + and, err := And(x, y) + check("And", and, err, []bool{true, false, false, false}) + or, err := Or(x, y) + check("Or", or, err, []bool{true, true, true, false}) + xor, err := Xor(x, y) + check("Xor", xor, err, []bool{false, true, true, false}) + not, err := Not(x) + check("Not", not, err, []bool{false, false, true, true}) + + // The refusals name the actual dtypes. + ints := mustFromInts(t, []int64{1, 0, 1, 0}, 4) + if _, err := And(x, ints); err == nil || + !strings.Contains(err.Error(), "operands must be bool arrays, got bool and int") { + t.Fatalf("And with an int operand: %v", err) + } + if _, err := Not(ints); err == nil || + !strings.Contains(err.Error(), "operands must be bool arrays, got int") { + t.Fatalf("Not on an int operand: %v", err) + } + short := narrowBools(t, []bool{true}, 1) + if _, err := Or(x, short); err == nil || !strings.Contains(err.Error(), "shape mismatch") { + t.Fatalf("Or with a shape mismatch: %v", err) + } +} + +// TestWhereBoolCondition pins Where's condition contract: an int mask or +// a bool condition, the pinned refusal wording for every other dtype, +// and promoted result dtypes written correctly, narrow ones included. +func TestWhereBoolCondition(t *testing.T) { + cond := narrowBools(t, []bool{true, false, true, false}, 4) + x8 := narrowInt8s(t, []int8{1, 2, 3, 4}, 4) + y8 := narrowInt8s(t, []int8{9, 9, 9, 9}, 4) + + out, err := Where(cond, x8, y8) + if err != nil { + t.Fatalf("Where with a bool condition: %v", err) + } + if out.Dtype() != Int8 { + t.Fatalf("Where bool-cond int8 answered %s, want int8", out.Dtype()) + } + if want := []int8{1, 9, 3, 9}; !slices.Equal(out.RawInt8s(), want) { + t.Fatalf("Where bool-cond int8 = %v, want %v", out.RawInt8s(), want) + } + + // A mixed integer-class pair promotes through the containment table + // and the fallback walk writes the promoted payload. + u8y := narrowUint8s(t, []uint8{9, 9, 9, 9}, 4) + icond := mustFromInts(t, []int64{1, 0, 1, 0}, 4) + mix, err := Where(icond, x8, u8y) + if err != nil { + t.Fatalf("Where int8 with uint8: %v", err) + } + if mix.Dtype() != Int16 { + t.Fatalf("Where int8 with uint8 answered %s, want int16", mix.Dtype()) + } + if want := []int16{1, 9, 3, 9}; !slices.Equal(mix.RawInt16s(), want) { + t.Fatalf("Where int8 with uint8 = %v, want %v", mix.RawInt16s(), want) + } + + // A bool condition over bool operands answers bool. + xb := narrowBools(t, []bool{true, false, true, false}, 4) + yb := narrowBools(t, []bool{false, false, false, false}, 4) + bb, err := Where(cond, xb, yb) + if err != nil { + t.Fatalf("Where bool over bool: %v", err) + } + if bb.Dtype() != Bool || !slices.Equal(bb.RawBools(), []bool{true, false, true, false}) { + t.Fatalf("Where bool over bool = %s %v", bb.Dtype(), bb.RawBools()) + } + + // Every other condition dtype keeps the pinned refusal wording. + fc := mustFromFloats(t, []float64{1, 0, 1, 0}, 4) + if _, err := Where(fc, x8, y8); err == nil || + !strings.Contains(err.Error(), "condition must be a bool or int array") { + t.Fatalf("Where with a float condition: %v", err) + } +} diff --git a/internal/core/mat.go b/internal/core/mat.go new file mode 100644 index 0000000..ec8e762 --- /dev/null +++ b/internal/core/mat.go @@ -0,0 +1,1345 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/engine" + "sync" +) + +// Split thresholds for the matmul, transpose and gather kernels: below +// them the work is smaller than the goroutine fan-out it would pay +// for. Every floor is a measured constraint: the sweep benchmarks pin +// the last size where the split loses to the calling goroutine and the +// first where it wins, and the constant sits between the two. +const ( + // matVecParallelMin is the minimum element work (n·k or k·m) one + // worker must carry before a 2-D×1-D or 1-D×2-D product spreads + // across workers, so the fan-out is capped by how much work each + // worker gets rather than by the total. Measured on the vector + // sweep: a dot product costs a couple of cycles per element, and a + // worker handed a few hundred of them pays more for its goroutine + // than it saves, so a 16,384-element product split thirty-two ways + // measures up to three times slower than the lone walk, while the + // same product on four workers matches it. + matVecParallelMin = 16_384 + // matMulParallelMin is the minimum element-update work (n·k·m) + // before the Int, Float32 and Float 2-D×2-D products spread their + // output rows across workers. Measured on the square sweep: at + // 48 cubed (110,592 updates) the 32-way row split loses 2 to 22 + // percent against the lone wide-panel walk, at 64 cubed (262,144) + // it wins 5 to 11 percent, and larger products keep widening the + // gap; 131,072 sits between the losing and winning sizes. + matMulParallelMin = 131_072 + // matMulComplexParallelMin is the Complex path's floor on the + // same n·k·m measure. A complex element-update costs about four + // real multiplies, so the fan-out pays for itself four times + // earlier: at 16 cubed (4,096 updates) the split still loses 47 + // percent, at 32 cubed (32,768) it already wins 33 percent. + matMulComplexParallelMin = 32_768 + // colParallelMin is the row count from which a column gather runs + // across workers. The gather reads one strided element per row, + // so the row count is the measure: shapes with equal rows·cols + // work split opposite ways (32,768 rows of 8 columns wins, 4,096 + // rows of 64 columns loses 3.7 times). Measured down the row + // axis: 16,384 rows still lose 7 to 31 percent, 32,768 rows win + // 9 to 16 percent. + colParallelMin = 32_768 + // transposeSplitMin is the minimum element work (rows·cols) + // before the tiled transpose spreads its tiles across workers. + // Measured on the square and rectangular sweep: serial tiles win + // 1.5 to 2.3 times at and below 128×128, 256×256 and 512×512 + // measure even, and of the two 65,536-element rectangles one + // wins and one loses by a fifth, so the split engages from that + // size up. + transposeSplitMin = 65_536 + // transposeTileMin is the row and column count from which the + // transpose takes the tiled 2-D walk; smaller matrices, and every + // rank above two, stay on the odometer. + transposeTileMin = 8 + // widenParallelMin is the minimum element count before the exact + // widening walks of denseFloatsLocal and complexPayload spread + // across workers. Measured on the conversion sweeps: 4,096 + // elements run twice as fast serial, 16,384 elements win about + // 20 percent parallel, and both float64 and complex128 widening + // widen that gap with size. + widenParallelMin = 16_384 +) + +// Matrix operations with complex support. MatMul +// supports exactly 2-D×2-D, 2-D×1-D and 1-D×2-D; there is no batch +// broadcasting: the inner dimensions must agree or the error +// names both shapes. Transpose, Reshape, Row and Col all copy: results +// never alias their receiver. + +// narrowRefused reports whether dt is one of the narrow element types +// the arithmetic surfaces do not take: bool and the small integers. +// Every entry point without a kernel for them refuses them loudly by +// name, so the caller converts with Astype instead of reading a +// payload nothing computed. +func narrowRefused(dt Dtype) bool { + return dt == Bool || narrowIntClass(dt) +} + +// Identity builds the n×n identity matrix of the given dtype. +func Identity(dt Dtype, n int) (*Array, error) { + if n < 0 { + return nil, errf("Identity: n must be zero or greater, got %d", n) + } + out, err := Zeros(dt, n, n) + if err != nil { + return nil, err + } + for i := range n { + switch dt { + case Int: + out.ints[i*n+i] = 1 + case Bool: + out.bools[i*n+i] = true + case Int8: + out.i8s[i*n+i] = 1 + case Uint8: + out.u8s[i*n+i] = 1 + case Int16: + out.i16s[i*n+i] = 1 + case Uint16: + out.u16s[i*n+i] = 1 + case Int32: + out.i32s[i*n+i] = 1 + case Uint32: + out.u32s[i*n+i] = 1 + case Float16: + out.halves[i*n+i] = halfOne + case Float32: + out.floats32[i*n+i] = 1 + case Float: + out.floats[i*n+i] = 1 + default: + out.complexes[i*n+i] = 1 + } + } + return out, nil +} + +// MatMul returns the matrix product of a and b: 2-D×2-D, a matrix times a +// vector, or a vector times a matrix. The result dtype follows the +// promotion ladder: int with int stays int, any float or complex operand +// promotes. Float16 operands are refused loudly: the product kernels are +// not offered for the half dtype yet, and Astype conversion is cheap. The +// narrow element types are refused the same way until their kernels +// arrive. +func MatMul2D(a, b *Array) (*Array, error) { + if a.Dtype() == Float16 || b.Dtype() == Float16 { + return nil, errf("MatMul: float16 operands are not supported; convert with Astype") + } + for _, op := range []*Array{a, b} { + if narrowRefused(op.Dtype()) { + return nil, errf("MatMul: dtype %s is not supported; convert with Astype", op.Dtype()) + } + } + switch { + case a.NDim() == 2 && b.NDim() == 2: + n, k := a.Shape()[0], a.Shape()[1] + if b.Shape()[0] != k { + return nil, matmulMismatch(a, b) + } + m := b.Shape()[1] + out := &Array{shape: []int{n, m}} + // Every 2-D kernel preserves the plain i-p-j walk per cell: + // each output element receives exactly its products in + // ascending p order, so any split over rows is bit-identical + //. All dtypes take the row-panel kernels, which share + // one pass over b across a panel of output rows. Scalar + // mul-then-add sustains at most two floating-point operations + // per cycle per core, and the panel walk reaches that rate + // while streaming b linearly; measured against it, 2-D + // output-tile grids over register accumulators lose to + // per-element spill and bound-check overhead, and the row + // split's re-reads of b stay well inside the L3 bandwidth. + // The vector products split only above matVecParallelMin, + // where the fan-out pays for itself; the 2-D products + // gate their row split the same way through matMulSplit, so a + // small product runs whole on the calling goroutine, whose + // lone-worker walk reaches the wide panels a split would keep + // narrow. + work := n * k * m + switch promote(a.Dtype(), b.Dtype()) { + case Int: + out.dt = Int + out.ints = make([]int64, n*m) + matMulSplit(n, work, matMulParallelMin, func(rs, re int) { + matMulIntRows(a.ints, b.ints, out.ints, rs, re, k, m) + }) + case Float32: + // Accumulate in float64 and round once: float32 + // products are exact in float64, so the sums are never + // less accurate than a float32 kernel. The payloads are + // read directly, widening exactly where floatAt would, and + // each worker narrows its own finished rows once. + out.dt = Float32 + out.floats32 = make([]float32, n*m) + aIs32 := a.Dtype() == Float32 + bIs32 := b.Dtype() == Float32 + matMulSplit(n, work, matMulParallelMin, func(rs, re int) { + panel := engine.GetFloat64Buf(4 * m) + defer engine.PutFloat64Buf(panel) + switch { + case aIs32 && bIs32: + matMulF32Rows(a.floats32, b.floats32, out.floats32, panel, rs, re, k, m) + case aIs32: + matMulF32Rows(a.floats32, b.ints, out.floats32, panel, rs, re, k, m) + case bIs32: + matMulF32Rows(a.ints, b.floats32, out.floats32, panel, rs, re, k, m) + default: + matMulF32Rows(a.ints, b.ints, out.floats32, panel, rs, re, k, m) + } + }) + case Float: + out.dt = Float + out.floats = make([]float64, n*m) + // Mixed real operands convert once at entry, in the + // same element order, so the kernel streams plain rows. + // A float64 operand's payload already is that row + // stream, so it skips the conversion copy. + af := a.floats + if a.Dtype() != Float { + af = denseFloatsLocal(a, n, k) + } + bf := b.floats + if b.Dtype() != Float { + bf = denseFloatsLocal(b, k, m) + } + // A lone worker runs the wide 8-row panel; a split + // range runs the narrow one (see matMulF64Rows). A + // product below the floor is a lone worker too. + wide := work < matMulParallelMin || workersFor(n) == 1 + matMulSplit(n, work, matMulParallelMin, func(rs, re int) { + matMulF64Rows(af, bf, out.floats, rs, re, k, m, wide) + }) + default: + out.dt = Complex + out.complexes = make([]complex128, n*m) + ac := complexPayload(a) + bc := complexPayload(b) + matMulSplit(n, work, matMulComplexParallelMin, func(rs, re int) { + matMulComplexRows(ac, bc, out.complexes, rs, re, k, m) + }) + } + return out, nil + + case a.NDim() == 2 && b.NDim() == 1: + n, k := a.Shape()[0], a.Shape()[1] + if b.Len() != k { + return nil, matmulMismatch(a, b) + } + out := &Array{shape: []int{n}} + switch promote(a.Dtype(), b.Dtype()) { + case Int: + out.dt = Int + out.ints = make([]int64, n) + parallelSplit(n, n*k, func(rs, re int) { + matVecRows(a.ints, b.ints, out.ints, rs, re, k) + }) + case Float32: + // Accumulate in float64 and round once. The raw + // payload walk applies for every real payload: each + // widening is exact, so the values and their order match + // accessor reads bit for bit. + out.dt = Float32 + out.floats32 = make([]float32, n) + aIs32 := a.Dtype() == Float32 + bIs32 := b.Dtype() == Float32 + parallelSplit(n, n*k, func(rs, re int) { + switch { + case aIs32 && bIs32: + matVecF32Rows(a.floats32, b.floats32, out.floats32, rs, re, k) + case aIs32: + matVecF32Rows(a.floats32, b.ints, out.floats32, rs, re, k) + case bIs32: + matVecF32Rows(a.ints, b.floats32, out.floats32, rs, re, k) + default: + matVecF32Rows(a.ints, b.ints, out.floats32, rs, re, k) + } + }) + case Float: + out.dt = Float + out.floats = make([]float64, n) + if a.Dtype() == Float && b.Dtype() == Float { + parallelSplit(n, n*k, func(rs, re int) { + matVecF64Rows(a.floats, b.floats, out.floats, rs, re, k) + }) + } else { + // The mixed walk converts through the same accessor the + // scalar loop read, so the widened values and their + // order are unchanged. + af, bf := matVecF64Operands(a, n, k, b) + parallelSplit(n, n*k, func(rs, re int) { + matVecF64Rows(af, bf, out.floats, rs, re, k) + }) + } + default: + out.dt = Complex + out.complexes = make([]complex128, n) + ac := complexPayload(a) + bc := complexPayload(b) + parallelSplit(n, n*k, func(rs, re int) { + matVecRows(ac, bc, out.complexes, rs, re, k) + }) + } + return out, nil + + case a.NDim() == 1 && b.NDim() == 2: + k := a.Len() + if b.Shape()[0] != k { + return nil, matmulMismatch(a, b) + } + m := b.Shape()[1] + out := &Array{shape: []int{m}} + switch promote(a.Dtype(), b.Dtype()) { + case Int: + out.dt = Int + out.ints = make([]int64, m) + parallelSplit(m, k*m, func(js, je int) { + vecMatCols(a.ints[:k], b.ints, out.ints[js:je], m, js, je) + }) + case Float32: + // Accumulate in float64 and round once, exactly as + // the 2-D×1-D path does. + out.dt = Float32 + out.floats32 = make([]float32, m) + aIs32 := a.Dtype() == Float32 + bIs32 := b.Dtype() == Float32 + parallelSplit(m, k*m, func(js, je int) { + acc := engine.GetFloat64Buf(je - js) + defer engine.PutFloat64Buf(acc) + window := out.floats32[js:je] + switch { + case aIs32 && bIs32: + vecMatF32Cols(a.floats32[:k], b.floats32, window, acc, m, js, je) + case aIs32: + vecMatF32Cols(a.floats32[:k], b.ints, window, acc, m, js, je) + case bIs32: + vecMatF32Cols(a.ints, b.floats32, window, acc, m, js, je) + default: + vecMatF32Cols(a.ints, b.ints, window, acc, m, js, je) + } + }) + case Float: + out.dt = Float + out.floats = make([]float64, m) + if a.Dtype() == Float && b.Dtype() == Float { + parallelSplit(m, k*m, func(js, je int) { + vecMatF64Cols(a.floats[:k], b.floats, out.floats[js:je], m, js, je) + }) + } else { + af, bf := vecMatF64Operands(a, k, b, m) + parallelSplit(m, k*m, func(js, je int) { + vecMatF64Cols(af, bf, out.floats[js:je], m, js, je) + }) + } + default: + out.dt = Complex + out.complexes = make([]complex128, m) + ac := complexPayload(a) + bc := complexPayload(b) + parallelSplit(m, k*m, func(js, je int) { + vecMatCols(ac[:k], bc, out.complexes[js:je], m, js, je) + }) + } + return out, nil + } + return nil, errf("MatMul: unsupported shapes %s and %s", shapeText(a.Shape()), shapeText(b.Shape())) +} + +func matmulMismatch(a, b *Array) error { + return errf("MatMul: shape mismatch %s vs %s: the inner dimensions must agree", + shapeText(a.Shape()), shapeText(b.Shape())) +} + +// parallelSplit runs fn over [0, n) with the vector products' worker +// policy: each worker must carry at least matVecParallelMin elements, +// and a workload that cannot fill one worker that far runs whole on the +// calling goroutine. The floor counts elements per worker rather than +// elements in total: a dot product is a couple of cycles per element, +// so a worker needs a few thousand of them before the goroutine's +// creation and its share of the join amortise, and a total-work floor +// would hand a bare 16,384-element product to thirty-two workers at a +// few hundred elements each, which measures up to three times slower +// than the lone walk. +func parallelSplit(n, work int, fn func(start, end int)) { + splitCapped(n, work, matVecParallelMin, fn) +} + +// matMulSplit is the 2-D twin of parallelSplit: it runs fn over the n +// output rows of a 2-D×2-D product serially while the element-update +// work (n·k·m) stays below floor and across workers above it. The +// Complex path passes its own, lower floor because a complex +// element-update costs about four times the arithmetic of a real one. +func matMulSplit(n, work, floor int, fn func(start, end int)) { + if work < floor { + fn(0, n) + return + } + engine.Parallel(n, fn) +} + +// splitCapped runs fn over the [0, items) range across at most +// work/floor workers, serially when the whole workload sits below floor +// or the capped fan-out falls to one. It is the scheduling primitive of +// the line-based folds (axis reductions, tiled transpose): those kernels +// stream about one element per instruction, so a spawned goroutine must +// carry tens of thousands of elements before its creation and +// synchronisation amortise, and uncapped fan-out pays that spawn cost +// core-count times on workloads a few workers swallow whole (measured on +// the axis and transpose sweeps: the 32-way split of a 256×256 +// transpose or a 544-line sum costs more than the walk itself). +// Chunks stay contiguous and every item belongs to exactly one chunk, +// so a kernel whose accumulator slots are single-writer produces +// identical results under any split, including none. +func splitCapped(items, work, floor int, fn func(start, end int)) { + w := workersFor(items) + if work < floor { + w = 1 + } else if byWork := work / floor; w > byWork { + w = max(byWork, 1) + } + if w <= 1 { + fn(0, items) + return + } + chunk := (items + w - 1) / w + var wg sync.WaitGroup + for ls := 0; ls < items; ls += chunk { + le := min(ls+chunk, items) + wg.Go(func() { + fn(ls, le) + }) + } + wg.Wait() +} + +// matMulF64Rows multiplies rows [rs, re) of the n×k matrix a by the +// k×m matrix b into the same rows of out, which must arrive zeroed. +// The kernel walks the p dimension with a panel of output rows sharing +// one pass over b: the innermost loop is a range walk over the b row, +// so the compiled loop is pure pointer arithmetic, and each output +// element still receives exactly the products a scalar walk would add, +// in ascending p order. The result is bit-identical down to signed +// zeros and NaN payloads. The wide flag picks the 8-row panel for a +// lone worker, where it has the whole L1 to itself; under a real +// worker split the SMT siblings share that L1, so the narrow 4-row +// panels win. Remainder rows fall to 2- and 1-row panels with the same +// per-element order. +func matMulF64Rows(a, b, out []float64, rs, re, k, m int, wide bool) { + matMulF64Tile(a, b, out, rs, re, k, m, 0, m, wide) +} + +// matMulF64Tile multiplies rows [rs, re) of a by columns [js, je) of b +// into the same rows and columns of out, which must arrive zeroed. The +// whole product is the [0, n) × [0, m) case, and the tile form is what +// keeps a worker's slice of the output independent of every other +// worker's: rows are walked in 8-, 4-, 2- and 1-row panels and the +// columns of the tile in blocks of matMulColBlockMax, so the panel's +// output lines and one b row stay resident while the p walk streams b +// once per panel. Every output element receives its addends in +// ascending p order whatever tile carries it, so the result does not +// depend on the split. +func matMulF64Tile(a, b, out []float64, rs, re, k, m, js, je int, wide bool) { + step := min(je-js, matMulColBlockMax) + i := rs + if wide { + for ; i+8 <= re; i += 8 { + for j0 := js; j0 < je; j0 += step { + matMulF64Panel8(a, b, out, i, k, m, j0, min(j0+step, je)) + } + } + } + for ; i+4 <= re; i += 4 { + for j0 := js; j0 < je; j0 += step { + matMulF64Panel4(a, b, out, i, k, m, j0, min(j0+step, je)) + } + } + for ; i+2 <= re; i += 2 { + for j0 := js; j0 < je; j0 += step { + matMulF64Panel2(a, b, out, i, k, m, j0, min(j0+step, je)) + } + } + for ; i < re; i++ { + for j0 := js; j0 < je; j0 += step { + matMulF64Panel1(a, b, out, i, k, m, j0, min(j0+step, je)) + } + } +} + +// matMulColBlockMax bounds the column span of one 8-row panel pass. +// Eight output rows occupy 64 bytes per column, so a row count beyond +// matMulColBlockMax would push the panel out of L1 even for a lone +// worker; wider matrices sweep m in blocks and keep a block of every +// row resident instead. +const matMulColBlockMax = 512 + +// matMulF64Panel8 multiplies the 8-row tile at i by the column block +// [j0, j1), accumulating straight into the zeroed output rows. The j +// walk runs over the symbolic block width w, which is also the length +// every row slice expression above declares, so the compiler drops +// every per-element bounds check. +func matMulF64Panel8(a, b, out []float64, i, k, m, j0, j1 int) { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + a4 := a[(i+4)*k : (i+5)*k] + a5 := a[(i+5)*k : (i+6)*k] + a6 := a[(i+6)*k : (i+7)*k] + a7 := a[(i+7)*k : (i+8)*k] + o0 := out[i*m : (i+1)*m] + o1 := out[(i+1)*m : (i+2)*m] + o2 := out[(i+2)*m : (i+3)*m] + o3 := out[(i+3)*m : (i+4)*m] + o4 := out[(i+4)*m : (i+5)*m] + o5 := out[(i+5)*m : (i+6)*m] + o6 := out[(i+6)*m : (i+7)*m] + o7 := out[(i+7)*m : (i+8)*m] + w := j1 - j0 + r0 := o0[j0:j1] + r1 := o1[j0:j1] + r2 := o2[j0:j1] + r3 := o3[j0:j1] + r4 := o4[j0:j1] + r5 := o5[j0:j1] + r6 := o6[j0:j1] + r7 := o7[j0:j1] + for p, av0 := range a0 { + av1, av2, av3 := a1[p], a2[p], a3[p] + av4, av5, av6, av7 := a4[p], a5[p], a6[p], a7[p] + brow := b[p*m+j0 : p*m+j1] + for j := range w { + bv := brow[j] + r0[j] = float64(av0*bv) + r0[j] + r1[j] = float64(av1*bv) + r1[j] + r2[j] = float64(av2*bv) + r2[j] + r3[j] = float64(av3*bv) + r3[j] + r4[j] = float64(av4*bv) + r4[j] + r5[j] = float64(av5*bv) + r5[j] + r6[j] = float64(av6*bv) + r6[j] + r7[j] = float64(av7*bv) + r7[j] + } + } +} + +// matMulF64Panel4 is the 4-row tile: its j walk runs over the symbolic +// block width, so its per-element bounds checks drop the same way. +func matMulF64Panel4(a, b, out []float64, i, k, m, j0, j1 int) { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + o0 := out[i*m : (i+1)*m] + o1 := out[(i+1)*m : (i+2)*m] + o2 := out[(i+2)*m : (i+3)*m] + o3 := out[(i+3)*m : (i+4)*m] + w := j1 - j0 + r0 := o0[j0:j1] + r1 := o1[j0:j1] + r2 := o2[j0:j1] + r3 := o3[j0:j1] + for p, av0 := range a0 { + av1, av2, av3 := a1[p], a2[p], a3[p] + brow := b[p*m+j0 : p*m+j1] + for j := range w { + bv := brow[j] + r0[j] = float64(av0*bv) + r0[j] + r1[j] = float64(av1*bv) + r1[j] + r2[j] = float64(av2*bv) + r2[j] + r3[j] = float64(av3*bv) + r3[j] + } + } +} + +// matMulF64Panel2 is the 2-row tile. +func matMulF64Panel2(a, b, out []float64, i, k, m, j0, j1 int) { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + o0 := out[i*m : (i+1)*m] + o1 := out[(i+1)*m : (i+2)*m] + w := j1 - j0 + r0 := o0[j0:j1] + r1 := o1[j0:j1] + for p, av0 := range a0 { + av1 := a1[p] + brow := b[p*m+j0 : p*m+j1] + for j := range w { + bv := brow[j] + r0[j] = float64(av0*bv) + r0[j] + r1[j] = float64(av1*bv) + r1[j] + } + } +} + +// matMulF64Panel1 is the single-row walk. Every panel above writes the +// product through float64(a*b) + c: that spelling is what stops the +// compiler from contracting the multiply and the add into one fused +// operation at GOAMD64 v3 and later, which is what keeps the panels +// bit-for-bit identical on every machine and every pinned level. +func matMulF64Panel1(a, b, out []float64, i, k, m, j0, j1 int) { + arow := a[i*k : i*k+k] + orow := out[i*m : (i+1)*m] + r0 := orow[j0:j1] + for p, av := range arow { + brow := b[p*m+j0 : p*m+j1] + for j, bv := range brow { + r0[j] = float64(av*bv) + r0[j] + } + } +} + +// matMulComplexRows multiplies rows [rs, re) of the complex n×k matrix +// a by the complex k×m matrix b into the same rows of out, with 4-row +// panels sharing one pass over b: the plain per-row walk reads the whole +// k×m b once for every output row, which for a complex payload is a +// sixteen-byte element stream four times the width of the arithmetic +// that consumes it. A panel quarters that stream and leaves the complex +// arithmetic as the limit. Every output element still receives its +// addends in ascending p order, so the sums are unchanged. +func matMulComplexRows(a, b, out []complex128, rs, re, k, m int) { + step := min(m, matMulComplexColBlockMax) + i := rs + for ; i+4 <= re; i += 4 { + for j0 := 0; j0 < m; j0 += step { + matMulComplexPanel4(a, b, out, i, k, m, j0, min(j0+step, m)) + } + } + for ; i < re; i++ { + for j0 := 0; j0 < m; j0 += step { + matMulComplexPanel1(a, b, out, i, k, m, j0, min(j0+step, m)) + } + } +} + +// matMulComplexColBlockMax bounds the column span of one complex panel +// pass: a four-row panel of sixteen-byte elements fills 64 bytes per +// column, one cache line, so a wider block would push the accumulator +// rows out of L1 beside the b row band. +const matMulComplexColBlockMax = 256 + +// matMulComplexPanel4 is the 4-row complex tile. +func matMulComplexPanel4(a, b, out []complex128, i, k, m, j0, j1 int) { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + o0 := out[i*m : (i+1)*m] + o1 := out[(i+1)*m : (i+2)*m] + o2 := out[(i+2)*m : (i+3)*m] + o3 := out[(i+3)*m : (i+4)*m] + r0 := o0[j0:j1] + r1 := o1[j0:j1] + r2 := o2[j0:j1] + r3 := o3[j0:j1] + for p, av0 := range a0 { + av1, av2, av3 := a1[p], a2[p], a3[p] + brow := b[p*m+j0 : p*m+j1] + for j, bv := range brow { + r0[j] += av0 * bv + r1[j] += av1 * bv + r2[j] += av2 * bv + r3[j] += av3 * bv + } + } +} + +// matMulComplexPanel1 is the single-row complex walk. +func matMulComplexPanel1(a, b, out []complex128, i, k, m, j0, j1 int) { + arow := a[i*k : i*k+k] + orow := out[i*m : (i+1)*m] + r0 := orow[j0:j1] + for p, av := range arow { + brow := b[p*m+j0 : p*m+j1] + for j, bv := range brow { + r0[j] += av * bv + } + } +} + +// matMulF32Rows is matMulF64Rows for a Float32 result. Float32 outputs +// cannot hold the running sums, so each tile accumulates in a float64 +// panel buffer (one m-wide row per output row, supplied zeroed by the +// caller) and narrows each finished row once. The operands keep +// their real payloads; every widening here is the exact conversion the +// accessor performs, so the products and their order are unchanged. The +// panel sweeps the whole row rather than blocking its columns: the +// measured sweep keeps a full row for every panel that fits L1, and the +// float32 rows that do not fit lose more to the strided b walk a block +// would create than they gain from the smaller panel. +func matMulF32Rows[A, B int64 | float32](a []A, b []B, out []float32, panel []float64, rs, re, k, m int) { + i := rs + for ; i+4 <= re; i += 4 { + matMulF32Panel4(a, b, out, panel, i, k, m) + } + for ; i < re; i++ { + matMulF32Panel1(a, b, out, panel, i, k, m) + } +} + +// matMulF32Panel4 multiplies the 4-row tile at i through the float64 +// panel, then narrows the four rows in one pass. +func matMulF32Panel4[A, B int64 | float32](a []A, b []B, out []float32, panel []float64, i, k, m int) { + clear(panel) + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + s0 := panel[0:m] + s1 := panel[m : 2*m] + s2 := panel[2*m : 3*m] + s3 := panel[3*m : 4*m] + for p, av0 := range a0 { + av0f := float64(av0) + av1, av2, av3 := float64(a1[p]), float64(a2[p]), float64(a3[p]) + brow := b[p*m : p*m+m] + for j, bv := range brow { + bvf := float64(bv) + s0[j] += av0f * bvf + s1[j] += av1 * bvf + s2[j] += av2 * bvf + s3[j] += av3 * bvf + } + } + o0 := out[i*m : (i+1)*m] + o1 := out[(i+1)*m : (i+2)*m] + o2 := out[(i+2)*m : (i+3)*m] + o3 := out[(i+3)*m : (i+4)*m] + for j, v := range s0 { + o0[j] = float32(v) + o1[j] = float32(s1[j]) + o2[j] = float32(s2[j]) + o3[j] = float32(s3[j]) + } +} + +// matMulF32Panel1 multiplies the single row at i through the float64 +// panel and narrows it once. +func matMulF32Panel1[A, B int64 | float32](a []A, b []B, out []float32, panel []float64, i, k, m int) { + clear(panel) + arow := a[i*k : i*k+k] + s0 := panel[0:m] + for p, av := range arow { + avf := float64(av) + brow := b[p*m : p*m+m] + for j, bv := range brow { + s0[j] += avf * float64(bv) + } + } + orow := out[i*m : (i+1)*m] + for j, v := range s0 { + orow[j] = float32(v) + } +} + +// matMulIntRows multiplies rows [rs, re) with 4-row panels sharing one +// pass over b. Int addition wraps exactly like the in-memory walk it +// replaces, and each column still receives its addends in ascending p +// order, so the wrapped sums are identical. +func matMulIntRows(a, b, out []int64, rs, re, k, m int) { + i := rs + for ; i+4 <= re; i += 4 { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + o0 := out[i*m : (i+1)*m] + o1 := out[(i+1)*m : (i+2)*m] + o2 := out[(i+2)*m : (i+3)*m] + o3 := out[(i+3)*m : (i+4)*m] + for p, av0 := range a0 { + av1, av2, av3 := a1[p], a2[p], a3[p] + brow := b[p*m : p*m+m] + for j, bv := range brow { + o0[j] += av0 * bv + o1[j] += av1 * bv + o2[j] += av2 * bv + o3[j] += av3 * bv + } + } + } + for ; i < re; i++ { + arow := a[i*k : i*k+k] + orow := out[i*m : (i+1)*m] + for p, av := range arow { + brow := b[p*m : p*m+m] + for j, bv := range brow { + orow[j] += av * bv + } + } + } +} + +// matVecF64Rows writes rows [rs, re) of a·b for the n×k matrix a and +// the k vector b, walking four rows at once: each row keeps its own +// ascending-p sum, so the four independent chains only add +// instruction-level parallelism, never reassociation. +func matVecF64Rows(a, b, out []float64, rs, re, k int) { + i := rs + for ; i+4 <= re; i += 4 { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + var s0, s1, s2, s3 float64 + for p := range k { + bv := b[p] + s0 += a0[p] * bv + s1 += a1[p] * bv + s2 += a2[p] * bv + s3 += a3[p] * bv + } + out[i], out[i+1], out[i+2], out[i+3] = s0, s1, s2, s3 + } + for ; i < re; i++ { + arow := a[i*k : i*k+k] + var s float64 + for p, av := range arow { + s += av * b[p] + } + out[i] = s + } +} + +// matVecF32Rows is matVecF64Rows for a Float32 result: one float64 sum +// per row, narrowed once, read straight from the real +// payloads. Four rows walk together like the float64 twin, so the four +// independent sum chains fill the pipeline a single row's latency-bound +// chain leaves idle. +func matVecF32Rows[A, B int64 | float32](a []A, b []B, out []float32, rs, re, k int) { + i := rs + for ; i+4 <= re; i += 4 { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + var s0, s1, s2, s3 float64 + for p := range k { + bv := float64(b[p]) + s0 += float64(a0[p]) * bv + s1 += float64(a1[p]) * bv + s2 += float64(a2[p]) * bv + s3 += float64(a3[p]) * bv + } + out[i], out[i+1], out[i+2], out[i+3] = float32(s0), float32(s1), float32(s2), float32(s3) + } + for ; i < re; i++ { + arow := a[i*k : i*k+k] + var s float64 + for p, av := range arow { + s += float64(av) * float64(b[p]) + } + out[i] = float32(s) + } +} + +// matVecRows is the row-dot walk for the int and complex payloads; +// each row's sum is independent, so a worker split cannot move a +// single addend. The rows walk together, as in matVecF64Rows, but the +// int payload takes four rows and the complex payload two: a complex +// chain carries two register-wide values, so four of them fill the +// register file and spill exactly what the unroll was hiding. +func matVecRows[T int64 | complex128](a, b, out []T, rs, re, k int) { + var zero T + rows := 4 + if _, isInt := any(zero).(int64); !isInt { + rows = 2 + } + i := rs + if rows == 4 { + for ; i+4 <= re; i += 4 { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + a2 := a[(i+2)*k : (i+3)*k] + a3 := a[(i+3)*k : (i+4)*k] + var s0, s1, s2, s3 T + for p := range k { + bv := b[p] + s0 += a0[p] * bv + s1 += a1[p] * bv + s2 += a2[p] * bv + s3 += a3[p] * bv + } + out[i], out[i+1], out[i+2], out[i+3] = s0, s1, s2, s3 + } + } else { + for ; i+2 <= re; i += 2 { + a0 := a[i*k : i*k+k] + a1 := a[(i+1)*k : (i+2)*k] + var s0, s1 T + for p := range k { + bv := b[p] + s0 += a0[p] * bv + s1 += a1[p] * bv + } + out[i], out[i+1] = s0, s1 + } + } + for ; i < re; i++ { + arow := a[i*k : i*k+k] + var s T + for p, av := range arow { + s += av * b[p] + } + out[i] = s + } +} + +// matVecF64Operands resolves a 2-D×1-D product's operands to float64 +// row streams: a float64 payload already is one, everything else +// converts once through the accessor's exact widening. +func matVecF64Operands(a *Array, n, k int, b *Array) ([]float64, []float64) { + var af, bf []float64 + if a.Dtype() == Float { + af = a.floats + } else { + af = denseFloatsLocal(a, n, k) + } + if b.Dtype() == Float { + bf = b.floats + } else { + bf = denseFloatsLocal(b, 1, k) + } + return af, bf +} + +// vecMatF64Cols multiplies a by columns [js, je) of b into out, which +// is the zeroed [js, je) window of the result: the b row is a range +// walk, the compiled loop is pure pointer arithmetic, and every output +// column receives its addends in ascending p order. +func vecMatF64Cols(a, b, out []float64, m, js, je int) { + if je-js <= vecMatColBandMax { + vecMatF64ColsNarrow(a, b, out, m, js, je) + return + } + for p, av := range a { + brow := b[p*m+js : p*m+je] + for j, bv := range brow { + out[j] += av * bv + } + } +} + +// vecMatColBandMax is the widest output band the column walk accumulates +// in registers, one cache line of float64. A band that narrow leaves the +// b rows' lines half read whatever the walk does, and then the per-p +// load and store of every accumulator in out costs more than the walk's +// own arithmetic: measured across the vector sweep, a sixteen-column +// output of long vectors runs fifteen times slower through the memory +// accumulators than through four register ones. +const vecMatColBandMax = 8 + +// vecMatF64ColsNarrow walks a narrow band through register +// accumulators, four columns at a time: each column still sums its +// addends in ascending p order, so the values are those the memory walk +// produces. +func vecMatF64ColsNarrow(a, b, out []float64, m, js, je int) { + w := je - js + j := 0 + for ; j+4 <= w; j += 4 { + var s0, s1, s2, s3 float64 + for p, av := range a { + brow := b[p*m+js+j : p*m+js+j+4] + s0 += av * brow[0] + s1 += av * brow[1] + s2 += av * brow[2] + s3 += av * brow[3] + } + out[j], out[j+1], out[j+2], out[j+3] = s0, s1, s2, s3 + } + for ; j < w; j++ { + var s float64 + for p, av := range a { + s += av * b[p*m+js+j] + } + out[j] = s + } +} + +// vecMatF32Cols is vecMatF64Cols for a Float32 result: the window +// accumulates in a float64 scratch (cleared by the caller) and narrows +// once at the end. +func vecMatF32Cols[A, B int64 | float32](a []A, b []B, out []float32, acc []float64, m, js, je int) { + clear(acc) + for p, av := range a { + avf := float64(av) + brow := b[p*m+js : p*m+je] + for j, bv := range brow { + acc[j] += avf * float64(bv) + } + } + for j, v := range acc { + out[j] = float32(v) + } +} + +// vecMatCols is the column walk for the int and complex payloads; +// each output column receives its addends in ascending p order, and +// int wrapping is unchanged by the visit order. +func vecMatCols[T int64 | complex128](a, b, out []T, m, js, je int) { + for p, av := range a { + brow := b[p*m+js : p*m+je] + for j, bv := range brow { + out[j] += av * bv + } + } +} + +// vecMatF64Operands resolves a 1-D×2-D product's operands to float64 +// row streams, the 1-D twin of matVecF64Operands. +func vecMatF64Operands(a *Array, k int, b *Array, m int) ([]float64, []float64) { + var af, bf []float64 + if a.Dtype() == Float { + af = a.floats + } else { + af = denseFloatsLocal(a, 1, k) + } + if b.Dtype() == Float { + bf = b.floats + } else { + bf = denseFloatsLocal(b, k, m) + } + return af, bf +} + +// complexPayload returns the array's elements as complex values, +// reading a complex payload directly and converting everything else +// once; every array this package produces stores elements at their +// flat payload index. The conversion walk splits across workers from +// widenParallelMin up: each element widens exactly as the scalar read +// does and lands in its own slot, so the values are unchanged. +func complexPayload(a *Array) []complex128 { + if a.Dtype() == Complex { + return a.complexes + } + out := make([]complex128, a.Len()) + if len(out) < widenParallelMin { + for i := range out { + out[i] = a.ComplexAt(i) + } + return out + } + engine.Parallel(len(out), func(ws, we int) { + for i := ws; i < we; i++ { + out[i] = a.ComplexAt(i) + } + }) + return out +} + +// Transpose returns a new array with the dimensions reversed; on 1-D it +// is a copy. It is infallible. +func Transpose(a *Array) *Array { + sh := a.Shape() + d := len(sh) + newShape := make([]int, d) + for i := range newShape { + newShape[i] = sh[d-1-i] + } + out := &Array{shape: newShape, dt: a.dt} + n := a.Len() + out.alloc(n) + // Rank-2 arrays take the tiled walk: cache-sized tiles keep both + // sides' lines alive and the tile range spreads across workers. + // The permutation is unchanged, only the visit order, so the copy + // stays bit-identical. Every other rank, and matrices below + // transposeTileMin in either dimension, keep the odometer below. + if d == 2 && sh[0] >= transposeTileMin && sh[1] >= transposeTileMin { + switch a.dt { + case Int: + transposeTiles(a.ints, out.ints, sh[0], sh[1]) + case Bool: + transposeTiles(a.bools, out.bools, sh[0], sh[1]) + case Int8: + transposeTiles(a.i8s, out.i8s, sh[0], sh[1]) + case Uint8: + transposeTiles(a.u8s, out.u8s, sh[0], sh[1]) + case Int16: + transposeTiles(a.i16s, out.i16s, sh[0], sh[1]) + case Uint16: + transposeTiles(a.u16s, out.u16s, sh[0], sh[1]) + case Int32: + transposeTiles(a.i32s, out.i32s, sh[0], sh[1]) + case Uint32: + transposeTiles(a.u32s, out.u32s, sh[0], sh[1]) + case Float16: + transposeTiles(a.halves, out.halves, sh[0], sh[1]) + case Float32: + transposeTiles(a.floats32, out.floats32, sh[0], sh[1]) + case Float: + transposeTiles(a.floats, out.floats, sh[0], sh[1]) + default: + transposeTiles(a.complexes, out.complexes, sh[0], sh[1]) + } + return out + } + // Under the reversed shape the destination flat index is the + // column-major index of the source coordinates: each advancing + // dimension contributes its source stride, so dst = Σ coord[k]·sw[k] + // is the same mapping the materialised index table used to hold, + // computed one pass earlier. The dtype dispatch keeps setFrom's + // per-element switch out of the copy loops. + sw := make([]int, d) + for k := range sw { + sw[k] = 1 + for _, s := range sh[:k] { + sw[k] *= s + } + } + coord := make([]int, d) + switch a.dt { + case Int: + for src := range n { + dst := 0 + for k := range coord { + dst += coord[k] * sw[k] + } + out.ints[dst] = a.ints[src] + advanceOdometer(coord, sh) + } + case Float16: + for src := range n { + dst := 0 + for k := range coord { + dst += coord[k] * sw[k] + } + out.halves[dst] = a.halves[src] + advanceOdometer(coord, sh) + } + case Float32: + for src := range n { + dst := 0 + for k := range coord { + dst += coord[k] * sw[k] + } + out.floats32[dst] = a.floats32[src] + advanceOdometer(coord, sh) + } + case Float: + for src := range n { + dst := 0 + for k := range coord { + dst += coord[k] * sw[k] + } + out.floats[dst] = a.floats[src] + advanceOdometer(coord, sh) + } + case Complex: + for src := range n { + dst := 0 + for k := range coord { + dst += coord[k] * sw[k] + } + out.complexes[dst] = a.complexes[src] + advanceOdometer(coord, sh) + } + default: + // Bool and the narrow integer widths carry no per-dtype loop + // here; setFrom writes them with the same values in the same + // order. + for src := range n { + dst := 0 + for k := range coord { + dst += coord[k] * sw[k] + } + out.setFrom(dst, a, src) + advanceOdometer(coord, sh) + } + } + return out +} + +// transposeTiles copies the rows×cols matrix src into its transpose dst +// through 32×32 tiles. The tile size keeps both sides' working set +// inside L1; the visit order is free because every destination slot is +// written exactly once. The tile range splits across workers only from +// transposeSplitMin up, and the fan-out is capped by splitCapped so +// every worker carries at least that many elements: below the floor the +// measured sweep keeps the whole walk faster on the calling goroutine, +// which still takes the tiled route, only alone. +func transposeTiles[T any](src, dst []T, rows, cols int) { + const tile = 32 + colTiles := (cols + tile - 1) / tile + tiles := ((rows + tile - 1) / tile) * colTiles + walk := func(ts, te int) { + var stage [tile * tile]T + for t := ts; t < te; t++ { + i0 := (t / colTiles) * tile + j0 := (t % colTiles) * tile + i1 := min(i0+tile, rows) + j1 := min(j0+tile, cols) + w := j1 - j0 + h := i1 - i0 + // Both proofs sit outside the loops they serve: every stage + // read below lands inside the h×w tile and every drow store + // inside its h-element run, so the inner walks run without + // per-element bounds checks. + _ = stage[h*w-1] + for i := i0; i < i1; i++ { + srow := src[i*cols+j0 : i*cols+j1] + copy(stage[(i-i0)*w:(i-i0)*w+w], srow) + } + for j := j0; j < j1; j++ { + drow := dst[j*rows+i0 : j*rows+i1] + _ = drow[h-1] + // si steps down the stage column: the tile's transposed + // neighbours sit w slots apart. + si := j - j0 + for i := range h { + drow[i] = stage[si] + si += w + } + } + } + } + splitCapped(tiles, rows*cols, transposeSplitMin, walk) +} + +// Reshape returns a copy with a new shape of the same element count. +func Reshape(a *Array, shape ...int) (*Array, error) { + total, sh, err := checkedDims(shape) + if err != nil { + return nil, err + } + if total != a.Len() { + return nil, errf("Reshape: %d elements do not fill the shape %s", a.Len(), shapeText(sh)) + } + // cloneArray carries every payload the dtype owns: the narrow + // element types have no five-slice clone to build from. + out := a.cloneArray() + out.shape = sh + return out, nil +} + +// Row returns a 1-D copy of row i; the array must be 2-D. +func Row(a *Array, i int) (*Array, error) { + if a.NDim() != 2 { + return nil, errf("Row: needs a 2-D array, got shape %s", shapeText(a.Shape())) + } + if i < 0 || i >= a.Shape()[0] { + return nil, errf("Row: index %d is out of range for %d rows", i, a.Shape()[0]) + } + cols := a.Shape()[1] + out := &Array{shape: []int{cols}, dt: a.Dtype()} + out.alloc(cols) + // A row is one contiguous payload block, so each dtype copies it in + // a single move. + switch a.dt { + case Int: + copy(out.ints, a.ints[i*cols:(i+1)*cols]) + case Bool: + copy(out.bools, a.bools[i*cols:(i+1)*cols]) + case Int8: + copy(out.i8s, a.i8s[i*cols:(i+1)*cols]) + case Uint8: + copy(out.u8s, a.u8s[i*cols:(i+1)*cols]) + case Int16: + copy(out.i16s, a.i16s[i*cols:(i+1)*cols]) + case Uint16: + copy(out.u16s, a.u16s[i*cols:(i+1)*cols]) + case Int32: + copy(out.i32s, a.i32s[i*cols:(i+1)*cols]) + case Uint32: + copy(out.u32s, a.u32s[i*cols:(i+1)*cols]) + case Float16: + copy(out.halves, a.halves[i*cols:(i+1)*cols]) + case Float32: + copy(out.floats32, a.floats32[i*cols:(i+1)*cols]) + case Float: + copy(out.floats, a.floats[i*cols:(i+1)*cols]) + default: + copy(out.complexes, a.complexes[i*cols:(i+1)*cols]) + } + return out, nil +} + +// Col returns a 1-D copy of column j; the array must be 2-D. +func Col(a *Array, j int) (*Array, error) { + if a.NDim() != 2 { + return nil, errf("Col: needs a 2-D array, got shape %s", shapeText(a.Shape())) + } + if j < 0 || j >= a.Shape()[1] { + return nil, errf("Col: index %d is out of range for %d columns", j, a.Shape()[1]) + } + rows, cols := a.Shape()[0], a.Shape()[1] + out := &Array{shape: []int{rows}, dt: a.Dtype()} + out.alloc(rows) + // Column reads stride by the row width; the dtype dispatch keeps + // the per-element accessor switch out of the gather, and tall + // matrices spread the strided reads across workers. + switch a.dt { + case Int: + colGather(a.ints, out.ints, rows, cols, j) + case Bool: + colGather(a.bools, out.bools, rows, cols, j) + case Int8: + colGather(a.i8s, out.i8s, rows, cols, j) + case Uint8: + colGather(a.u8s, out.u8s, rows, cols, j) + case Int16: + colGather(a.i16s, out.i16s, rows, cols, j) + case Uint16: + colGather(a.u16s, out.u16s, rows, cols, j) + case Int32: + colGather(a.i32s, out.i32s, rows, cols, j) + case Uint32: + colGather(a.u32s, out.u32s, rows, cols, j) + case Float16: + colGather(a.halves, out.halves, rows, cols, j) + case Float32: + colGather(a.floats32, out.floats32, rows, cols, j) + case Float: + colGather(a.floats, out.floats, rows, cols, j) + default: + colGather(a.complexes, out.complexes, rows, cols, j) + } + return out, nil +} + +// colGather copies column j into dst, serially below colParallelMin +// and across workers above it: the gather is a pure permutation, so +// the split cannot move a single element. The floor counts rows, not +// rows·cols: the walk reads one strided element per row, and the sweep +// splits equal-work shapes opposite ways down the row axis. +func colGather[T any](src, dst []T, rows, cols, j int) { + if rows < colParallelMin { + for r := range rows { + dst[r] = src[r*cols+j] + } + return + } + engine.Parallel(rows, func(rs, re int) { + for r := rs; r < re; r++ { + dst[r] = src[r*cols+j] + } + }) +} + +// denseFloatsLocal copies an array's elements into a flat float64 +// slice, widening int and float32 elements exactly. The walk splits +// across workers by rows from widenParallelMin up: every element +// widens exactly as the scalar read does and lands in its own slot, +// so the row streams are unchanged. +func denseFloatsLocal(a *Array, rows, cols int) []float64 { + out := make([]float64, rows*cols) + if rows*cols < widenParallelMin { + for i := range out { + out[i] = a.FloatAt(i) + } + return out + } + engine.Parallel(rows, func(rs, re int) { + for r := rs; r < re; r++ { + row := out[r*cols : (r+1)*cols] + for c := range row { + row[c] = a.FloatAt(r*cols + c) + } + } + }) + return out +} diff --git a/internal/core/mat_kernel_bench_test.go b/internal/core/mat_kernel_bench_test.go new file mode 100644 index 0000000..1213d19 --- /dev/null +++ b/internal/core/mat_kernel_bench_test.go @@ -0,0 +1,159 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Kernel-level companions to the shape benchmarks in bench_test.go: +// they pin the panel-walk matrix product, the vector products, the +// tiled transpose and the column gather at fixed sizes, so a +// regression in the mat.go kernels is visible independently of the +// mixed workloads. The serial variants measure the inner loops alone; +// the parallel variants measure the worker split on top. + +// benchMatMulKernelF64 runs the float64 2-D×2-D product at n×n with the +// worker count pinned to workers (0 keeps the default policy). +func benchMatMulKernelF64(b *testing.B, n, workers int) { + b.Helper() + prev := SetNumCPU(workers) + defer SetNumCPU(prev) + a, _ := FromFloats(make([]float64, n*n), n, n) + c, _ := FromFloats(make([]float64, n*n), n, n) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelMatMulF64Serial512 isolates the wide-panel kernel: one +// worker, so the numbers are pure inner-loop throughput. +func BenchmarkKernelMatMulF64Serial512(b *testing.B) { + benchMatMulKernelF64(b, 512, 1) +} + +// BenchmarkKernelMatMulF64Parallel512 measures the narrow-panel kernel +// under the default worker split. +func BenchmarkKernelMatMulF64Parallel512(b *testing.B) { + benchMatMulKernelF64(b, 512, 0) +} + +// BenchmarkKernelMatMulF64Small128 watches the panel edges: 128 rows +// give every worker exactly one 4-row panel. +func BenchmarkKernelMatMulF64Small128(b *testing.B) { + benchMatMulKernelF64(b, 128, 0) +} + +// benchMatMulKernelF32 is the float32 twin of benchMatMulKernelF64. +func benchMatMulKernelF32(b *testing.B, n, workers int) { + b.Helper() + prev := SetNumCPU(workers) + defer SetNumCPU(prev) + af, _ := FromFloats(make([]float64, n*n), n, n) + bf, _ := FromFloats(make([]float64, n*n), n, n) + a, _ := Astype(af, Float32) + c, _ := Astype(bf, Float32) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelMatMulF32Serial512 isolates the float32 kernel's inner +// loops and the round-once narrowing. +func BenchmarkKernelMatMulF32Serial512(b *testing.B) { + benchMatMulKernelF32(b, 512, 1) +} + +// BenchmarkKernelMatMulF32Parallel512 measures the float32 kernel under +// the default worker split. +func BenchmarkKernelMatMulF32Parallel512(b *testing.B) { + benchMatMulKernelF32(b, 512, 0) +} + +// BenchmarkKernelMatVecSerial measures the 2-D×1-D row dots at 512×512 +// on one worker. +func BenchmarkKernelMatVecSerial(b *testing.B) { + prev := SetNumCPU(1) + defer SetNumCPU(prev) + m, _ := FromFloats(make([]float64, 512*512), 512, 512) + v, _ := FromFloats(make([]float64, 512), 512) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(m, v); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelMatVecParallel measures the same row dots with the +// worker split above matVecParallelMin. +func BenchmarkKernelMatVecParallel(b *testing.B) { + m, _ := FromFloats(make([]float64, 512*512), 512, 512) + v, _ := FromFloats(make([]float64, 512), 512) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(m, v); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelVecMatParallel measures the 1-D×2-D column walk with +// the worker split. +func BenchmarkKernelVecMatParallel(b *testing.B) { + v, _ := FromFloats(make([]float64, 512), 512) + m, _ := FromFloats(make([]float64, 512*512), 512, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(v, m); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelTransposeTiled256 is the shape bench_test.go uses; the +// tiled walk owns this size outright. +func BenchmarkKernelTransposeTiled256(b *testing.B) { + a, _ := FromFloats(make([]float64, 256*256), 256, 256) + b.ReportAllocs() + for b.Loop() { + Transpose(a) + } +} + +// BenchmarkKernelTransposeOdd171 watches the tile edges with a size +// that is not a multiple of the 32-wide tile. +func BenchmarkKernelTransposeOdd171(b *testing.B) { + a, _ := FromFloats(make([]float64, 171*171), 171, 171) + b.ReportAllocs() + for b.Loop() { + Transpose(a) + } +} + +// BenchmarkKernelTransposeRank3 keeps the odometer path pinned: rank +// three never takes the tiled walk. +func BenchmarkKernelTransposeRank3(b *testing.B) { + a, _ := FromFloats(make([]float64, 16*16*16), 16, 16, 16) + b.ReportAllocs() + for b.Loop() { + Transpose(a) + } +} + +// BenchmarkKernelColParallelGather pins the parallel column gather at a +// row count above colParallelMin. +func BenchmarkKernelColParallelGather(b *testing.B) { + a, _ := FromFloats(make([]float64, 8192*16), 8192, 16) + b.ReportAllocs() + for b.Loop() { + if _, err := Col(a, 3); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/mat_kernel_test.go b/internal/core/mat_kernel_test.go new file mode 100644 index 0000000..81c2200 --- /dev/null +++ b/internal/core/mat_kernel_test.go @@ -0,0 +1,192 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// The products below are pinned against the scalar sum computed in this +// file, not against another kernel: every output element receives its +// addends over p ascending whatever panel, unroll or column band carries +// it, so each walk must reproduce that order exactly. The samples are +// integer-valued, so the comparison is exact. + +// TestComplexMatMulPanelWalks pins the complex product's four-row panel +// and its single-row tail: five rows make the panel run once and the +// tail once. +func TestComplexMatMulPanelWalks(t *testing.T) { + a := mustFromComplexes(t, []complex128{ + 1, -2, + 3, 4, + -5, 6, + 7, -8, + 9, 10, + }, 5, 2) + b := mustFromComplexes(t, []complex128{ + 1, 0, 2, + 0, 1, -3, + }, 2, 3) + + got, err := MatMul2D(a, b) + if err != nil { + t.Fatalf("MatMul complex 5x2 by 2x3: %v", err) + } + if sh := got.Shape(); len(sh) != 2 || sh[0] != 5 || sh[1] != 3 { + t.Fatalf("MatMul complex shape: %v, want [5 3]", sh) + } + want := naiveMatMulCells(a.RawComplexes(), b.RawComplexes(), 5, 2, 3) + for i := range want { + if got.RawComplexes()[i] != want[i] { + t.Fatalf("MatMul complex cell %d = %v, want %v", + i, got.RawComplexes()[i], want[i]) + } + } + + // Four rows exactly: the panel with no remainder row. + c := mustFromComplexes(t, []complex128{ + 2, 1, + 0, -3, + 4, 5, + 6, -7, + }, 4, 2) + d := mustFromComplexes(t, []complex128{ + 1, 2, + 3, -1, + }, 2, 2) + got4, err := MatMul2D(c, d) + if err != nil { + t.Fatalf("MatMul complex 4x2 by 2x2: %v", err) + } + want4 := naiveMatMulCells(c.RawComplexes(), d.RawComplexes(), 4, 2, 2) + for i := range want4 { + if got4.RawComplexes()[i] != want4[i] { + t.Fatalf("MatMul complex panel cell %d = %v, want %v", + i, got4.RawComplexes()[i], want4[i]) + } + } +} + +// TestComplexMatVecRows pins the two-row complex unroll of the +// matrix-vector product: five rows make the unroll run twice and the +// single-row tail once. +func TestComplexMatVecRows(t *testing.T) { + m := mustFromComplexes(t, []complex128{ + 1, 2, -3, + 4, -5, 6, + 7, 8, 9, + -1, 2, 3, + 4, 5, -6, + }, 5, 3) + v := mustFromComplexes(t, []complex128{2, -1, 3}, 3) + + got, err := MatMul2D(m, v) + if err != nil { + t.Fatalf("MatMul complex 5x3 by 3: %v", err) + } + if sh := got.Shape(); len(sh) != 1 || sh[0] != 5 { + t.Fatalf("MatMul complex matrix-vector shape: %v, want [5]", sh) + } + want := naiveMatVec(m.RawComplexes(), v.RawComplexes(), 5, 3) + for i := range want { + if got.RawComplexes()[i] != want[i] { + t.Fatalf("MatMul complex matrix-vector row %d = %v, want %v", + i, got.RawComplexes()[i], want[i]) + } + } +} + +// TestVecMatColumnWalks pins the vector-matrix product's column walks. +// The float64 payload accumulates four columns at a time in registers +// while the band is no wider than vecMatColBandMax and falls back to the +// accumulator slots above it; the complex payload always walks the +// slots. Every column sums its addends over p ascending. +func TestVecMatColumnWalks(t *testing.T) { + const k = 3 + v := mustFromFloats(t, []float64{2, -3, 5}, k) + // Four and six columns take the register walk (six with a + // remainder), nine columns the slot walk. + for _, cols := range []int{4, 6, 9} { + vals := make([]float64, k*cols) + for i := range vals { + vals[i] = float64(i%7) - 3 + 0.5 + } + m := mustFromFloats(t, vals, k, cols) + got, err := MatMul2D(v, m) + if err != nil { + t.Fatalf("MatMul vector by %dx%d: %v", k, cols, err) + } + if sh := got.Shape(); len(sh) != 1 || sh[0] != cols { + t.Fatalf("MatMul vector by %dx%d shape: %v, want [%d]", k, cols, sh, cols) + } + want := naiveVecMat(v.RawFloats(), m.RawFloats(), k, cols) + for j := range want { + if got.RawFloats()[j] != want[j] { + t.Fatalf("MatMul vector by %dx%d column %d = %v, want %v", + k, cols, j, got.RawFloats()[j], want[j]) + } + } + } + + // The complex payload takes the slot walk for every band. + cv := mustFromComplexes(t, []complex128{2, -1, 3}, k) + cvals := make([]complex128, k*4) + for i := range cvals { + cvals[i] = complex(float64(i%5)-2, float64(i%3)-1) + } + cm := mustFromComplexes(t, cvals, k, 4) + cgot, err := MatMul2D(cv, cm) + if err != nil { + t.Fatalf("MatMul complex vector by %dx4: %v", k, err) + } + cwant := naiveVecMat(cv.RawComplexes(), cm.RawComplexes(), k, 4) + for j := range cwant { + if cgot.RawComplexes()[j] != cwant[j] { + t.Fatalf("MatMul complex vector-matrix column %d = %v, want %v", + j, cgot.RawComplexes()[j], cwant[j]) + } + } +} + +// naiveMatMulCells returns the n×m product of an n×k matrix and a k×m +// matrix, each cell summed over p ascending. +func naiveMatMulCells[T complex128 | float64](a, b []T, n, k, m int) []T { + out := make([]T, n*m) + for i := range n { + for j := range m { + var s T + for p := range k { + s += a[i*k+p] * b[p*m+j] + } + out[i*m+j] = s + } + } + return out +} + +// naiveMatVec returns the n-vector an n×k matrix multiplies a k-vector +// into, each row summed over p ascending. +func naiveMatVec[T complex128 | float64](a, v []T, n, k int) []T { + out := make([]T, n) + for i := range n { + var s T + for p := range k { + s += a[i*k+p] * v[p] + } + out[i] = s + } + return out +} + +// naiveVecMat returns the m-vector a k-vector multiplies a k×m matrix +// into, each column summed over p ascending. +func naiveVecMat[T complex128 | float64](v, m []T, k, cols int) []T { + out := make([]T, cols) + for j := range cols { + var s T + for p := range k { + s += v[p] * m[p*cols+j] + } + out[j] = s + } + return out +} diff --git a/internal/core/mat_tile_bench_test.go b/internal/core/mat_tile_bench_test.go new file mode 100644 index 0000000..34daa33 --- /dev/null +++ b/internal/core/mat_tile_bench_test.go @@ -0,0 +1,101 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Shape companions to the matmul benchmarks in bench_test.go: they pin +// the large square product and the rectangular extremes the row-panel +// kernels serve, so a dispatch or kernel change in mat.go is judged on +// the shapes that stress different sides of the split. + +// benchMatMulSized runs the float64 n×n product with the worker count +// pinned to workers (0 keeps the default policy). +func benchMatMulSized(b *testing.B, n, workers int) { + b.Helper() + prev := SetNumCPU(workers) + defer SetNumCPU(prev) + a, _ := FromFloats(make([]float64, n*n), n, n) + c, _ := FromFloats(make([]float64, n*n), n, n) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMatMul1024 is the large square product, one size above +// BenchmarkMatMulParallel: its b payload is four times the L3 slice +// each CCX serves, so it watches the row split's shared reads of b. +func BenchmarkMatMul1024(b *testing.B) { benchMatMulSized(b, 1024, 0) } + +// benchMatMulRect runs rectangular n×k·k×m float64 products under the +// default worker policy. The four extremes pin the panel edges: tall +// and wide shapes with a thin k, and squat and skinny shapes whose b +// or a payload dwarfs the other operand. +func benchMatMulRect(b *testing.B, n, k, m int) { + b.Helper() + a, _ := FromFloats(make([]float64, n*k), n, k) + c, _ := FromFloats(make([]float64, k*m), k, m) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMatMulTall512x64x2048(b *testing.B) { benchMatMulRect(b, 512, 64, 2048) } +func BenchmarkMatMulWide2048x64x512(b *testing.B) { benchMatMulRect(b, 2048, 64, 512) } +func BenchmarkMatMulFat64x512x2048(b *testing.B) { benchMatMulRect(b, 64, 512, 2048) } +func BenchmarkMatMulSkinny2048x512x64(b *testing.B) { benchMatMulRect(b, 2048, 512, 64) } + +// BenchmarkMatMulOdd130 watches the remainder rows and columns: 130 is +// past the 4-row panel stride in both dimensions. +func BenchmarkMatMulOdd130(b *testing.B) { benchMatMulRect(b, 130, 130, 130) } + +// Small-matrix benchmarks guard the parallel spawn floors: the +// training-loop shapes pay goroutine start-up per call, so the floors +// must keep them on the calling goroutine without regressing the large +// splits. +func BenchmarkMatMulSmall32(b *testing.B) { benchMatMulSquare(b, 32) } + +func BenchmarkMatMulSmall64(b *testing.B) { benchMatMulSquare(b, 64) } + +func BenchmarkMatMulMid128(b *testing.B) { benchMatMulSquare(b, 128) } + +func BenchmarkMatMulF32Small64(b *testing.B) { benchMatMulSquareF32(b, 64) } + +func benchMatMulSquare(b *testing.B, n int) { + b.Helper() + a, _ := FromFloats(make([]float64, n*n), n, n) + for i := range a.RawFloats() { + a.RawFloats()[i] = float64(i%7) - 3 + } + c, _ := FromFloats(make([]float64, n*n), n, n) + for i := range c.RawFloats() { + c.RawFloats()[i] = float64(i%5) - 2 + } + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} + +func benchMatMulSquareF32(b *testing.B, n int) { + b.Helper() + af, _ := FromFloats(make([]float64, n*n), n, n) + a, _ := Astype(af, Float32) + bf, _ := FromFloats(make([]float64, n*n), n, n) + c, _ := Astype(bf, Float32) + b.ReportAllocs() + for b.Loop() { + if _, err := MatMul2D(a, c); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/math_bench_extra_test.go b/internal/core/math_bench_extra_test.go new file mode 100644 index 0000000..958428a --- /dev/null +++ b/internal/core/math_bench_extra_test.go @@ -0,0 +1,98 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Size-sweep guards for the element-wise dispatch policy. The parallelMin +// threshold sends workloads whose per-worker chunk falls below the floor +// to the calling goroutine; the small and medium sizes here must show +// that policy winning (no spawn cost), while 1M and 4M must keep scaling +// across workers. Run with -bench before and after any change to the +// threshold in runtime.go or the specialised kernels in mathfunc.go. + +func benchAddSize(b *testing.B, n int) { + a, _ := FromFloats(make([]float64, n), n) + c, _ := FromFloats(make([]float64, n), n) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkAdd1k(b *testing.B) { benchAddSize(b, 1_000) } + +func BenchmarkAdd10k(b *testing.B) { benchAddSize(b, 10_000) } + +func BenchmarkAdd100k(b *testing.B) { benchAddSize(b, 100_000) } + +func BenchmarkAdd1M(b *testing.B) { benchAddSize(b, 1_000_000) } + +func BenchmarkAdd4M(b *testing.B) { benchAddSize(b, 4_194_304) } + +// BenchmarkExp100k guards the transcendental family at the threshold +// boundary: realFunc shares the element-wise threshold, so this size +// records what the policy costs the expensive per-element maps. +func BenchmarkExp100k(b *testing.B) { + a, _ := FromFloats(make([]float64, 100_000), 100_000) + for i := range a.Len() { + a.RawFloats()[i] = float64(i%97) + 1 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Exp(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkExp1k guards the small end of the transcendental family: +// their per-element cost is high enough that even a thousand elements +// may still feed every worker, which the floor must not contradict. +func BenchmarkExp1k(b *testing.B) { + a, _ := FromFloats(make([]float64, 1_000), 1_000) + for i := range a.Len() { + a.RawFloats()[i] = float64(i%97) + 1 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Exp(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkAddComplex100k guards the complex128 path of the arithmetic +// maps: each element moves three times the bytes of a float64 add, so +// the serial crossovers of the two paths differ and the floor must +// serve both. +func BenchmarkAddComplex100k(b *testing.B) { + a, _ := FromComplexes(make([]complex128, 100_000), 100_000) + c, _ := FromComplexes(make([]complex128, 100_000), 100_000) + b.ReportAllocs() + for b.Loop() { + if _, err := Add(a, c); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSqrtFloat32 pins the float32 path of the Sqrt specialisation: +// one widening, one intrinsic square root and one narrowing per element, +// never a call through a func value. +func BenchmarkSqrtFloat32(b *testing.B) { + a, _ := FromFloats(make([]float64, 1<<20), 1<<20) + for i := range a.Len() { + a.RawFloats()[i] = float64(i%97) + 1 + } + f32, _ := Astype(a, Float32) + b.ReportAllocs() + for b.Loop() { + if _, err := Sqrt(f32); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/mathfunc.go b/internal/core/mathfunc.go new file mode 100644 index 0000000..bbd64c9 --- /dev/null +++ b/internal/core/mathfunc.go @@ -0,0 +1,427 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "math" + +// Element-wise math functions. The transcendental family is +// real-only for now: int arrays convert to float, complex arrays error +// loudly. IEEE domain rules apply: log and sqrt of negatives yield NaN, +// never an error. Rounding keeps each dtype: int arrays are identity +// copies, float arrays map to float. + +// Abs returns the element-wise absolute value; a complex array yields its +// float magnitudes. Bool and the narrow integer widths are refused: they +// carry no kernel here, and this no-error entry answers nil, +// the refusal shape the no-error constructors carry; the caller converts +// with Astype first. +func Abs(a *Array) *Array { + if narrowRefused(a.dt) { + return nil + } + out := &Array{shape: a.Shape()} + switch a.dt { + case Int: + out.dt = Int + out.ints = make([]int64, a.Len()) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.ints[s:e], out.ints[s:e] + for i := range os { + v := as[i] + if v < 0 { + v = -v // wraps at the minimum value + } + os[i] = v + } + }) + case Float16: + out.dt = Float16 + out.halves = make([]uint16, a.Len()) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + // Abs of a half is exact: widen, clear the sign in + // float64, narrow. + os[i] = HalfFromFloat64(math.Abs(HalfToFloat64(as[i]))) + } + }) + case Float32: + out.dt = Float32 + out.floats32 = make([]float32, a.Len()) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = float32(math.Abs(float64(as[i]))) + } + }) + case Float: + out.dt = Float + out.floats = make([]float64, a.Len()) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats[s:e], out.floats[s:e] + for i := range os { + os[i] = math.Abs(as[i]) + } + }) + default: + // The magnitude of a complex element costs a math.Hypot, so this + // walk amortises a spawn at the math-floor chunk, not the + // arithmetic floor. + out.dt = Float + out.floats = make([]float64, a.Len()) + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.complexes[s:e], out.floats[s:e] + for i := range os { + os[i] = math.Hypot(real(as[i]), imag(as[i])) + } + }) + } + return out +} + +// Exp returns e raised to each element. +func Exp(a *Array) (*Array, error) { return a.realFunc("Exp", math.Exp) } + +// Log returns the natural logarithm of each element; negatives yield NaN. +func Log(a *Array) (*Array, error) { return a.realFunc("Log", math.Log) } + +// Log2 returns the base-2 logarithm of each element; negatives yield NaN. +func Log2(a *Array) (*Array, error) { return a.realFunc("Log2", math.Log2) } + +// Log10 returns the base-10 logarithm of each element; negatives yield +// NaN. +func Log10(a *Array) (*Array, error) { return a.realFunc("Log10", math.Log10) } + +// Sqrt returns the square root of each element; negatives yield NaN. +// Specialised per dtype: math.Sqrt called directly lowers to the +// SQRTSD intrinsic with no call, while a loop through realFunc's func +// value pays an indirect call per element that dwarfs the operation +// itself. The per-dtype loops read through realFunc's accessors +// element for element, so the outputs, the error and the dtype +// promotion are identical. +func Sqrt(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("Sqrt: complex arrays are not supported") + } + if narrowRefused(a.dt) { + return nil, errf("Sqrt: dtype %s is not supported; convert with Astype", a.dt) + } + out := &Array{shape: a.Shape(), dt: a.dt} + if a.dt == Int { + out.dt = Float + } + out.alloc(a.Len()) + switch a.dt { + case Int: + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.floats[i] = math.Sqrt(a.floatAt(i)) + } + }) + case Float16: + // float16 keeps float16, computed in float64 and narrowed once: + // the widening is exact, so the narrowed result is the correctly + // rounded square root. + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.halves[i] = HalfFromFloat64(math.Sqrt(HalfToFloat64(a.halves[i]))) + } + }) + case Float32: + // float32 keeps float32, computed in float64 and rounded once + //: the widening is exact, so the narrowed result is the + // correctly rounded square root. + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.floats32[i] = float32(math.Sqrt(float64(a.floats32[i]))) + } + }) + default: + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + os := out.floats[s:e] + as := a.floats[s:e] + for i := range os { + os[i] = math.Sqrt(as[i]) + } + }) + } + return out, nil +} + +// Sin returns the sine of each element (radians). +func Sin(a *Array) (*Array, error) { return a.realFunc("Sin", math.Sin) } + +// Cos returns the cosine of each element (radians). +func Cos(a *Array) (*Array, error) { return a.realFunc("Cos", math.Cos) } + +// Tan returns the tangent of each element (radians). +func Tan(a *Array) (*Array, error) { return a.realFunc("Tan", math.Tan) } + +// Floor rounds each element down; an int array is an identity copy. +func Floor(a *Array) (*Array, error) { return a.roundFunc("Floor", math.Floor) } + +// Ceil rounds each element up; an int array is an identity copy. +func Ceil(a *Array) (*Array, error) { return a.roundFunc("Ceil", math.Ceil) } + +// Round rounds each element half away from zero; an int array is an +// identity copy. +func Round(a *Array) (*Array, error) { return a.roundFunc("Round", math.Round) } + +// Trunc rounds each element toward zero; an int array is an identity copy. +func Trunc(a *Array) (*Array, error) { return a.roundFunc("Trunc", math.Trunc) } + +// Pow returns the element-wise power; integer-class operands promote +// to an integer dtype and stay integer there (wrapping on overflow, +// and a negative exponent is refused at every width: the reciprocal +// has no integer result), any float or complex operand promotes per +// the ladder (int operands convert), and negative bases with +// fractional exponents follow math.Pow's IEEE NaN. +func Pow(a, b *Array) (*Array, error) { + if intClass(a.dt) && intClass(b.dt) { + // An integer-class pair computes in integer space whatever + // promote answers, and a negative exponent has no integer + // result at any width. The refusal scans the exponent operand + // through intAt, whose widening of the whole class is exact, + // so no mixed pair (an Int base with a narrow exponent, a + // Uint32 base with a negative narrow exponent) slips past the + // gate into powInt, whose loop answers 1 for a negative + // exponent. + for i := range b.Len() { + if e := b.intAt(i); e < 0 { + return nil, errf("Pow: negative exponent %d at element %d has no int result", e, i) + } + } + } + return elementwiseOp(a, b, "Pow", pairPow) +} + +// mathPowComplex raises a complex base to a complex exponent via +// exp(y·log(x)): the principal branch. A zero base with a non-positive +// exponent follows IEEE NaN/Inf conventions. +func mathPowComplex(x, y complex128) complex128 { + if x == 0 { + if imag(y) == 0 && real(y) > 0 { + return 0 + } + return complex(math.NaN(), math.NaN()) + } + return cexp(cmul(y, clog(x))) +} + +// cexp, clog and cmul are the complex transcendals needed by Pow, +// implemented here to keep the library dependency-free. +func cexp(z complex128) complex128 { + e := math.Exp(real(z)) + return complex(e*math.Cos(imag(z)), e*math.Sin(imag(z))) +} + +func clog(z complex128) complex128 { + return complex(math.Log(math.Hypot(real(z), imag(z))), math.Atan2(imag(z), real(z))) +} + +func cmul(a, b complex128) complex128 { + return complex(real(a)*real(b)-imag(a)*imag(b), real(a)*imag(b)+imag(a)*real(b)) +} + +// PowI raises each element to an int exponent; int arrays stay int with +// wrapping on overflow (a negative exponent is an error), float32 arrays +// keep float32 computed in float64, float64 arrays stay float64, +// complex arrays use exact repeated squaring (a negative exponent takes +// the reciprocal). +func PowI(a *Array, n int64) (*Array, error) { + switch a.dt { + case Int: + if n < 0 { + return nil, errf("PowI: a negative exponent has no int result") + } + out := &Array{shape: a.Shape(), dt: Int, ints: make([]int64, a.Len())} + // Bounded by Len, not by the payload: a rebased view carries a + // payload longer than its extent. + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.ints[s:e], out.ints[s:e] + for i := range os { + os[i] = powInt(as[i], n) + } + }) + return out, nil + case Float16: + out := &Array{shape: a.Shape(), dt: Float16, halves: make([]uint16, a.Len())} + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(math.Pow(HalfToFloat64(as[i]), float64(n))) + } + }) + return out, nil + case Float32: + out := &Array{shape: a.Shape(), dt: Float32, floats32: make([]float32, a.Len())} + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = float32(math.Pow(float64(as[i]), float64(n))) + } + }) + return out, nil + case Float: + return a.realFunc("PowI", func(x float64) float64 { return math.Pow(x, float64(n)) }) + case Complex: + if n < 0 { + // A negative power is the reciprocal of the positive one. + out := &Array{shape: a.Shape(), dt: Complex, complexes: make([]complex128, a.Len())} + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.complexes[s:e], out.complexes[s:e] + for i := range os { + p := powComplexI(as[i], -n) + os[i] = complex(1, 0) / p + } + }) + return out, nil + } + out := &Array{shape: a.Shape(), dt: Complex, complexes: make([]complex128, a.Len())} + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.complexes[s:e], out.complexes[s:e] + for i := range os { + os[i] = powComplexI(as[i], n) + } + }) + return out, nil + } + return nil, errf("PowI: dtype %s is not supported; convert with Astype", a.dt) +} + +// powComplexI raises z to a non-negative integer power by repeated +// squaring: exact integer arithmetic in complex128. +func powComplexI(z complex128, n int64) complex128 { + result := complex(1, 0) + base := z + for n > 0 { + if n&1 == 1 { + result *= base + } + base *= base + n >>= 1 + } + return result +} + +// powInt raises x to a non-negative y by repeated squaring, wrapping on +// overflow like every int operation. +func powInt(x, y int64) int64 { + result := int64(1) + base := x + for y > 0 { + if y&1 == 1 { + result *= base + } + base *= base + y >>= 1 + } + return result +} + +// realFunc applies a real function element-wise: int arrays convert to +// float64, float16 and float32 arrays keep their dtype (computed in +// float64, narrowed once), complex arrays error. +func (a *Array) realFunc(name string, f func(float64) float64) (*Array, error) { + if a.dt == Complex { + return nil, errf("%s: complex arrays are not supported", name) + } + if narrowRefused(a.dt) { + // The transcendental family widens integer arrays to float64; + // bool and the narrow widths carry no kernel here. + return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt) + } + out := &Array{shape: a.Shape(), dt: a.dt} + if a.dt == Int { + out.dt = Float + } + out.alloc(a.Len()) + if out.dt == Float16 { + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(f(HalfToFloat64(as[i]))) + } + }) + return out, nil + } + if out.dt == Float32 { + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = float32(f(float64(as[i]))) + } + }) + return out, nil + } + // A float64 payload reads directly; int keeps the widening accessor. + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + os := out.floats[s:e] + if a.dt == Float { + as := a.floats[s:e] + for i := range os { + os[i] = f(as[i]) + } + return + } + for i := s; i < e; i++ { + os[i-s] = f(a.floatAt(i)) + } + }) + return out, nil +} + +// roundFunc applies a rounding function element-wise: float arrays map to +// their own width, int arrays are identity copies, complex arrays error. +func (a *Array) roundFunc(name string, f func(float64) float64) (*Array, error) { + if a.dt == Complex { + return nil, errf("%s: complex arrays have no rounding", name) + } + if narrowRefused(a.dt) { + // Integer arrays are identity copies under this family; bool and + // the narrow widths carry no kernel here. + return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt) + } + if a.dt == Int { + ints, _, _, _, _ := a.cloneData() + return &Array{shape: a.Shape(), dt: Int, ints: ints}, nil + } + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(a.Len()) + if a.dt == Float16 { + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(f(HalfToFloat64(as[i]))) + } + }) + return out, nil + } + if a.dt == Float32 { + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = float32(f(float64(as[i]))) + } + }) + return out, nil + } + parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { + as, os := a.floats[s:e], out.floats[s:e] + for i := range os { + os[i] = f(as[i]) + } + }) + return out, nil +} + +// Tanh returns the hyperbolic tangent of each element. +func Tanh(a *Array) (*Array, error) { return a.realFunc("Tanh", math.Tanh) } + +// Sigmoid returns the logistic function 1/(1+e^−x) of each element. +func Sigmoid(a *Array) (*Array, error) { + return a.realFunc("Sigmoid", func(x float64) float64 { + return 1 / (1 + math.Exp(-x)) + }) +} diff --git a/internal/core/mathfunc_test.go b/internal/core/mathfunc_test.go new file mode 100644 index 0000000..1a86cea --- /dev/null +++ b/internal/core/mathfunc_test.go @@ -0,0 +1,259 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +func TestMathFuncs(t *testing.T) { + f := mustFromFloats(t, []float64{1.0, 4.0}, 2) + + sqrt := mustOk(Sqrt(f)) + if v, _ := FloatAt(sqrt, 0); v != 1.0 { + t.Fatalf("Sqrt: %v", v) + } + log2 := mustOk(Log2(f)) + if v, _ := FloatAt(log2, 1); v != 2.0 { + t.Fatalf("Log2: %v", v) + } + exp := mustOk(Exp(mustFromFloats(t, []float64{0}, 1))) + if v, _ := FloatAt(exp, 0); math.Abs(v-1) > 1e-12 { + t.Fatalf("Exp(0) must be 1: %v", v) + } + log := mustOk(Log(mustFromFloats(t, []float64{1}, 1))) + if v, _ := FloatAt(log, 0); v != 0 { + t.Fatalf("Log(1): %v", v) + } + log10 := mustOk(Log10(mustFromFloats(t, []float64{10}, 1))) + if v, _ := FloatAt(log10, 0); v != 1 { + t.Fatalf("Log10(10): %v", v) + } + sin := mustOk(Sin(mustFromFloats(t, []float64{0}, 1))) + if v, _ := FloatAt(sin, 0); v != 0 { + t.Fatalf("Sin(0): %v", v) + } + cos := mustOk(Cos(mustFromFloats(t, []float64{0}, 1))) + if v, _ := FloatAt(cos, 0); v != 1 { + t.Fatalf("Cos(0): %v", v) + } + tan := mustOk(Tan(mustFromFloats(t, []float64{0}, 1))) + if v, _ := FloatAt(tan, 0); v != 0 { + t.Fatalf("Tan(0): %v", v) + } + + // IEEE domain rules: negatives yield NaN, never an error. + neg := mustFromFloats(t, []float64{-1}, 1) + sq, err := Sqrt(neg) + if err != nil { + t.Fatalf("Sqrt(-1): %v", err) + } + if v, _ := FloatAt(sq, 0); !math.IsNaN(v) { + t.Fatalf("Sqrt(-1): %v", v) + } + lg, _ := Log(neg) + if v, _ := FloatAt(lg, 0); !math.IsNaN(v) { + t.Fatalf("Log(-1): %v", v) + } + + // int arrays convert for the transcendental family. + i := mustFromInts(t, []int64{4}, 1) + isqrtArr := mustOk(Sqrt(i)) + isqrt, _ := FloatAt(isqrtArr, 0) + if isqrt != 2 || isqrtArr.Dtype() != Float { + t.Fatalf("Sqrt on int: %s", isqrtArr) + } + + // complex arrays error for the real family. + c := mustFromComplexes(t, []complex128{1}, 1) + if _, err := Exp(c); err == nil || !strings.Contains(err.Error(), "not supported") { + t.Fatalf("Exp complex: %v", err) + } + if _, err := Sin(c); err == nil { + t.Fatalf("Sin complex must error") + } + if _, err := Floor(c); err == nil || !strings.Contains(err.Error(), "no rounding") { + t.Fatalf("Floor complex: %v", err) + } +} + +func TestAbs(t *testing.T) { + i := mustFromInts(t, []int64{-5, 3}, 2) + ai := Abs(i) + if ai.Dtype() != Int { + t.Fatalf("Abs int dtype: %s", ai.Dtype()) + } + if v, _ := IntAt(ai, 0); v != 5 { + t.Fatalf("Abs int: %d", v) + } + + f := mustFromFloats(t, []float64{-2.5}, 1) + if v, _ := FloatAt(Abs(f), 0); v != 2.5 { + t.Fatalf("Abs float: %v", v) + } + + // Complex magnitude is real. + c := mustFromComplexes(t, []complex128{complex(3, 4)}, 1) + ac := Abs(c) + if ac.Dtype() != Float { + t.Fatalf("Abs complex dtype: %s", ac.Dtype()) + } + if v, _ := FloatAt(ac, 0); v != 5 { + t.Fatalf("Abs complex: %v", v) + } +} + +func TestRounding(t *testing.T) { + f := mustFromFloats(t, []float64{2.7, -2.7}, 2) + floor := mustOk(Floor(f)) + if v, _ := FloatAt(floor, 0); v != 2 { + t.Fatalf("Floor: %v", v) + } + if v, _ := FloatAt(floor, 1); v != -3 { + t.Fatalf("Floor neg: %v", v) + } + ceil := mustOk(Ceil(f)) + if v, _ := FloatAt(ceil, 0); v != 3 { + t.Fatalf("Ceil: %v", v) + } + round := mustOk(Round(f)) + if v, _ := FloatAt(round, 0); v != 3 { + t.Fatalf("Round: %v", v) + } + if v, _ := FloatAt(round, 1); v != -3 { + t.Fatalf("Round half away: %v", v) + } + trunc := mustOk(Trunc(f)) + if v, _ := FloatAt(trunc, 0); v != 2 { + t.Fatalf("Trunc: %v", v) + } + + // Int arrays are identity copies. + i := mustFromInts(t, []int64{7}, 1) + ir, err := Floor(i) + if err != nil || ir.Dtype() != Int { + t.Fatalf("Floor on int: %s %v", ir, err) + } + if v, _ := IntAt(ir, 0); v != 7 { + t.Fatalf("Floor int identity: %d", v) + } +} + +func TestPow(t *testing.T) { + a := mustFromInts(t, []int64{2, 3}, 2) + e := mustFromInts(t, []int64{10, 2}, 2) + + p, err := Pow(a, e) + if err != nil { + t.Fatalf("Pow: %v", err) + } + if p.Dtype() != Int { + t.Fatalf("Pow dtype: %s", p.Dtype()) + } + if v, _ := IntAt(p, 0); v != 1024 { + t.Fatalf("Pow 2^10: %d", v) + } + if v, _ := IntAt(p, 1); v != 9 { + t.Fatalf("Pow 3^2: %d", v) + } + + // Float promotion. + pf, err := Pow(mustFromFloats(t, []float64{4.0, 9.0}, 2), e) + if err != nil { + t.Fatalf("Pow promote: %v", err) + } + if pf.Dtype() != Float { + t.Fatalf("Pow promote dtype: %s", pf.Dtype()) + } + if v, _ := FloatAt(pf, 0); v != 1048576 { + t.Fatalf("Pow 4^10: %v", v) + } + if v, _ := FloatAt(pf, 1); v != 81 { + t.Fatalf("Pow 9^2: %v", v) + } + + // Scalar exponent. + pi, err := PowI(a, 3) + if err != nil { + t.Fatalf("PowI: %v", err) + } + if v, _ := IntAt(pi, 0); v != 8 { + t.Fatalf("PowI 2^3: %d", v) + } + pif, err := PowI(mustFromFloats(t, []float64{9.0}, 1), 2) + if err != nil { + t.Fatalf("PowI float: %v", err) + } + if v, _ := FloatAt(pif, 0); v != 81 { + t.Fatalf("PowI float value: %v", v) + } + + // Negative int exponent on an int array errors; on float it works. + if _, err := PowI(a, -1); err == nil || !strings.Contains(err.Error(), "negative exponent") { + t.Fatalf("PowI negative: %v", err) + } + if _, err := PowI(mustFromFloats(t, []float64{2}, 1), -2); err != nil { + t.Fatalf("PowI negative on float: %v", err) + } + // PowI carries no narrow kernel: the refusal names + // the dtype and the conversion the contract asks for. + if _, err := PowI(narrowInt8s(t, []int8{2, 3}, 2), 2); err == nil || + !strings.Contains(err.Error(), "convert with Astype") || + !strings.Contains(err.Error(), "int8") { + t.Fatalf("PowI int8: %v", err) + } + + // Complex powers use exact repeated squaring on complex128. + c := mustFromComplexes(t, []complex128{2 + 3i}, 1) + if _, err := Pow(c, c); err != nil { + t.Fatalf("Pow complex: %v", err) + } + squared := mustOk(PowI(c, 2)) + want := (2 + 3i) * (2 + 3i) // −5+12i + if squared.RawComplexes()[0] != want { + t.Fatalf("PowI complex = %v, want %v", squared.RawComplexes()[0], want) + } +} + +// mustOk unwraps a call that must succeed; a panic here fails the test. +func mustOk(a *Array, err error) *Array { + if err != nil { + panic(err) + } + return a +} + +// TestPowDenseFloatOperands pins the dense float base against a dense +// float exponent, the pair whose kernel holds both payloads as slices; +// an int exponent and a promoted pair take the closure walk instead. +func TestPowDenseFloatOperands(t *testing.T) { + base := mustFromFloats(t, []float64{4, 9, 2, 2}, 2, 2) + exp := mustFromFloats(t, []float64{10, 2, 3, -1}, 2, 2) + got, err := Pow(base, exp) + if err != nil { + t.Fatalf("Pow float: %v", err) + } + if got.Dtype() != Float { + t.Fatalf("Pow float dtype: %s", got.Dtype()) + } + for i, want := range []float64{1048576, 81, 8, 0.5} { + if v := got.RawFloats()[i]; v != want { + t.Fatalf("Pow float [%d] = %v, want %v", i, v, want) + } + } + // The float32 pair takes the same kernel one width down. + b32 := mustFromFloat32s(t, []float32{4, 2}, 2) + e32 := mustFromFloat32s(t, []float32{2, -1}, 2) + p32, err := Pow(b32, e32) + if err != nil { + t.Fatalf("Pow float32: %v", err) + } + for i, want := range []float32{16, 0.5} { + if v := p32.RawFloat32s()[i]; v != want { + t.Fatalf("Pow float32 [%d] = %v, want %v", i, v, want) + } + } +} diff --git a/internal/core/matrix2.go b/internal/core/matrix2.go new file mode 100644 index 0000000..b9ae12b --- /dev/null +++ b/internal/core/matrix2.go @@ -0,0 +1,171 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +// Matrix utilities: Trace, Diagonal, Kron. These are the +// smaller linear-algebra helpers that sit alongside MatMul / Solve / +// Inv. Trace reduces a 2-D matrix along the main diagonal; Diagonal +// extracts one diagonal as a 1-D array (offset > 0 picks a super- +// diagonal, offset < 0 picks a sub-diagonal); Kron computes the +// Kronecker product of two 2-D matrices, with a result whose shape is +// the outer product of the inputs' shapes. + +// Trace returns the sum along the main diagonal of a 2-D square +// matrix. The result is always float64: int arrays promote, complex +// matrices are answered by TraceComplex. The narrow integer widths and +// bool accumulate in int64 from exact widenings before the single +// widening into the float64 result, the exactness rule the scalar +// reductions carry. +func Trace(a *Array) (float64, error) { + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return 0, errf("Trace: needs a square 2-D matrix, got shape %s", shapeText(a.Shape())) + } + if a.Dtype() == Complex { + return 0, errf("Trace: complex matrices are answered by TraceComplex") + } + n := a.Shape()[0] + if narrowRefused(a.Dtype()) { + var is int64 + for i := range n { + is += a.intAt(i*n + i) + } + return float64(is), nil + } + var s float64 + for i := range n { + s += a.FloatAt(i*n + i) + } + return s, nil +} + +// TraceComplex returns the main-diagonal sum of a complex square matrix. +func TraceComplex(a *Array) (complex128, error) { + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return 0, errf("TraceComplex: needs a square 2-D matrix, got shape %s", shapeText(a.Shape())) + } + if a.Dtype() != Complex { + return 0, errf("TraceComplex: needs a complex matrix") + } + var s complex128 + n := a.Shape()[0] + for i := range n { + s += a.ComplexAt(i*n + i) + } + return s, nil +} + +// Diagonal returns the elements on the offset-th diagonal of a 2-D +// matrix as a 1-D array. offset = 0 is the main diagonal, offset > 0 a +// super-diagonal, offset < 0 a sub-diagonal. The result's dtype follows +// the input. +func Diagonal(a *Array, offset int) (*Array, error) { + if a.NDim() != 2 { + return nil, errf("Diagonal: needs a 2-D array, got shape %s", shapeText(a.Shape())) + } + rows, cols := a.Shape()[0], a.Shape()[1] + var startRow, startCol int + var length int + switch { + case offset >= 0: + startRow = 0 + startCol = offset + length = min(rows, cols-offset) + default: + startRow = -offset + startCol = 0 + length = min(rows+offset, cols) + } + if length <= 0 { + return nil, errf("Diagonal: offset %d is out of range for shape %s", offset, shapeText(a.Shape())) + } + out := &Array{shape: []int{length}, dt: a.Dtype()} + out.alloc(length) + for i := range length { + out.setFrom(i, a, (startRow+i)*cols+(startCol+i)) + } + return out, nil +} + +// Kron returns the Kronecker product of two 2-D matrices. For shapes +// (m, n) and (p, q) the result has shape (m*p, n*q) and element +// (i*p+k, j*q+l) = a[i,j] * b[k,l]. The result dtype follows the +// promotion ladder: int with int stays int, any float or complex +// operand promotes. +func Kron(a, b *Array) (*Array, error) { + if a.NDim() != 2 || b.NDim() != 2 { + return nil, errf("Kron: both inputs must be 2-D, got shapes %s and %s", + shapeText(a.shape), shapeText(b.shape)) + } + for _, op := range []*Array{a, b} { + if narrowRefused(op.Dtype()) { + return nil, errf("Kron: dtype %s is not supported; convert with Astype", op.Dtype()) + } + } + ma, na := a.Shape()[0], a.Shape()[1] + mb, nb := b.Shape()[0], b.Shape()[1] + dt := promote(a.Dtype(), b.Dtype()) + // The outer-product shape must go through the same validation as + // every other shape: unchecked products wrap on oversized inputs. + total, outShape, terr := checkedDims([]int{ma * mb, na * nb}) + if terr != nil { + return nil, base.WrapErr("Kron", terr) + } + out := &Array{shape: outShape, dt: dt} + out.alloc(total) + if dt == Complex { + for i := range ma { + for j := range na { + for k := range mb { + for l := range nb { + out.complexes[((i*mb+k)*(na*nb))+j*nb+l] = + a.ComplexAt(i*na+j) * b.ComplexAt(k*nb+l) + } + } + } + } + return out, nil + } + if dt == Int { + // int x int stays exact: products must not round-trip through + // float64, which loses low bits above 2^53. + for i := range ma { + for j := range na { + for k := range mb { + for l := range nb { + outRow := (i*mb + k) * (na * nb) + outCol := j*nb + l + out.ints[outRow+outCol] = a.ints[i*na+j] * b.ints[k*nb+l] + } + } + } + } + return out, nil + } + for i := range ma { + for j := range na { + av := a.FloatAt(i*na + j) + for k := range mb { + for l := range nb { + // Both operands widen to float64, the product + // narrows once per element, so an int + // operand loses at most one rounding, like every + // sibling kernel. + prod := av * b.FloatAt(k*nb+l) + pos := (i*mb+k)*(na*nb) + j*nb + l + switch out.dt { + case Float16: + out.halves[pos] = HalfFromFloat64(prod) + case Float32: + out.floats32[pos] = float32(prod) + default: + out.floats[pos] = prod + } + } + } + } + } + return out, nil +} diff --git a/internal/core/misc_bench_test.go b/internal/core/misc_bench_test.go new file mode 100644 index 0000000..124ed67 --- /dev/null +++ b/internal/core/misc_bench_test.go @@ -0,0 +1,80 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Miscellaneous bulk-path benchmarks: the generator fills, the sort +// family and the bilinear 2-D interpolation, at sizes where their +// inner loops dominate. + +func BenchmarkGeneratorFloats1M(b *testing.B) { + g := NewGenerator(42) + b.ReportAllocs() + for b.Loop() { + if _, err := Floats(g, 1<<20); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkGeneratorNormal1M(b *testing.B) { + g := NewGenerator(42) + b.ReportAllocs() + for b.Loop() { + if _, err := Normal(g, 1<<20, 0, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSort1M(b *testing.B) { + a, _ := FromFloats(make([]float64, 1<<20), 1<<20) + g := NewGenerator(7) + src, _ := Floats(g, 1<<20) + copy(a.RawFloats(), src.RawFloats()) + b.ResetTimer() + b.ReportAllocs() + for b.Loop() { + if _, err := Sort(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkArgSort1M(b *testing.B) { + a, _ := FromFloats(make([]float64, 1<<20), 1<<20) + g := NewGenerator(7) + src, _ := Floats(g, 1<<20) + copy(a.RawFloats(), src.RawFloats()) + b.ResetTimer() + b.ReportAllocs() + for b.Loop() { + if _, err := ArgSort(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkInterp2D samples one million bilinear queries against a +// 1024×1024 grid. +func BenchmarkInterp2D(b *testing.B) { + grid, _ := FromFloats(make([]float64, 1024*1024), 1024, 1024) + for i := range grid.RawFloats() { + grid.RawFloats()[i] = float64(i%97) / 97 + } + n := 1 << 20 + xs, _ := FromFloats(make([]float64, n), n) + ys, _ := FromFloats(make([]float64, n), n) + for i := range n { + xs.RawFloats()[i] = float64(i%1023) + 0.5 + ys.RawFloats()[i] = float64((i*7)%1023) + 0.25 + } + b.ReportAllocs() + for b.Loop() { + if _, err := Interpolate2D(grid, xs, ys, 0, 0, 1, 1); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/narrow_compare_alloc_test.go b/internal/core/narrow_compare_alloc_test.go new file mode 100644 index 0000000..d79856f --- /dev/null +++ b/internal/core/narrow_compare_alloc_test.go @@ -0,0 +1,622 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "fmt" + "math" + "math/big" + "testing" +) + +// The comparison family's result representation, measured and proved. +// Production answers a bool mask: one byte per element, the predicate +// itself. The arms named Legacy write the int64 0/1 mask the earlier +// representation kept, through local helpers over the same operands, +// so each Mask-versus-Legacy pair prices the same predicate in the two +// payloads. The pin at the foot of the file is the precision proof: +// over every relation, every arity and every operand class, the bool +// answer widened to int equals the plain-Go predicate computed inline. + +// cmpBenchLen is the element count the benchmark payloads fill. +const cmpBenchLen = 1 << 20 + +// narrowLtLegacyMask writes x < y as the 0/1 int64 mask of the earlier +// representation, windowed per chunk like the kernel it mirrors. +func narrowLtLegacyMask[T int8 | uint8](x, y []T, out []int64) { + parallelMin(len(out), elementwiseMinPerWorker, func(s, e int) { + xs, ys, os := x[s:e], y[s:e], out[s:e] + for i := range os { + if xs[i] < ys[i] { + os[i] = 1 + } + } + }) +} + +// mixedLtLegacyMask writes the int64-widened x < y of two different +// narrow widths as the 0/1 int64 mask. +func mixedLtLegacyMask[A, B int8 | uint8 | int16 | uint16 | int32 | uint32](x []A, y []B, out []int64) { + parallelMin(len(out), elementwiseMinPerWorker, func(s, e int) { + xs, ys, os := x[s:e], y[s:e], out[s:e] + for i := range os { + if int64(xs[i]) < int64(ys[i]) { + os[i] = 1 + } + } + }) +} + +// geScalarLegacyMask writes x >= v as the 0/1 int64 mask. The +// production GeI over a float64 payload answers applyIntFloat's exact +// relation, whose 0/1 this mask reproduces at the benchmark's scalar: +// NaN compares false, an infinity orders, and every float64 at or +// above 2^63 sits above the scalar. +func geScalarLegacyMask(xs []float64, v float64, out []int64) { + parallelMin(len(out), elementwiseMinPerWorker, func(s, e int) { + fs, os := xs[s:e], out[s:e] + for i := range os { + if fs[i] >= v { + os[i] = 1 + } + } + }) +} + +func BenchmarkCmpLtInt8Mask(b *testing.B) { + a, c := benchNarrowPair(Int8) + b.ReportAllocs() + for b.Loop() { + if _, err := Lt(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCmpLtInt8Legacy(b *testing.B) { + a, c := benchNarrowPair(Int8) + b.ReportAllocs() + for b.Loop() { + out := make([]int64, narrowBenchLen) + narrowLtLegacyMask(a.RawInt8s(), c.RawInt8s(), out) + } +} + +func BenchmarkCmpLtUint8Mask(b *testing.B) { + a, c := benchNarrowPair(Uint8) + b.ReportAllocs() + for b.Loop() { + if _, err := Lt(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCmpLtUint8Legacy(b *testing.B) { + a, c := benchNarrowPair(Uint8) + b.ReportAllocs() + for b.Loop() { + out := make([]int64, narrowBenchLen) + narrowLtLegacyMask(a.RawUint8s(), c.RawUint8s(), out) + } +} + +func BenchmarkCmpMixedLtI8U16Mask(b *testing.B) { + a, _ := benchNarrowPair(Int8) + c, _ := benchNarrowPair(Uint16) + b.ReportAllocs() + for b.Loop() { + if _, err := Lt(a, c); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCmpMixedLtI8U16Legacy(b *testing.B) { + a, _ := benchNarrowPair(Int8) + c, _ := benchNarrowPair(Uint16) + b.ReportAllocs() + for b.Loop() { + out := make([]int64, narrowBenchLen) + mixedLtLegacyMask(a.RawInt8s(), c.RawUint16s(), out) + } +} + +func BenchmarkCmpGeIFloat64Mask(b *testing.B) { + af, _ := benchNarrowPair(Float) + b.ReportAllocs() + for b.Loop() { + if _, err := GeI(af, 3000000001); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCmpGeIFloat64Legacy(b *testing.B) { + af, _ := benchNarrowPair(Float) + b.ReportAllocs() + for b.Loop() { + out := make([]int64, narrowBenchLen) + geScalarLegacyMask(af.RawFloats(), 3000000001, out) + } +} + +// The precision proof. Every relation of the comparison family, over +// array-against-array, array-against-int-scalar and array-against- +// float-scalar arities, over every operand class, must answer the bool +// the plain-Go predicate answers: the mask widened to int is checked +// element for element against a reference computed inline from the raw +// payloads. The value sets carry NaN, both infinities, both zeros, the +// 2^53 and 2^63 boundaries of int-float mixing and the width edges of +// every narrow class; the lengths split into parallelMin chunks exactly, +// with a remainder, and not at all. + +// proofRels pairs each public relation with its scalar twins. +var proofRels = []struct { + name string + op cmpOp + pair func(a, b *Array) (*Array, error) + scaI func(a *Array, v int64) (*Array, error) + scaF func(a *Array, v float64) (*Array, error) +}{ + {"Eq", opEq, Eq, EqI, EqF}, + {"Ne", opNe, Ne, NeI, NeF}, + {"Lt", opLt, Lt, LtI, LtF}, + {"Le", opLe, Le, LeI, LeF}, + {"Gt", opGt, Gt, GtI, GtF}, + {"Ge", opGe, Ge, GeI, GeF}, +} + +// proofNum carries one widened element: the integer class as its int64 +// value, the float classes as their float64 value, bool as 0/1. +type proofNum struct { + isInt bool + i int64 + f float64 +} + +func proofInt(v int64) proofNum { return proofNum{isInt: true, i: v} } +func proofFloat(v float64) proofNum { return proofNum{f: v} } + +// refIntOp answers op over two int64 values, plain Go. +func refIntOp(op cmpOp, x, y int64) bool { + switch op { + case opEq: + return x == y + case opNe: + return x != y + case opLt: + return x < y + case opLe: + return x <= y + case opGt: + return x > y + default: + return x >= y + } +} + +// refFloatOp answers op over two float64 values, plain Go: the IEEE +// comparison, NaN included. +func refFloatOp(op cmpOp, x, y float64) bool { + switch op { + case opEq: + return x == y + case opNe: + return x != y + case opLt: + return x < y + case opLe: + return x <= y + case opGt: + return x > y + default: + return x >= y + } +} + +// refRatOp answers op(x int, y float) over the rationals, exactly: the +// float converts losslessly, so no neighbour above 2^53 collapses and +// no boundary at 2^63 blurs. NaN and the infinities never reach here. +func refRatOp(op cmpOp, x int64, y float64) bool { + rx := new(big.Rat).SetInt64(x) + ry := new(big.Rat).SetFloat64(y) + if ry == nil { + panic("refRatOp: non-finite y") + } + switch rx.Cmp(ry) { + case 0: + return op == opEq || op == opLe || op == opGe + case -1: + return op == opNe || op == opLt || op == opLe + default: + return op == opNe || op == opGt || op == opGe + } +} + +// proofRel answers op(x, y) over the widened values exactly: NaN kills +// every relation but Ne, the infinities order on the extended real +// line, like classes compare in their own arithmetic and a mixed +// int-float pair compares over the rationals. +func proofRel(op cmpOp, x, y proofNum) bool { + if (!x.isInt && math.IsNaN(x.f)) || (!y.isInt && math.IsNaN(y.f)) { + return op == opNe + } + xInf, yInf := proofInfSign(x), proofInfSign(y) + if xInf != 0 || yInf != 0 { + return refInfRel(op, xInf, yInf) + } + switch { + case x.isInt && y.isInt: + return refIntOp(op, x.i, y.i) + case !x.isInt && !y.isInt: + return refFloatOp(op, x.f, y.f) + case x.isInt: + return refRatOp(op, x.i, y.f) + default: + // float x against int y: y op x, the mirrored relation, is the + // same predicate over the same rationals. + return refRatOp(op.mirror(), y.i, x.f) + } +} + +// proofInfSign reports +1, -1 or 0 as the widened value is an infinity. +func proofInfSign(v proofNum) int { + if v.isInt { + return 0 + } + switch { + case math.IsInf(v.f, 1): + return 1 + case math.IsInf(v.f, -1): + return -1 + } + return 0 +} + +// refInfRel compares with at least one infinity: -Inf sits below every +// finite, +Inf above every finite, and each infinity equals itself +// alone. +func refInfRel(op cmpOp, xInf, yInf int) bool { + var c int + switch { + case xInf == yInf: + c = 0 + case xInf == 1, yInf == -1: + c = 1 + default: + c = -1 + } + switch op { + case opEq: + return c == 0 + case opNe: + return c != 0 + case opLt: + return c < 0 + case opLe: + return c <= 0 + case opGt: + return c > 0 + default: + return c >= 0 + } +} + +// proofValue widens element i of a contiguous 1-D array exactly, from +// the raw payload: the integer class into int64, the float classes into +// float64, bool as 0/1. +func proofValue(a *Array, i int) proofNum { + switch a.dt { + case Int: + return proofInt(a.ints[i]) + case Int8: + return proofInt(int64(a.i8s[i])) + case Uint8: + return proofInt(int64(a.u8s[i])) + case Int16: + return proofInt(int64(a.i16s[i])) + case Uint16: + return proofInt(int64(a.u16s[i])) + case Int32: + return proofInt(int64(a.i32s[i])) + case Uint32: + return proofInt(int64(a.u32s[i])) + case Bool: + return proofInt(boolToInt64(a.bools[i])) + case Float16: + return proofFloat(HalfToFloat64(a.halves[i])) + case Float32: + return proofFloat(float64(a.floats32[i])) + case Float: + return proofFloat(a.floats[i]) + default: + panic(fmt.Sprintf("proofValue: %s carries no real width", a.dt)) + } +} + +// proofClasses lists every operand class the proof runs over, with two +// deterministic fills per class so the pair arity compares distinct +// arrays. Each value set carries the class's width edges, and the +// float sets carry NaN, both infinities, both zeros and the 2^53 and +// 2^63 boundaries the int-float mixing decides on. +var proofClasses = []struct { + name string + dt Dtype + complex bool + fill func(t *testing.T, n, rot int) *Array +}{ + {"Int", Int, false, func(t *testing.T, n, rot int) *Array { + set := []int64{0, 1, -1, 2, -2, 1 << 53, 1<<53 + 1, -(1<<53 + 1), + 1 << 60, 1<<60 + 1, 1<<60 + 64, math.MaxInt64, math.MinInt64} + return proofMust(t, "FromInts")(FromInts(proofCycle(set, n, rot), n)) + }}, + {"Int8", Int8, false, func(t *testing.T, n, rot int) *Array { + set := []int8{math.MinInt8, math.MaxInt8, -1, 0, 1, 2, -2} + return proofMust(t, "FromInt8s")(FromInt8s(proofCycle(set, n, rot), n)) + }}, + {"Uint8", Uint8, false, func(t *testing.T, n, rot int) *Array { + set := []uint8{0, 1, math.MaxUint8, 2, math.MaxUint8 - 1} + return proofMust(t, "FromUint8s")(FromUint8s(proofCycle(set, n, rot), n)) + }}, + {"Int16", Int16, false, func(t *testing.T, n, rot int) *Array { + set := []int16{math.MinInt16, math.MaxInt16, -1, 0, 1, 2, -2} + return proofMust(t, "FromInt16s")(FromInt16s(proofCycle(set, n, rot), n)) + }}, + {"Uint16", Uint16, false, func(t *testing.T, n, rot int) *Array { + set := []uint16{0, 1, math.MaxUint16, 2, math.MaxUint16 - 1} + return proofMust(t, "FromUint16s")(FromUint16s(proofCycle(set, n, rot), n)) + }}, + {"Int32", Int32, false, func(t *testing.T, n, rot int) *Array { + set := []int32{math.MinInt32, math.MaxInt32, -1, 0, 1, 2, -2} + return proofMust(t, "FromInt32s")(FromInt32s(proofCycle(set, n, rot), n)) + }}, + {"Uint32", Uint32, false, func(t *testing.T, n, rot int) *Array { + set := []uint32{0, 1, math.MaxUint32, 2, math.MaxUint32 - 1} + return proofMust(t, "FromUint32s")(FromUint32s(proofCycle(set, n, rot), n)) + }}, + {"Float", Float, false, func(t *testing.T, n, rot int) *Array { + set := []float64{math.NaN(), math.Inf(1), math.Inf(-1), 0, + math.Copysign(0, -1), 1, -1, 0.5, -0.5, 9007199254740992, + 9007199254740993, 9223372036854775808, -9223372036854775808, + 1e300, -1e300, 3} + return proofMust(t, "FromFloats")(FromFloats(proofCycle(set, n, rot), n)) + }}, + {"Float32", Float32, false, func(t *testing.T, n, rot int) *Array { + set := []float32{float32(math.NaN()), float32(math.Inf(1)), + float32(math.Inf(-1)), 0, float32(math.Copysign(0, -1)), 1, -1, + 0.5, 3.5, math.MaxFloat32, -math.MaxFloat32} + return proofMust(t, "FromFloat32s")(FromFloat32s(proofCycle(set, n, rot), n)) + }}, + {"Float16", Float16, false, func(t *testing.T, n, rot int) *Array { + set := []float64{math.NaN(), math.Inf(1), math.Inf(-1), 0, + math.Copysign(0, -1), 1, -1, 0.5, 2, 65504, -65504} + return proofMust(t, "FromFloat16s")(FromFloat16s(proofCycle(set, n, rot), n)) + }}, + {"Bool", Bool, false, func(t *testing.T, n, rot int) *Array { + v := make([]bool, n) + for i := range v { + v[i] = (i+rot)%2 == 0 + } + return proofMust(t, "FromBools")(FromBools(v, n)) + }}, + {"Complex", Complex, true, func(t *testing.T, n, rot int) *Array { + set := []complex128{0, 1, 1i, 1 + 2i, 3, complex(math.NaN(), 0), + complex(0, math.Inf(1))} + return proofMust(t, "FromComplexes")(FromComplexes(proofCycle(set, n, rot), n)) + }}, +} + +// proofCycle fills n elements from set, starting rot entries in, so rot +// 0 and rot 1 hand the pair arity two different arrays. Every fourth +// element of the rot 1 fill repeats rot 0's value, so the pair arity +// sees equal elements too: the Eq answer, the Lt and Le boundary and +// the Ne negation all rest on them. +func proofCycle[T int64 | int8 | uint8 | int16 | uint16 | int32 | uint32 | float64 | float32 | complex128](set []T, n, rot int) []T { + v := make([]T, n) + for i := range v { + j := (i + rot*3) % len(set) + if rot == 1 && i%4 == 0 { + j = i % len(set) + } + v[i] = set[j] + } + return v +} + +// proofMust adapts a constructor's pair to a fill body, failing the +// test on the constructor's error. +func proofMust(t *testing.T, what string) func(*Array, error) *Array { + t.Helper() + return func(a *Array, err error) *Array { + t.Helper() + if err != nil { + t.Fatalf("%s: %v", what, err) + } + if a == nil { + t.Fatalf("%s: nil array", what) + } + return a + } +} + +// TestCompareBoolMaskMatchesInlinePredicates is the no-precision-loss +// pin: over every relation, every arity, every operand class, the +// bool mask widened to int equals the reference predicate computed +// inline from the raw payloads. +func TestCompareBoolMaskMatchesInlinePredicates(t *testing.T) { + prev := SetNumCPU(4) + defer SetNumCPU(prev) + // pinChunkLen splits into four workers with a remainder over the + // 1024 per-worker floor, 4096 into four exact ones, 7 keeps one + // worker on the serial path. + for _, n := range []int{pinChunkLen, 4096, 7} { + proofRun(t, n) + } +} + +// proofRun walks the whole relation-by-arity-by-class matrix at one +// length. +func proofRun(t *testing.T, n int) { + t.Helper() + for _, cls := range proofClasses { + a := cls.fill(t, n, 0) + b := cls.fill(t, n, 1) + for _, rel := range proofRels { + if cls.complex && rel.op.isOrdering() { + continue + } + got, err := rel.pair(a, b) + if err != nil { + t.Fatalf("%s %s pair: %v", cls.name, rel.name, err) + } + label := fmt.Sprintf("%s/%s pair n=%d", cls.name, rel.name, n) + proofCheckMask(t, label, got, n, func(i int) bool { + if cls.complex { + return proofComplexRel(rel.op, proofComplex(a, i), proofComplex(b, i)) + } + return proofRel(rel.op, proofValue(a, i), proofValue(b, i)) + }) + } + for _, sv := range []int64{3, 0, -1} { + for _, rel := range proofRels { + if cls.complex && rel.op.isOrdering() { + continue + } + got, err := rel.scaI(a, sv) + if err != nil { + t.Fatalf("%s %sI(%d): %v", cls.name, rel.name, sv, err) + } + label := fmt.Sprintf("%s/%sI scalar %d n=%d", cls.name, rel.name, sv, n) + proofCheckMask(t, label, got, n, func(i int) bool { + if cls.complex { + return proofComplexRel(rel.op, proofComplex(a, i), complex(float64(sv), 0)) + } + return proofRel(rel.op, proofValue(a, i), proofInt(sv)) + }) + } + } + for _, wv := range []float64{2.5, math.NaN(), math.Inf(1), math.Inf(-1), 0, math.Copysign(0, -1)} { + for _, rel := range proofRels { + if cls.complex && rel.op.isOrdering() { + continue + } + got, err := rel.scaF(a, wv) + if err != nil { + t.Fatalf("%s %sF(%v): %v", cls.name, rel.name, wv, err) + } + label := fmt.Sprintf("%s/%sF scalar %v n=%d", cls.name, rel.name, wv, n) + proofCheckMask(t, label, got, n, func(i int) bool { + if cls.complex { + return proofComplexRel(rel.op, proofComplex(a, i), complex(wv, 0)) + } + return proofRel(rel.op, proofValue(a, i), proofFloat(wv)) + }) + } + } + } + proofMixedPairs(t, n) + proofIntFloatPairs(t, n) +} + +// proofComplexRel answers Eq or Ne over two complex values, plain Go: +// the component equality the contract states, NaN parts unequal to +// everything. +func proofComplexRel(op cmpOp, x, y complex128) bool { + if op == opNe { + return x != y + } + return x == y +} + +func proofComplex(a *Array, i int) complex128 { return a.complexes[i] } + +// proofMixedPairs holds every ordered pair of different narrow widths: +// both sides widen exactly into int64, so the reference compares the +// int64 widenings. +func proofMixedPairs(t *testing.T, n int) { + t.Helper() + var narrow []*struct { + name string + dt Dtype + } + for _, cls := range proofClasses { + if narrowIntClass(cls.dt) { + narrow = append(narrow, &struct { + name string + dt Dtype + }{cls.name, cls.dt}) + } + } + for _, ca := range narrow { + for _, cb := range narrow { + if ca.dt == cb.dt { + continue + } + a := proofFill(t, ca.name, n, 0) + b := proofFill(t, cb.name, n, 1) + for _, rel := range proofRels { + got, err := rel.pair(a, b) + if err != nil { + t.Fatalf("%s/%s %s pair: %v", ca.name, cb.name, rel.name, err) + } + label := fmt.Sprintf("%s/%s %s pair n=%d", ca.name, cb.name, rel.name, n) + proofCheckMask(t, label, got, n, func(i int) bool { + return proofRel(rel.op, proofValue(a, i), proofValue(b, i)) + }) + } + } + } +} + +// proofIntFloatPairs holds the int-float mixing both ways, the path +// applyIntFloat keeps exact while the float64 widening would round the +// int side above 2^53. +func proofIntFloatPairs(t *testing.T, n int) { + t.Helper() + for _, pair := range [][2]string{{"Int", "Float"}, {"Float", "Int"}} { + a := proofFill(t, pair[0], n, 0) + b := proofFill(t, pair[1], n, 1) + for _, rel := range proofRels { + got, err := rel.pair(a, b) + if err != nil { + t.Fatalf("%s/%s %s pair: %v", pair[0], pair[1], rel.name, err) + } + label := fmt.Sprintf("%s/%s %s pair n=%d", pair[0], pair[1], rel.name, n) + proofCheckMask(t, label, got, n, func(i int) bool { + return proofRel(rel.op, proofValue(a, i), proofValue(b, i)) + }) + } + } +} + +// proofFill builds one class's array by name. +func proofFill(t *testing.T, name string, n, rot int) *Array { + t.Helper() + for _, cls := range proofClasses { + if cls.name == name { + return cls.fill(t, n, rot) + } + } + t.Fatalf("no proof class %s", name) + return nil +} + +// proofCheckMask holds one mask against the reference: the dtype must +// be bool and every element, widened to int, must equal the predicate's +// answer. +func proofCheckMask(t *testing.T, label string, m *Array, n int, ref func(i int) bool) { + t.Helper() + if m.Dtype() != Bool { + t.Fatalf("%s: dtype %s, want bool", label, m.Dtype()) + } + if m.Len() != n { + t.Fatalf("%s: %d elements, want %d", label, m.Len(), n) + } + for i := range n { + if g, _ := IntAt(m, i); (g != 0) != ref(i) { + t.Fatalf("%s: element %d = %d, want %v", label, i, g, ref(i)) + } + } +} diff --git a/internal/core/narrow_dispatch_pin_test.go b/internal/core/narrow_dispatch_pin_test.go new file mode 100644 index 0000000..d7e8ff7 --- /dev/null +++ b/internal/core/narrow_dispatch_pin_test.go @@ -0,0 +1,432 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// The narrow dispatch walks slice every parallel chunk, so a kernel +// that indexes its source from zero instead of from its own chunk +// start corrupts everything past the first chunk while small fixtures +// stay green. These pins run lengths that split into four chunks under +// a pinned worker count and compare every element against an +// independently computed reference; the narrow arithmetic reference +// computes in int64 and narrows on store, which is the keeping-kind +// contract. + +// pinChunkLen is four chunks of 1250, above the 1024-element floor. +const pinChunkLen = 5000 + +func pinNarrowFixtures() (ints []int64, i8 []int8, u8 []uint8, i16 []int16, u16 []uint16, i32 []int32, u32 []uint32) { + n := pinChunkLen + ints, i8, u8, i16, u16, i32, u32 = + make([]int64, n), make([]int8, n), make([]uint8, n), + make([]int16, n), make([]uint16, n), make([]int32, n), make([]uint32, n) + for i := range n { + ints[i] = int64(i)*37 - 1000 + i8[i] = int8(i*29 - 60) // wraps through the width on store + u8[i] = uint8(i * 31) + i16[i] = int16(i*701 - 9000) + u16[i] = uint16(i * 1301) + i32[i] = int32(i*90001 - 7) + u32[i] = uint32(3000000000 + i*13) // above 2^31, so a narrowed widening loses the high bits + } + return ints, i8, u8, i16, u16, i32, u32 +} + +func TestScalarMapKeepingChunks(t *testing.T) { + prev := SetNumCPU(4) + defer SetNumCPU(prev) + ints, i8, u8, i16, u16, i32, u32 := pinNarrowFixtures() + n := pinChunkLen + const addV, subV, mulV = 7, 3, 5 + + ai, err := FromInts(ints, n) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + got := AddI(ai, addV) + for i := range n { + want := ints[i] + addV + if g, _ := IntAt(got, i); g != want { + t.Fatalf("AddI(Int)[%d] = %d, want %d", i, g, want) + } + } + got = SubI(ai, subV) + for i := range n { + want := ints[i] - subV + if g, _ := IntAt(got, i); g != want { + t.Fatalf("SubI(Int)[%d] = %d, want %d", i, g, want) + } + } + got = MulI(ai, mulV) + for i := range n { + want := ints[i] * mulV + if g, _ := IntAt(got, i); g != want { + t.Fatalf("MulI(Int)[%d] = %d, want %d", i, g, want) + } + } + + a8, err := FromInt8s(i8, n) + if err != nil { + t.Fatalf("FromInt8s: %v", err) + } + g8 := AddI(a8, addV).RawInt8s() + for i := range n { + want := int8(int64(i8[i]) + addV) + if g8[i] != want { + t.Fatalf("AddI(Int8)[%d] = %d, want %d", i, g8[i], want) + } + } + g8 = SubI(a8, subV).RawInt8s() + for i := range n { + want := int8(int64(i8[i]) - subV) + if g8[i] != want { + t.Fatalf("SubI(Int8)[%d] = %d, want %d", i, g8[i], want) + } + } + g8 = MulI(a8, mulV).RawInt8s() + for i := range n { + want := int8(int64(i8[i]) * mulV) + if g8[i] != want { + t.Fatalf("MulI(Int8)[%d] = %d, want %d", i, g8[i], want) + } + } + + au8, err := FromUint8s(u8, n) + if err != nil { + t.Fatalf("FromUint8s: %v", err) + } + gu8 := MulI(au8, mulV).RawUint8s() + for i := range n { + want := uint8(int64(u8[i]) * mulV) + if gu8[i] != want { + t.Fatalf("MulI(Uint8)[%d] = %d, want %d", i, gu8[i], want) + } + } + + a16, err := FromInt16s(i16, n) + if err != nil { + t.Fatalf("FromInt16s: %v", err) + } + g16 := AddI(a16, addV).RawInt16s() + for i := range n { + want := int16(int64(i16[i]) + addV) + if g16[i] != want { + t.Fatalf("AddI(Int16)[%d] = %d, want %d", i, g16[i], want) + } + } + + au16, err := FromUint16s(u16, n) + if err != nil { + t.Fatalf("FromUint16s: %v", err) + } + gu16 := SubI(au16, subV).RawUint16s() + for i := range n { + want := uint16(int64(u16[i]) - subV) + if gu16[i] != want { + t.Fatalf("SubI(Uint16)[%d] = %d, want %d", i, gu16[i], want) + } + } + + a32, err := FromInt32s(i32, n) + if err != nil { + t.Fatalf("FromInt32s: %v", err) + } + g32 := MulI(a32, mulV).RawInt32s() + for i := range n { + want := int32(int64(i32[i]) * mulV) + if g32[i] != want { + t.Fatalf("MulI(Int32)[%d] = %d, want %d", i, g32[i], want) + } + } + + au32, err := FromUint32s(u32, n) + if err != nil { + t.Fatalf("FromUint32s: %v", err) + } + gu32 := AddI(au32, addV).RawUint32s() + for i := range n { + want := uint32(int64(u32[i]) + addV) + if gu32[i] != want { + t.Fatalf("AddI(Uint32)[%d] = %d, want %d", i, gu32[i], want) + } + } + + // The widening float64 walk (AddF) reads the same payloads. + gf := AddF(a8, float64(addV)).RawFloats() + for i := range n { + want := float64(i8[i]) + float64(addV) + if gf[i] != want { + t.Fatalf("AddF(Int8)[%d] = %v, want %v", i, gf[i], want) + } + } + gf = AddF(ai, float64(addV)).RawFloats() + for i := range n { + want := float64(ints[i]) + float64(addV) + if gf[i] != want { + t.Fatalf("AddF(Int)[%d] = %v, want %v", i, gf[i], want) + } + } +} + +func TestCmpMixedNarrowAndBool(t *testing.T) { + prev := SetNumCPU(4) + defer SetNumCPU(prev) + ints, i8, u8, i16, u16, i32, u32 := pinNarrowFixtures() + n := pinChunkLen + bl := make([]bool, n) + for i := range n { + bl[i] = i%3 == 0 + } + + ab, err := FromBools(bl, n) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + bl2 := make([]bool, n) + for i := range n { + bl2[i] = i%2 == 0 + } + ab2, err := FromBools(bl2, n) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + a8, _ := FromInt8s(i8, n) + au8, _ := FromUint8s(u8, n) + a16, _ := FromInt16s(i16, n) + au16, _ := FromUint16s(u16, n) + a32, _ := FromInt32s(i32, n) + au32, _ := FromUint32s(u32, n) + ai, _ := FromInts(ints, n) + + // Every pair below walks the mixed kernel (two narrow widths) or + // the bool-plus-narrow case, which must widen the boolean to 0/1 + // and never read a payload the dtype does not carry. + widen := map[*Array]func(int) int64{ + ab: func(i int) int64 { + if bl[i] { + return 1 + } + return 0 + }, + a8: func(i int) int64 { return int64(i8[i]) }, + au8: func(i int) int64 { return int64(u8[i]) }, + a16: func(i int) int64 { return int64(i16[i]) }, + au16: func(i int) int64 { return int64(u16[i]) }, + a32: func(i int) int64 { return int64(i32[i]) }, + au32: func(i int) int64 { return int64(u32[i]) }, + ai: func(i int) int64 { return ints[i] }, + } + widen[ab2] = func(i int) int64 { + if bl2[i] { + return 1 + } + return 0 + } + pairs := [][2]*Array{ + {ab, a8}, {ab, au32}, {a8, a16}, {au8, a32}, {a32, au16}, {au32, a8}, + {a8, a8}, {ab, ab}, {a8, ai}, + {ab, ab2}, {a8, ab}, {au32, ab}, {ab, ai}, + } + rels := []struct { + name string + fn func(*Array, *Array) (*Array, error) + want func(int64, int64) bool + }{ + {"Eq", Eq, func(x, y int64) bool { return x == y }}, + {"Ne", Ne, func(x, y int64) bool { return x != y }}, + {"Lt", Lt, func(x, y int64) bool { return x < y }}, + {"Le", Le, func(x, y int64) bool { return x <= y }}, + {"Gt", Gt, func(x, y int64) bool { return x > y }}, + {"Ge", Ge, func(x, y int64) bool { return x >= y }}, + } + for _, p := range pairs { + fa, fb := widen[p[0]], widen[p[1]] + for _, r := range rels { + mask, merr := r.fn(p[0], p[1]) + if merr != nil { + t.Fatalf("%s mixed: %v", r.name, merr) + } + for i := range n { + want := int64(0) + if r.want(fa(i), fb(i)) { + want = 1 + } + if g, _ := IntAt(mask, i); g != want { + t.Fatalf("%s mixed [%d] = %d, want %d (widened %d, %d)", + r.name, i, g, want, fa(i), fb(i)) + } + } + } + } + + // The scalar relation over a bool payload is the same widening + // question against an int scalar. + mask, merr := LtI(ab, 1) + if merr != nil { + t.Fatalf("LtI on bool: %v", merr) + } + for i := range n { + b := int64(0) + if bl[i] { + b = 1 + } + want := int64(0) + if b < 1 { + want = 1 + } + if g, _ := IntAt(mask, i); g != want { + t.Fatalf("LtI(bool)[%d] = %d, want %d", i, g, want) + } + } +} + +func TestBoolLogicChunks(t *testing.T) { + prev := SetNumCPU(4) + defer SetNumCPU(prev) + const n = pinChunkLen + x := make([]bool, n) + y := make([]bool, n) + for i := range n { + x[i] = i%3 == 0 + y[i] = i%2 == 0 + } + ax, err := FromBools(x, n) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + ay, err := FromBools(y, n) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + and, aerr := And(ax, ay) + if aerr != nil { + t.Fatalf("And: %v", aerr) + } + or, oerr := Or(ax, ay) + if oerr != nil { + t.Fatalf("Or: %v", oerr) + } + xor, xerr := Xor(ax, ay) + if xerr != nil { + t.Fatalf("Xor: %v", xerr) + } + not, nerr := Not(ax) + if nerr != nil { + t.Fatalf("Not: %v", nerr) + } + gotAnd, gotOr, gotXor, gotNot := and.RawBools(), or.RawBools(), xor.RawBools(), not.RawBools() + for i := range n { + if gotAnd[i] != (x[i] && y[i]) { + t.Fatalf("And[%d] = %v, want %v", i, gotAnd[i], x[i] && y[i]) + } + if gotOr[i] != (x[i] || y[i]) { + t.Fatalf("Or[%d] = %v, want %v", i, gotOr[i], x[i] || y[i]) + } + if gotXor[i] != (x[i] != y[i]) { + t.Fatalf("Xor[%d] = %v, want %v", i, gotXor[i], x[i] != y[i]) + } + if gotNot[i] != !x[i] { + t.Fatalf("Not[%d] = %v, want %v", i, gotNot[i], !x[i]) + } + } +} + +func TestScalarCmpBoolFloatWidening(t *testing.T) { + prev := SetNumCPU(4) + defer SetNumCPU(prev) + const n = pinChunkLen + bl := make([]bool, n) + for i := range n { + bl[i] = i%3 == 0 + } + ab, err := FromBools(bl, n) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + // A bool against a float scalar widens to 0/1 in float64. + ltf, lerr := LtF(ab, 0.5) + if lerr != nil { + t.Fatalf("LtF on bool: %v", lerr) + } + gef, gerr := GeF(ab, 0.5) + if gerr != nil { + t.Fatalf("GeF on bool: %v", gerr) + } + for i := range n { + want := int64(0) + if !bl[i] { + want = 1 + } + if g, _ := IntAt(ltf, i); g != want { + t.Fatalf("LtF(bool)[%d] = %d, want %d", i, g, want) + } + want = 0 + if bl[i] { + want = 1 + } + if g, _ := IntAt(gef, i); g != want { + t.Fatalf("GeF(bool)[%d] = %d, want %d", i, g, want) + } + } +} + +func TestDotIntegerKindsChunked(t *testing.T) { + prev := SetNumCPU(4) + defer SetNumCPU(prev) + _, i8, _, _, _, i32, u32 := pinNarrowFixtures() + n := pinChunkLen + j8 := make([]int8, n) + j32 := make([]int32, n) + ju32 := make([]uint32, n) + for i := range n { + j8[i] = int8(i*7 + 5) + j32[i] = int32(i*90001 + 17) + ju32[i] = uint32(2000000000 + i*29) + } + // The fold computes int64 products with wrapping addition over a + // partition fixed by the length; wrapping addition is associative, + // so the sequential sum pins the value whatever the partition. The + // uint32 pair wraps its products, which the pin exercises above + // 2^31 in both widenings. + a8, _ := FromInt8s(i8, n) + b8, _ := FromInt8s(j8, n) + d8, err := Dot(a8, b8) + if err != nil { + t.Fatalf("Dot int8: %v", err) + } + var s8 int64 + for i := range n { + s8 += int64(i8[i]) * int64(j8[i]) + } + if d8.i != s8 { + t.Fatalf("Dot int8 = %d, want %d", d8.i, s8) + } + a32, _ := FromInt32s(i32, n) + b32, _ := FromInt32s(j32, n) + d32, err := Dot(a32, b32) + if err != nil { + t.Fatalf("Dot int32: %v", err) + } + var s32 int64 + for i := range n { + s32 += int64(i32[i]) * int64(j32[i]) + } + if d32.i != s32 { + t.Fatalf("Dot int32 = %d, want %d", d32.i, s32) + } + au32, _ := FromUint32s(u32, n) + bu32, _ := FromUint32s(ju32, n) + du32, err := Dot(au32, bu32) + if err != nil { + t.Fatalf("Dot uint32: %v", err) + } + var su32 int64 + for i := range n { + su32 += int64(u32[i]) * int64(ju32[i]) + } + if du32.i != su32 { + t.Fatalf("Dot uint32 = %d, want %d", du32.i, su32) + } +} diff --git a/internal/core/narrowdt.go b/internal/core/narrowdt.go new file mode 100644 index 0000000..f07d3f2 --- /dev/null +++ b/internal/core/narrowdt.go @@ -0,0 +1,325 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +// The narrow element types: bool and the small integers. They exist so +// byte payloads, masks and the file formats' short integers ride in the +// width they were born in instead of widening to int64 in transit. The +// constructors here mirror the established pair: From…s copies its +// values, …FromArray takes ownership of the caller's slice. The byte +// serialisation stays deliberately Int-only: Bytes answers +// only for an Int array and FromBytes builds Int, so a narrow payload +// converts through Astype first rather than riding Bytes natively. + +// FromBools builds a bool array from vals, copying them. +func FromBools(vals []bool, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + bools := make([]bool, len(vals)) + copy(bools, vals) + return &Array{shape: sh, dt: Bool, bools: bools}, nil +} + +// FromInt8s builds an int8 array from vals, copying them. +func FromInt8s(vals []int8, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + out := make([]int8, len(vals)) + copy(out, vals) + return &Array{shape: sh, dt: Int8, i8s: out}, nil +} + +// FromUint8s builds a uint8 array from vals, copying them. +func FromUint8s(vals []uint8, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + out := make([]uint8, len(vals)) + copy(out, vals) + return &Array{shape: sh, dt: Uint8, u8s: out}, nil +} + +// FromInt16s builds an int16 array from vals, copying them. +func FromInt16s(vals []int16, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + out := make([]int16, len(vals)) + copy(out, vals) + return &Array{shape: sh, dt: Int16, i16s: out}, nil +} + +// FromUint16s builds a uint16 array from vals, copying them. +func FromUint16s(vals []uint16, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + out := make([]uint16, len(vals)) + copy(out, vals) + return &Array{shape: sh, dt: Uint16, u16s: out}, nil +} + +// FromInt32s builds an int32 array from vals, copying them. +func FromInt32s(vals []int32, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + out := make([]int32, len(vals)) + copy(out, vals) + return &Array{shape: sh, dt: Int32, i32s: out}, nil +} + +// FromUint32s builds a uint32 array from vals, copying them. +func FromUint32s(vals []uint32, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + out := make([]uint32, len(vals)) + copy(out, vals) + return &Array{shape: sh, dt: Uint32, u32s: out}, nil +} + +// BoolsFromArray builds a bool array that takes ownership of vals, with +// the FloatsFromArray contract. +func BoolsFromArray(vals []bool, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Bool, bools: vals}, nil +} + +// Int8sFromArray builds an int8 array that takes ownership of vals, +// with the FloatsFromArray contract. +func Int8sFromArray(vals []int8, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Int8, i8s: vals}, nil +} + +// Uint8sFromArray builds a uint8 array that takes ownership of vals, +// with the FloatsFromArray contract. +func Uint8sFromArray(vals []uint8, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Uint8, u8s: vals}, nil +} + +// Int16sFromArray builds an int16 array that takes ownership of vals, +// with the FloatsFromArray contract. +func Int16sFromArray(vals []int16, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Int16, i16s: vals}, nil +} + +// Uint16sFromArray builds a uint16 array that takes ownership of vals, +// with the FloatsFromArray contract. +func Uint16sFromArray(vals []uint16, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Uint16, u16s: vals}, nil +} + +// Int32sFromArray builds an int32 array that takes ownership of vals, +// with the FloatsFromArray contract. +func Int32sFromArray(vals []int32, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Int32, i32s: vals}, nil +} + +// Uint32sFromArray builds a uint32 array that takes ownership of vals, +// with the FloatsFromArray contract. +func Uint32sFromArray(vals []uint32, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Uint32, u32s: vals}, nil +} + +// RawBools returns the bool payload with the RawFloats contract. +func (a *Array) RawBools() []bool { return a.bools } + +// RawInt8s returns the int8 payload with the RawFloats contract. +func (a *Array) RawInt8s() []int8 { return a.i8s } + +// RawUint8s returns the uint8 payload with the RawFloats contract. +func (a *Array) RawUint8s() []uint8 { return a.u8s } + +// RawInt16s returns the int16 payload with the RawFloats contract. +func (a *Array) RawInt16s() []int16 { return a.i16s } + +// RawUint16s returns the uint16 payload with the RawFloats contract. +func (a *Array) RawUint16s() []uint16 { return a.u16s } + +// RawInt32s returns the int32 payload with the RawFloats contract. +func (a *Array) RawInt32s() []int32 { return a.i32s } + +// RawUint32s returns the uint32 payload with the RawFloats contract. +func (a *Array) RawUint32s() []uint32 { return a.u32s } + +// boolAt returns element i as a bool: the boolean payload's own value, +// and everywhere else the test against zero. A NaN compares unequal to +// zero, so a NaN element reads true, the same masked-read decision the +// int mask path makes on it. +func (a *Array) boolAt(i int) bool { + if a.strides != nil { + i = a.physIndex(i) + } + switch a.dt { + case Bool: + return a.bools[i] + case Int: + return a.ints[i] != 0 + case Float16: + return a.halves[i]&0x7FFF != 0 + case Float32: + return a.floats32[i] != 0 + case Float: + return a.floats[i] != 0 + case Complex: + return a.complexes[i] != 0 + case Int8: + return a.i8s[i] != 0 + case Uint8: + return a.u8s[i] != 0 + case Int16: + return a.i16s[i] != 0 + case Uint16: + return a.u16s[i] != 0 + case Int32: + return a.i32s[i] != 0 + default: + return a.u32s[i] != 0 + } +} + +// BoolAt returns the element at the given index as a bool: the boolean +// payload's own value, and every other dtype read against zero. It is +// the widening reader the small integer types share with the internal +// accessors; unlike IntAt and FloatAt it takes any dtype, because the +// mask question makes sense over all of them. +func BoolAt(a *Array, index ...int) (bool, error) { + off, err := flatIndex(a.shape, index) + if err != nil { + return false, err + } + return a.boolAt(off), nil +} + +// intClassOrder lists the integer-class dtypes in containment order, +// the axis the promotion table below is indexed by. +var intClassOrder = [...]Dtype{Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int} + +// intClassIndex maps an integer-class dtype to its row in the promotion +// table; an unlisted ordinal lands on Int, the recorded default. +func intClassIndex(d Dtype) int { + for i, c := range intClassOrder { + if c == d { + return i + } + } + return len(intClassOrder) - 1 +} + +// intClass reports whether d belongs to the integer class: Bool or one +// of the integer dtypes, everything promote resolves by containment. +func intClass(d Dtype) bool { + switch d { + case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int: + return true + } + return false +} + +// cloneArray returns a deep copy of the array: same shape, same dtype, +// own payload. The narrow element types clone here, where their payload +// fields live; cloneData keeps serving the five legacy slices to the +// callers that consume them individually. +func (a *Array) cloneArray() *Array { + out := &Array{shape: a.Shape(), dt: a.dt} + n := a.Len() + out.alloc(n) + if a.strides != nil { + for i := range n { + out.setFrom(i, a, i) + } + return out + } + switch a.dt { + case Int: + copy(out.ints, a.ints[:n]) + case Float16: + copy(out.halves, a.halves[:n]) + case Float32: + copy(out.floats32, a.floats32[:n]) + case Float: + copy(out.floats, a.floats[:n]) + case Complex: + copy(out.complexes, a.complexes[:n]) + case Bool: + copy(out.bools, a.bools[:n]) + case Int8: + copy(out.i8s, a.i8s[:n]) + case Uint8: + copy(out.u8s, a.u8s[:n]) + case Int16: + copy(out.i16s, a.i16s[:n]) + case Uint16: + copy(out.u16s, a.u16s[:n]) + case Int32: + copy(out.i32s, a.i32s[:n]) + default: + copy(out.u32s, a.u32s[:n]) + } + return out +} + +// intPromote answers the smallest dtype of the integer class whose +// value range contains both operands' ranges. Mixed signedness pairs +// therefore widen instead of losing negative values: int8 with uint8 +// answers int16, int16 with uint16 answers int32, int32 with uint32 +// answers int. promote answers a same-dtype pair before the table is +// ever consulted, and the signed rows widen their own kind to the next +// width as the containment rule dictates; the table itself is +// symmetric. +var intPromote = [len(intClassOrder)][len(intClassOrder)]Dtype{ + // Bool row: every integer dtype contains {0, 1}. + {Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32, Int}, + // Int8 [-128, 127]: with Uint8 needs the 16-bit signed width. + {Int8, Int16, Int16, Int16, Int32, Int32, Int, Int}, + // Uint8 [0, 255]: contained by every wider dtype of either sign. + {Uint8, Int16, Uint8, Int16, Uint16, Int32, Uint32, Int}, + // Int16: with Uint16 needs 32-bit signed; with Uint32 needs Int. + {Int16, Int16, Int16, Int32, Int32, Int32, Int, Int}, + // Uint16 [0, 65535]: contained by Int32 and everything wider. + {Uint16, Int32, Uint16, Int32, Uint16, Int32, Uint32, Int}, + // Int32: with Uint32 needs the 64-bit signed width. + {Int32, Int32, Int32, Int32, Int32, Int, Int, Int}, + // Uint32 [0, 2^32-1]: contained by Int. + {Uint32, Int, Uint32, Int, Uint32, Int, Uint32, Int}, + // Int contains the whole class. + {Int, Int, Int, Int, Int, Int, Int, Int}, +} diff --git a/internal/core/onehot.go b/internal/core/onehot.go new file mode 100644 index 0000000..a7d485b --- /dev/null +++ b/internal/core/onehot.go @@ -0,0 +1,57 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// One-hot encoding. The sole producer here is the lookup pathway: +// a differentiable row-gather collapses to one-hot·weights, which the +// existing matrix-product backward already differentiates exactly, so +// embedding gradients need no dedicated kernel. + +// OneHot expands integer class codes into indicator vectors: every code +// becomes a length-classes row of zeros with a single one at the code's +// position, appended as a NEW LAST AXIS, so shape (…) turns into +// (…, classes). Codes must be an int array and each value inside +// [0, classes). The result is float32: indicators are exactly +// representable there and downstream weight matrices stay in the fast +// float32 lane. +func OneHot(codes *Array, classes int) (*Array, error) { + if codes.dt != Int { + return nil, errf("OneHot: codes must be an int array, got %s", codes.dt) + } + if classes < 1 { + return nil, errf("OneHot: classes must be positive, got %d", classes) + } + n := codes.Len() + for i := range n { + if c := codes.ints[i]; c < 0 || int(c) >= classes { + return nil, errf("OneHot: code %d out of range [0, %d) at position %d", c, classes, i) + } + } + // Bound the row length before multiplying: a wrapped product must be + // an error, not a negative allocation. + if classes != 0 && n > math.MaxInt/classes { + return nil, errf("OneHot: %d rows of width %d hold more elements than fit in an index", n, classes) + } + total, shape, terr := checkedDims(append(codes.Shape(), classes)) + if terr != nil { + return nil, base.WrapErr("OneHot", terr) + } + + out := &Array{shape: shape, dt: Float32} + out.floats32 = make([]float32, total) + engine.Parallel(n, func(rs, re int) { + for i := rs; i < re; i++ { + out.floats32[i*classes+int(codes.ints[i])] = 1 + } + }) + return out, nil +} diff --git a/internal/core/onehot_test.go b/internal/core/onehot_test.go new file mode 100644 index 0000000..d8f8986 --- /dev/null +++ b/internal/core/onehot_test.go @@ -0,0 +1,59 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +func TestOneHotEncoding(t *testing.T) { + codes := &Array{shape: []int{4}, dt: Int, ints: []int64{0, 2, 1, 2}} + + hot, err := OneHot(codes, 3) + if err != nil { + t.Fatal(err) + } + if got := hot.Shape(); got[0] != 4 || got[1] != 3 { + t.Fatalf("shape: %v", got) + } + want := [][]float64{{1, 0, 0}, {0, 0, 1}, {0, 1, 0}, {0, 0, 1}} + for i := range 4 { + for j := range 3 { + if g := float64(hot.RawFloat32s()[i*3+j]); g != want[i][j] { + t.Fatalf("hot[%d][%d] = %v, want %v", i, j, g, want[i][j]) + } + } + } + + // A single-code array keeps its leading dimension and gains the + // class axis at the end. + solo, _ := FromInts([]int64{1}, 1) + hotSolo, err := OneHot(solo, 2) + if err != nil { + t.Fatal(err) + } + if got := hotSolo.Shape(); got[0] != 1 || got[1] != 2 { + t.Fatalf("solo shape: %v", got) + } +} + +func TestOneHotErrors(t *testing.T) { + floatCodes, _ := FromFloats([]float64{1}, 1) + if _, err := OneHot(floatCodes, 3); err == nil { + t.Fatal("float codes accepted") + } + + intCodes := &Array{shape: []int{2}, dt: Int, ints: []int64{0, 5}} + if _, err := OneHot(intCodes, 3); err == nil { + t.Fatal("out-of-range code accepted") + } + + negative := &Array{shape: []int{1}, dt: Int, ints: []int64{-1}} + if _, err := OneHot(negative, 3); err == nil { + t.Fatal("negative code accepted") + } + + valid := &Array{shape: []int{1}, dt: Int, ints: []int64{0}} + if _, err := OneHot(valid, 0); err == nil { + t.Fatal("zero classes accepted") + } +} diff --git a/internal/core/ops.go b/internal/core/ops.go new file mode 100644 index 0000000..c2f2d1a --- /dev/null +++ b/internal/core/ops.go @@ -0,0 +1,1094 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "math" + +// Element-wise arithmetic. All operations are +// package-level functions, not methods. Shapes must agree exactly +//: a mismatch is a loud error naming both shapes. The promotion +// ladder is int to float16 to float32 to float64 to complex128: int +// with int stays int (wrapping like Go's int64 arithmetic), float16 +// keeps float16 against int, +// float32 keeps float32 against int or float16, any float64 operand +// promotes, any complex operand promotes. Div is true division (float +// or complex, IEEE behaviour on zero); Quo is integer division and +// rejects both float and complex operands. + +// Add returns the element-wise sum of two arrays of the same shape. +func Add(a, b *Array) (*Array, error) { return elementwiseOp(a, b, "Add", pairAdd) } + +// Sub returns the element-wise difference of two arrays of the same +// shape. +func Sub(a, b *Array) (*Array, error) { return elementwiseOp(a, b, "Sub", pairSub) } + +// Mul returns the element-wise product of two arrays of the same shape. +func Mul(a, b *Array) (*Array, error) { return elementwiseOp(a, b, "Mul", pairMul) } + +// Div returns the element-wise true division of two arrays of the same +// shape. The result is float for real operands and complex when a complex +// operand takes part; division by zero yields ±Inf or NaN per IEEE-754, +// never an error. +func Div(a, b *Array) (*Array, error) { + return elementwiseDiv(a, b) +} + +// Quo returns the element-wise integer division of two int arrays of the +// same shape. A float or complex operand, or a zero divisor element, is +// an error. +func Quo(a, b *Array) (*Array, error) { + if a.dt != Int || b.dt != Int { + return nil, errf("Quo: needs int arrays, got %s and %s", a.dt, b.dt) + } + if !sameShape(a.shape, b.shape) { + return nil, errf("Quo: shape mismatch %s vs %s", shapeText(a.shape), shapeText(b.shape)) + } + ints, _, _, _, _ := a.cloneData() + // The zero check rides the division walk: it visits the divisors in + // the same ascending order the separate scan did, so a zero reports + // the same element index, and a leading zero reports before any + // quotient is written. + for i := range ints { + d := b.ints[i] + if d == 0 { + return nil, errf("Quo: integer division by zero at element %d", i) + } + ints[i] /= d + } + return &Array{shape: a.Shape(), dt: Int, ints: ints}, nil +} + +// scalarIntOp selects the int-class operation a mapKeeping walk applies. +type scalarIntOp uint8 + +const ( + addScalarInt scalarIntOp = iota + subScalarInt + mulScalarInt +) + +// AddI adds an int scalar element-wise; each dtype keeps its kind. +func AddI(a *Array, v int64) *Array { + return a.mapKeeping(v, addScalarInt, + func(x, s float64) float64 { return x + s }, + func(x, s complex128) complex128 { return x + s }) +} + +// SubI subtracts an int scalar element-wise; each dtype keeps its kind. +func SubI(a *Array, v int64) *Array { + return a.mapKeeping(v, subScalarInt, + func(x, s float64) float64 { return x - s }, + func(x, s complex128) complex128 { return x - s }) +} + +// MulI multiplies by an int scalar element-wise; each dtype keeps its +// kind. +func MulI(a *Array, v int64) *Array { + return a.mapKeeping(v, mulScalarInt, + func(x, s float64) float64 { return x * s }, + func(x, s complex128) complex128 { return x * s }) +} + +// AddF adds a float scalar element-wise; integer arrays become +// float64, floating arrays keep their dtype, and complex arrays stay +// complex. +func AddF(a *Array, v float64) *Array { + return a.mapReal(v, + func(x float64) float64 { return x + v }, + func(x complex128) complex128 { return x + complex(v, 0) }) +} + +// SubF subtracts a float scalar element-wise; integer arrays become +// float64, floating arrays keep their dtype, and complex arrays stay +// complex. +func SubF(a *Array, v float64) *Array { + return a.mapReal(v, + func(x float64) float64 { return x - v }, + func(x complex128) complex128 { return x - complex(v, 0) }) +} + +// MulF multiplies by a float scalar element-wise; integer arrays +// become float64, floating arrays keep their dtype, and complex arrays +// stay complex. +func MulF(a *Array, v float64) *Array { + return a.mapReal(v, + func(x float64) float64 { return x * v }, + func(x complex128) complex128 { return x * complex(v, 0) }) +} + +// DivI divides by an int scalar as true division; integer arrays +// become float64, floating arrays keep their dtype, and complex arrays +// stay complex, with IEEE behaviour when v is zero. +func DivI(a *Array, v int64) *Array { + return a.mapReal(float64(v), + func(x float64) float64 { return x / float64(v) }, + func(x complex128) complex128 { return x / complex(float64(v), 0) }) +} + +// DivF divides by a float scalar as true division; integer arrays +// become float64, floating arrays keep their dtype, and complex arrays +// stay complex, with IEEE behaviour when v is zero. +func DivF(a *Array, v float64) *Array { + return a.mapReal(v, + func(x float64) float64 { return x / v }, + func(x complex128) complex128 { return x / complex(v, 0) }) +} + +// AddC adds a complex scalar element-wise; the result is always complex. +func AddC(a *Array, v complex128) *Array { + return a.mapComplex(func(x complex128) complex128 { return x + v }) +} + +// SubC subtracts a complex scalar element-wise; the result is always +// complex. +func SubC(a *Array, v complex128) *Array { + return a.mapComplex(func(x complex128) complex128 { return x - v }) +} + +// MulC multiplies by a complex scalar element-wise; the result is always +// complex. +func MulC(a *Array, v complex128) *Array { + return a.mapComplex(func(x complex128) complex128 { return x * v }) +} + +// DivC divides by a complex scalar element-wise; the result is always +// complex. +func DivC(a *Array, v complex128) *Array { + return a.mapComplex(func(x complex128) complex128 { return x / v }) +} + +// QuoI divides an int array by an int scalar with integer division; a +// float or complex array, or a zero scalar, is an error. +func QuoI(a *Array, v int64) (*Array, error) { + if a.dt != Int { + return nil, errf("QuoI: needs an int array, got %s", a.dt) + } + if v == 0 { + return nil, errf("QuoI: integer division by zero") + } + ints, _, _, _, _ := a.cloneData() + parallelMin(len(ints), elementwiseMinPerWorker, func(s, e int) { + xs := ints[s:e] + for i := range xs { + xs[i] /= v + } + }) + return &Array{shape: a.Shape(), dt: Int, ints: ints}, nil +} + +// pairOp names an element-wise binary operation. The kernels write the +// operation as a direct expression per dtype, so the element loop holds +// an add or a multiply rather than a call through a func value, which +// neither inlines nor lets the compiler keep the payloads in registers +// across the loop. Only the mixed-dtype pair, whose walk pays an +// accessor call per element whatever the body does, keeps the closure +// form below. +type pairOp uint8 + +const ( + pairAdd pairOp = iota + pairSub + pairMul + pairMin + pairMax + pairPow +) + +// complexOK reports whether the operation accepts complex operands. The +// extrema have no ordering to compare, so they reject them; every other +// operation carries them. +func (o pairOp) complexOK() bool { return o != pairMin && o != pairMax } + +// narrowIntClass reports whether dt is one of the small integer +// payloads: the dtypes whose same-width arithmetic runs in the element +// type itself. +func narrowIntClass(dt Dtype) bool { + switch dt { + case Int8, Uint8, Int16, Uint16, Int32, Uint32: + return true + } + return false +} + +// opClosures returns op as the closure triple elementwise takes. It +// serves the mixed-dtype pair, where the two operands need different +// accessors and the walk rebuilds each element anyway; the expression is +// the one the specialised kernels write, so the results agree bit for +// bit. +func opClosures(op pairOp) (func(x, y int64) int64, func(x, y float64) float64, func(x, y complex128) complex128) { + switch op { + case pairAdd: + return func(x, y int64) int64 { return x + y }, + func(x, y float64) float64 { return x + y }, + func(x, y complex128) complex128 { return x + y } + case pairSub: + return func(x, y int64) int64 { return x - y }, + func(x, y float64) float64 { return x - y }, + func(x, y complex128) complex128 { return x - y } + case pairMul: + return func(x, y int64) int64 { return x * y }, + func(x, y float64) float64 { return x * y }, + func(x, y complex128) complex128 { return x * y } + case pairMin: + return func(x, y int64) int64 { return min(x, y) }, + func(x, y float64) float64 { return min(x, y) }, nil + case pairMax: + return func(x, y int64) int64 { return max(x, y) }, + func(x, y float64) float64 { return max(x, y) }, nil + default: // pairPow + return powInt, + func(x, y float64) float64 { return math.Pow(x, y) }, + func(x, y complex128) complex128 { return mathPowComplex(x, y) } + } +} + +// elementwiseOp applies op pairwise under the promotion ladder, with the +// operation written into each dtype's loop as a direct expression. The +// two same-width payloads are read as slices, so the loop is a load, an +// operation and a store per element; a mixed pair and a strided view +// fall back to elementwise, whose accessor walk sees the same values in +// the same order. +func elementwiseOp(a, b *Array, name string, op pairOp) (*Array, error) { + if !sameShape(a.shape, b.shape) { + return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape)) + } + dt := promote(a.dt, b.dt) + if dt == Bool { + // Bool carries no arithmetic: the loud refusal sits here because + // every arithmetic entry (Add, Sub, Mul, Div, Pow, the pairwise + // extrema) funnels through this throat with its own name. + return nil, errf("%s: bool arrays have no arithmetic", name) + } + if dt == Complex && !op.complexOK() { + return nil, errf("%s: complex arrays have no ordering", name) + } + if op == pairPow && intClass(dt) { + // A negative exponent has no integer result: the same refusal + // Pow makes for int operands, asked of the whole integer class + // through intAt's exact widenings. promote answers an + // integer-class dtype exactly when both operands are + // integer-class, so a gate on the narrow widths alone would + // miss the mixed pair whose promoted dtype is Int. + for i := range b.Len() { + if e := b.intAt(i); e < 0 { + return nil, errf("%s: negative exponent %d at element %d has no int result", name, e, i) + } + } + } + if a.dt != dt || b.dt != dt || !a.isContiguous() || !b.isContiguous() { + ii, ff, cc := opClosures(op) + return elementwise(a, b, name, ii, ff, cc) + } + out := &Array{shape: a.Shape(), dt: dt} + out.alloc(a.Len()) + n := a.Len() + switch dt { + case Int8: + narrowPairRun(op, a.i8s, b.i8s, out.i8s) + case Uint8: + narrowPairRun(op, a.u8s, b.u8s, out.u8s) + case Int16: + narrowPairRun(op, a.i16s, b.i16s, out.i16s) + case Uint16: + narrowPairRun(op, a.u16s, b.u16s, out.u16s) + case Int32: + narrowPairRun(op, a.i32s, b.i32s, out.i32s) + case Uint32: + narrowPairRun(op, a.u32s, b.u32s, out.u32s) + case Int: + ai, bi, oi := a.ints, b.ints, out.ints + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + as, bs, os := ai[s:e], bi[s:e], oi[s:e] + switch op { + case pairAdd: + for i := range os { + os[i] = as[i] + bs[i] + } + case pairSub: + for i := range os { + os[i] = as[i] - bs[i] + } + case pairMul: + for i := range os { + os[i] = as[i] * bs[i] + } + case pairMin: + for i := range os { + os[i] = min(as[i], bs[i]) + } + case pairMax: + for i := range os { + os[i] = max(as[i], bs[i]) + } + default: // pairPow + for i := range os { + os[i] = powInt(as[i], bs[i]) + } + } + }) + case Float16: + // Float16 mirrors float32 one step down the ladder: the + // arithmetic runs in float64, where every half value is exact, + // and narrows once per element. + ai, bi, oi := a.halves, b.halves, out.halves + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + as, bs, os := ai[s:e], bi[s:e], oi[s:e] + switch op { + case pairAdd: + for i := range os { + os[i] = HalfFromFloat64(HalfToFloat64(as[i]) + HalfToFloat64(bs[i])) + } + case pairSub: + for i := range os { + os[i] = HalfFromFloat64(HalfToFloat64(as[i]) - HalfToFloat64(bs[i])) + } + case pairMul: + for i := range os { + os[i] = HalfFromFloat64(HalfToFloat64(as[i]) * HalfToFloat64(bs[i])) + } + case pairMin: + for i := range os { + os[i] = HalfFromFloat64(min(HalfToFloat64(as[i]), HalfToFloat64(bs[i]))) + } + case pairMax: + for i := range os { + os[i] = HalfFromFloat64(max(HalfToFloat64(as[i]), HalfToFloat64(bs[i]))) + } + default: // pairPow + for i := range os { + os[i] = HalfFromFloat64(math.Pow(HalfToFloat64(as[i]), HalfToFloat64(bs[i]))) + } + } + }) + case Float32: + // The widening to float64 is exact and the result narrows once, + // so the widened spelling is the accessor expression unchanged. + ai, bi, oi := a.floats32, b.floats32, out.floats32 + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + as, bs, os := ai[s:e], bi[s:e], oi[s:e] + switch op { + case pairAdd: + for i := range os { + os[i] = float32(float64(as[i]) + float64(bs[i])) + } + case pairSub: + for i := range os { + os[i] = float32(float64(as[i]) - float64(bs[i])) + } + case pairMul: + for i := range os { + os[i] = float32(float64(as[i]) * float64(bs[i])) + } + case pairMin: + for i := range os { + os[i] = float32(min(float64(as[i]), float64(bs[i]))) + } + case pairMax: + for i := range os { + os[i] = float32(max(float64(as[i]), float64(bs[i]))) + } + default: // pairPow + for i := range os { + os[i] = float32(math.Pow(float64(as[i]), float64(bs[i]))) + } + } + }) + case Float: + ai, bi, oi := a.floats, b.floats, out.floats + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + as, bs, os := ai[s:e], bi[s:e], oi[s:e] + switch op { + case pairAdd: + for i := range os { + os[i] = as[i] + bs[i] + } + case pairSub: + for i := range os { + os[i] = as[i] - bs[i] + } + case pairMul: + for i := range os { + os[i] = as[i] * bs[i] + } + case pairMin: + for i := range os { + os[i] = min(as[i], bs[i]) + } + case pairMax: + for i := range os { + os[i] = max(as[i], bs[i]) + } + default: // pairPow + for i := range os { + os[i] = math.Pow(as[i], bs[i]) + } + } + }) + default: + ai, bi, oi := a.complexes, b.complexes, out.complexes + parallelMin(n, elementwiseMinPerWorker, func(s, e int) { + as, bs, os := ai[s:e], bi[s:e], oi[s:e] + switch op { + case pairAdd: + for i := range os { + os[i] = as[i] + bs[i] + } + case pairSub: + for i := range os { + os[i] = as[i] - bs[i] + } + case pairMul: + for i := range os { + os[i] = as[i] * bs[i] + } + default: // pairPow, the only complex operation left + for i := range os { + os[i] = mathPowComplex(as[i], bs[i]) + } + } + }) + } + return out, nil +} + +// narrowPairRun applies op to two same-width integer payloads in the +// element type itself: Go's arithmetic wraps on overflow at every width, +// exactly the semantics the int64 kernels carry, and the pow widens +// exactly and narrows on store. +func narrowPairRun[T int8 | uint8 | int16 | uint16 | int32 | uint32](op pairOp, x, y, dst []T) { + parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) { + as, bs, os := x[s:e], y[s:e], dst[s:e] + switch op { + case pairAdd: + for i := range os { + os[i] = as[i] + bs[i] + } + case pairSub: + for i := range os { + os[i] = as[i] - bs[i] + } + case pairMul: + for i := range os { + os[i] = as[i] * bs[i] + } + case pairMin: + for i := range os { + os[i] = min(as[i], bs[i]) + } + case pairMax: + for i := range os { + os[i] = max(as[i], bs[i]) + } + default: // pairPow + for i := range os { + os[i] = T(powInt(int64(as[i]), int64(bs[i]))) + } + } + }) +} + +// elementwise applies an operation pairwise under the promotion ladder. +// A nil cc marks an ordering, which complex operands reject. float16 and +// float32 operands compute in float64 and round once: their arithmetic +// is exact in float64. The operation arrives as closures, which +// the mixed-dtype pair of elementwiseOp and the extrema family use. +func elementwise(a, b *Array, name string, + ii func(x, y int64) int64, ff func(x, y float64) float64, cc func(x, y complex128) complex128, +) (*Array, error) { + if !sameShape(a.shape, b.shape) { + return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape)) + } + dt := promote(a.dt, b.dt) + if dt == Bool { + return nil, errf("%s: bool arrays have no arithmetic", name) + } + if dt == Complex && cc == nil { + return nil, errf("%s: complex arrays have no ordering", name) + } + out := &Array{shape: a.Shape(), dt: dt} + out.alloc(a.Len()) + // A stride table is a mechanism no public constructor sets, but the + // accessor walk is the one reader that resolves it, so every + // raw-payload branch below is gated on both operands being dense. + dense := a.isContiguous() && b.isContiguous() + switch dt { + case Int8: + narrowElementwise(a, b, out, func(x *Array) []int8 { return x.i8s }, dense, ii) + case Uint8: + narrowElementwise(a, b, out, func(x *Array) []uint8 { return x.u8s }, dense, ii) + case Int16: + narrowElementwise(a, b, out, func(x *Array) []int16 { return x.i16s }, dense, ii) + case Uint16: + narrowElementwise(a, b, out, func(x *Array) []uint16 { return x.u16s }, dense, ii) + case Int32: + narrowElementwise(a, b, out, func(x *Array) []int32 { return x.i32s }, dense, ii) + case Uint32: + narrowElementwise(a, b, out, func(x *Array) []uint32 { return x.u32s }, dense, ii) + case Int: + // Dense same-width operands stream their raw payloads; a mixed + // integer pair promoted to Int (int32 with uint32, any narrow + // operand with Int) reads through intAt, which widens every + // integer-class payload exactly. This mirrors the float32 + // branch's operand gates: the raw read is only sound when the + // operand really carries the Int payload. + aI, bI := a.dt == Int, b.dt == Int + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.ints[s:e] + switch { + case dense && aI && bI: + as, bs := a.ints[s:e], b.ints[s:e] + for i := range os { + os[i] = ii(as[i], bs[i]) + } + case dense && aI: + as := a.ints[s:e] + for i := range os { + os[i] = ii(as[i], b.intAt(s+i)) + } + case dense && bI: + bs := b.ints[s:e] + for i := range os { + os[i] = ii(a.intAt(s+i), bs[i]) + } + default: + for i := s; i < e; i++ { + os[i-s] = ii(a.intAt(i), b.intAt(i)) + } + } + }) + case Float16: + // Float16 mirrors float32 one step down the ladder: the + // arithmetic runs in float64, where every half value is exact, + // and narrows once per element. Dense same-width operands stream + // their raw payloads; everything strided or mixed keeps the + // accessor fallback. + aH, bH := a.dt == Float16, b.dt == Float16 + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.halves[s:e] + switch { + case dense && aH && bH: + as, bs := a.halves[s:e], b.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(ff(HalfToFloat64(as[i]), HalfToFloat64(bs[i]))) + } + case dense && aH: + as := a.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(ff(HalfToFloat64(as[i]), b.floatAt(s+i))) + } + case dense && bH: + bs := b.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(ff(a.floatAt(s+i), HalfToFloat64(bs[i]))) + } + default: + for i := s; i < e; i++ { + os[i-s] = HalfFromFloat64(ff(a.floatAt(i), b.floatAt(i))) + } + } + }) + case Float32: + // Dense same-width operands stream their raw payloads: float32 + // widens exactly, so the values and their order match accessor + // reads bit for bit. Everything strided or mixed keeps the + // accessor fallback. + a32, b32 := a.dt == Float32, b.dt == Float32 + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.floats32[s:e] + switch { + case dense && a32 && b32: + as, bs := a.floats32[s:e], b.floats32[s:e] + for i := range os { + os[i] = float32(ff(float64(as[i]), float64(bs[i]))) + } + case dense && a32: + as := a.floats32[s:e] + for i := range os { + os[i] = float32(ff(float64(as[i]), b.floatAt(s+i))) + } + case dense && b32: + bs := b.floats32[s:e] + for i := range os { + os[i] = float32(ff(a.floatAt(s+i), float64(bs[i]))) + } + default: + for i := s; i < e; i++ { + os[i-s] = float32(ff(a.floatAt(i), b.floatAt(i))) + } + } + }) + case Float: + aF, bF := a.dt == Float, b.dt == Float + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.floats[s:e] + switch { + case dense && aF && bF: + as, bs := a.floats[s:e], b.floats[s:e] + for i := range os { + os[i] = ff(as[i], bs[i]) + } + case dense && aF: + as := a.floats[s:e] + for i := range os { + os[i] = ff(as[i], b.floatAt(s+i)) + } + case dense && bF: + bs := b.floats[s:e] + for i := range os { + os[i] = ff(a.floatAt(s+i), bs[i]) + } + default: + for i := s; i < e; i++ { + os[i-s] = ff(a.floatAt(i), b.floatAt(i)) + } + } + }) + default: + aC, bC := a.dt == Complex, b.dt == Complex + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.complexes[s:e] + switch { + case dense && aC && bC: + as, bs := a.complexes[s:e], b.complexes[s:e] + for i := range os { + os[i] = cc(as[i], bs[i]) + } + case dense && aC: + as := a.complexes[s:e] + for i := range os { + os[i] = cc(as[i], b.complexAt(s+i)) + } + case dense && bC: + bs := b.complexes[s:e] + for i := range os { + os[i] = cc(a.complexAt(s+i), bs[i]) + } + default: + for i := s; i < e; i++ { + os[i-s] = cc(a.complexAt(i), b.complexAt(i)) + } + } + }) + } + return out, nil +} + +// narrowElementwise is elementwise's walk for a narrow integer promoted +// result: dense same-width operands stream their own payloads, every +// mixed or strided read goes through intAt, and the int64 closure value +// narrows on store with Go's conversion, the implicit-store cast +// setConverted performs. +func narrowElementwise[T int8 | uint8 | int16 | uint16 | int32 | uint32]( + a, b, out *Array, payload func(*Array) []T, dense bool, ii func(x, y int64) int64, +) { + aN, bN := a.dt == out.dt, b.dt == out.dt + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := payload(out)[s:e] + switch { + case dense && aN && bN: + as, bs := payload(a)[s:e], payload(b)[s:e] + for i := range os { + os[i] = T(ii(int64(as[i]), int64(bs[i]))) + } + case dense && aN: + as := payload(a)[s:e] + for i := range os { + os[i] = T(ii(int64(as[i]), b.intAt(s+i))) + } + case dense && bN: + bs := payload(b)[s:e] + for i := range os { + os[i] = T(ii(a.intAt(s+i), int64(bs[i]))) + } + default: + for i := s; i < e; i++ { + os[i-s] = T(ii(a.intAt(i), b.intAt(i))) + } + } + }) +} + +// elementwiseDiv applies true division; the result follows the promotion +// ladder, except that every integer-class pair divides into float64. +func elementwiseDiv(a, b *Array) (*Array, error) { + if !sameShape(a.shape, b.shape) { + return nil, errf("Div: shape mismatch %s vs %s", shapeText(a.shape), shapeText(b.shape)) + } + dt := promote(a.dt, b.dt) + if dt == Bool { + return nil, errf("Div: bool arrays have no arithmetic") + } + if intClass(dt) { + // Integer true division has always answered float64 for an int + // pair; bool and the narrow integer widths take the same route, + // their widenings exact and the division run in float64. + dt = Float + } + out := &Array{shape: a.Shape(), dt: dt} + out.alloc(a.Len()) + // The same contiguity gate the elementwise fallback carries: a + // strided operand reads through the accessors, never in payload + // order. + dense := a.isContiguous() && b.isContiguous() + switch dt { + case Float16: + // Dense same-width float16 operands stream their raw payloads; + // the division runs in float64 and narrows once per element, the + // mirror of the float32 path below. + aH, bH := a.dt == Float16, b.dt == Float16 + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.halves[s:e] + switch { + case dense && aH && bH: + as, bs := a.halves[s:e], b.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(HalfToFloat64(as[i]) / HalfToFloat64(bs[i])) + } + case dense && aH: + as := a.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(HalfToFloat64(as[i]) / b.floatAt(s+i)) + } + case dense && bH: + bs := b.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(a.floatAt(s+i) / HalfToFloat64(bs[i])) + } + default: + for i := s; i < e; i++ { + os[i-s] = HalfFromFloat64(a.floatAt(i) / b.floatAt(i)) + } + } + }) + case Float32: + // Dense same-width float32 operands stream their raw payloads; + // the widening is exact, so nothing shifts a bit. Everything + // strided or mixed keeps the accessor fallback. + a32, b32 := a.dt == Float32, b.dt == Float32 + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.floats32[s:e] + switch { + case dense && a32 && b32: + as, bs := a.floats32[s:e], b.floats32[s:e] + for i := range os { + os[i] = float32(float64(as[i]) / float64(bs[i])) + } + case dense && a32: + as := a.floats32[s:e] + for i := range os { + os[i] = float32(float64(as[i]) / b.floatAt(s+i)) + } + case dense && b32: + bs := b.floats32[s:e] + for i := range os { + os[i] = float32(a.floatAt(s+i) / float64(bs[i])) + } + default: + for i := s; i < e; i++ { + os[i-s] = float32(a.floatAt(i) / b.floatAt(i)) + } + } + }) + case Float: + aF, bF := a.dt == Float, b.dt == Float + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.floats[s:e] + switch { + case dense && aF && bF: + as, bs := a.floats[s:e], b.floats[s:e] + for i := range os { + os[i] = as[i] / bs[i] + } + case dense && aF: + as := a.floats[s:e] + for i := range os { + os[i] = as[i] / b.floatAt(s+i) + } + case dense && bF: + bs := b.floats[s:e] + for i := range os { + os[i] = a.floatAt(s+i) / bs[i] + } + default: + for i := s; i < e; i++ { + os[i-s] = a.floatAt(i) / b.floatAt(i) + } + } + }) + default: + aC, bC := a.dt == Complex, b.dt == Complex + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.complexes[s:e] + switch { + case dense && aC && bC: + as, bs := a.complexes[s:e], b.complexes[s:e] + for i := range os { + os[i] = as[i] / bs[i] + } + case dense && aC: + as := a.complexes[s:e] + for i := range os { + os[i] = as[i] / b.complexAt(s+i) + } + case dense && bC: + bs := b.complexes[s:e] + for i := range os { + os[i] = a.complexAt(s+i) / bs[i] + } + default: + for i := s; i < e; i++ { + os[i-s] = a.complexAt(i) / b.complexAt(i) + } + } + }) + } + return out, nil +} + +// mapKeeping applies a scalar operation that keeps each real dtype and +// stays complex on complex arrays. The int-class operation arrives as an +// enum, so every element loop below holds the operation itself rather +// than a call through a func value. +func (a *Array) mapKeeping(v int64, op scalarIntOp, + ff func(x, s float64) float64, cc func(x, s complex128) complex128, +) *Array { + if a.dt == Bool { + // Bool carries no arithmetic, and this keeping-kind surface has + // no error channel; nil is the refusal shape the no-error + // constructors already use. + return nil + } + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(a.Len()) + switch a.dt { + case Int: + src, dst := a.ints, out.ints + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + xs, os := src[s:e], dst[s:e] + switch op { + case addScalarInt: + for i := range os { + os[i] = xs[i] + v + } + case subScalarInt: + for i := range os { + os[i] = xs[i] - v + } + default: // mulScalarInt + for i := range os { + os[i] = xs[i] * v + } + } + }) + case Int8, Uint8, Int16, Uint16, Int32, Uint32: + // The keeping-kind walk for the narrow integer widths: the int64 + // operation narrows on store with Go's conversion, the same + // implicit-store cast setConverted performs. + switch a.dt { + case Int8: + narrowMapKeeping(a, out, v, op, func(x *Array) []int8 { return x.i8s }) + case Uint8: + narrowMapKeeping(a, out, v, op, func(x *Array) []uint8 { return x.u8s }) + case Int16: + narrowMapKeeping(a, out, v, op, func(x *Array) []int16 { return x.i16s }) + case Uint16: + narrowMapKeeping(a, out, v, op, func(x *Array) []uint16 { return x.u16s }) + case Int32: + narrowMapKeeping(a, out, v, op, func(x *Array) []int32 { return x.i32s }) + default: + narrowMapKeeping(a, out, v, op, func(x *Array) []uint32 { return x.u32s }) + } + case Float16: + fv := float64(v) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(ff(HalfToFloat64(as[i]), fv)) + } + }) + case Float32: + fv := float64(v) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = float32(ff(float64(as[i]), fv)) + } + }) + case Float: + fv := float64(v) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats[s:e], out.floats[s:e] + for i := range os { + os[i] = ff(as[i], fv) + } + }) + default: + cv := complex(float64(v), 0) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.complexes[s:e], out.complexes[s:e] + for i := range os { + os[i] = cc(as[i], cv) + } + }) + } + return out +} + +// narrowMapKeeping is mapKeeping's keeping-kind walk for a narrow +// integer payload: each element computes in int64 from the exact +// widening and narrows on store, the operation written into the loop. +func narrowMapKeeping[T int8 | uint8 | int16 | uint16 | int32 | uint32]( + a, out *Array, v int64, op scalarIntOp, pick func(*Array) []T, +) { + src, dst := pick(a), pick(out) + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + xs, os := src[s:e], dst[s:e] + switch op { + case addScalarInt: + for i := range os { + os[i] = T(int64(xs[i]) + v) + } + case subScalarInt: + for i := range os { + os[i] = T(int64(xs[i]) - v) + } + default: // mulScalarInt + for i := range os { + os[i] = T(int64(xs[i]) * v) + } + } + }) +} + +// mapReal applies a scalar operation: integer arrays widen to float64, +// float16 computes in float64 and narrows, float32 arrays keep +// float32 (the scalar is weak and narrows), float64 stays +// float64, complex stays complex. +func (a *Array) mapReal(v float64, f func(float64) float64, c func(complex128) complex128) *Array { + if a.dt == Complex { + return a.mapComplex(c) + } + out := &Array{shape: a.Shape(), dt: a.dt} + if intClass(a.dt) { + // Every integer-class array, bool and the narrow widths included, + // widens to float64 under this surface's contract; the accessor + // walk below reads each widening exactly. + out.dt = Float + } + out.alloc(a.Len()) + if a.dt == Float16 { + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.halves[s:e], out.halves[s:e] + for i := range os { + os[i] = HalfFromFloat64(f(HalfToFloat64(as[i]))) + } + }) + return out + } + if a.dt == Float32 { + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + as, os := a.floats32[s:e], out.floats32[s:e] + for i := range os { + os[i] = float32(f(float64(as[i]))) + } + }) + return out + } + // A float64 payload reads directly; the integer-class payloads read + // their own slices when dense, each widening exactly the one the + // accessor hands over; an integer-class view keeps the accessor + // walk. + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.floats[s:e] + switch { + case a.dt == Float: + as := a.floats[s:e] + for i := range os { + os[i] = f(as[i]) + } + case a.strides != nil: + for i := s; i < e; i++ { + os[i-s] = f(a.floatAt(i)) + } + default: + intRealRun(a, os, s, f) + } + }) + return out +} + +// intRealRun is mapReal's widening walk for a dense integer-class +// payload: the dtype dispatch sits outside the element loop and the +// float64 operation f reads each exact widening. +func intRealRun(a *Array, os []float64, s int, f func(float64) float64) { + switch a.dt { + case Int: + as := a.ints[s : s+len(os)] + for i := range os { + os[i] = f(float64(as[i])) + } + case Bool: + as := a.bools[s : s+len(os)] + for i := range os { + w := 0.0 + if as[i] { + w = 1 + } + os[i] = f(w) + } + case Int8: + as := a.i8s[s : s+len(os)] + for i := range os { + os[i] = f(float64(as[i])) + } + case Uint8: + as := a.u8s[s : s+len(os)] + for i := range os { + os[i] = f(float64(as[i])) + } + case Int16: + as := a.i16s[s : s+len(os)] + for i := range os { + os[i] = f(float64(as[i])) + } + case Uint16: + as := a.u16s[s : s+len(os)] + for i := range os { + os[i] = f(float64(as[i])) + } + case Int32: + as := a.i32s[s : s+len(os)] + for i := range os { + os[i] = f(float64(as[i])) + } + default: // Uint32 + as := a.u32s[s : s+len(os)] + for i := range os { + os[i] = f(float64(as[i])) + } + } +} + +// mapComplex applies an operation whose result is always complex. +func (a *Array) mapComplex(c func(complex128) complex128) *Array { + out := &Array{shape: a.Shape(), dt: Complex} + out.complexes = make([]complex128, a.Len()) + if a.dt == Complex { + ac := a.complexes + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + os := out.complexes[s:e] + as := ac[s:e] + for i := range os { + os[i] = c(as[i]) + } + }) + return out + } + parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out.complexes[i] = c(a.complexAt(i)) + } + }) + return out +} diff --git a/internal/core/ops_test.go b/internal/core/ops_test.go new file mode 100644 index 0000000..f0ce119 --- /dev/null +++ b/internal/core/ops_test.go @@ -0,0 +1,420 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "slices" + "strings" + "testing" +) + +func intsOf(t *testing.T, a *Array) []int64 { + t.Helper() + out := make([]int64, a.Len()) + for i := range out { + out[i], _ = IntAt(a, i) + } + return out +} + +func floatsOf(t *testing.T, a *Array) []float64 { + t.Helper() + out := make([]float64, a.Len()) + for i := range out { + out[i], _ = FloatAt(a, i) + } + return out +} + +func TestElementwiseIntStaysInt(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3}, 3) + b := mustFromInts(t, []int64{4, 5, 6}, 3) + + sum, err := Add(a, b) + if err != nil { + t.Fatalf("Add: %v", err) + } + if sum.Dtype() != Int { + t.Fatalf("int + int must stay int, got %s", sum.Dtype()) + } + want := []int64{5, 7, 9} + got := intsOf(t, sum) + for i := range want { + if got[i] != want[i] { + t.Fatalf("Add: %v", got) + } + } + + diff, _ := Sub(b, a) + if diff.Dtype() != Int || intsOf(t, diff)[0] != 3 { + t.Fatalf("Sub: %s %v", diff, intsOf(t, diff)) + } + + prod, _ := Mul(a, b) + if prod.Dtype() != Int || intsOf(t, prod)[2] != 18 { + t.Fatalf("Mul: %s %v", prod, intsOf(t, prod)) + } + + // int arithmetic wraps like Go's int64. + big := mustFromInts(t, []int64{math.MaxInt64}, 1) + one := mustFromInts(t, []int64{1}, 1) + wrapped, _ := Add(big, one) + if v, _ := IntAt(wrapped, 0); v != math.MinInt64 { + t.Fatalf("wrap: %d", v) + } +} + +func TestElementwisePromotesToFloat(t *testing.T) { + i := mustFromInts(t, []int64{1, 2}, 2) + f := mustFromFloats(t, []float64{0.5, 1.5}, 2) + + sum, err := Add(i, f) + if err != nil { + t.Fatalf("Add: %v", err) + } + if sum.Dtype() != Float { + t.Fatalf("int + float must promote, got %s", sum.Dtype()) + } + got := floatsOf(t, sum) + if got[0] != 1.5 || got[1] != 3.5 { + t.Fatalf("Add promoted: %v", got) + } +} + +func TestDivIsTrueDivision(t *testing.T) { + a := mustFromInts(t, []int64{1, 7, -7}, 3) + b := mustFromInts(t, []int64{2, 2, 2}, 3) + + q, err := Div(a, b) + if err != nil { + t.Fatalf("Div: %v", err) + } + if q.Dtype() != Float { + t.Fatalf("Div must always yield float, got %s", q.Dtype()) + } + got := floatsOf(t, q) + if got[0] != 0.5 || got[1] != 3.5 || got[2] != -3.5 { + t.Fatalf("Div: %v", got) + } + + // Zero divisors are IEEE, never errors. + num := mustFromFloats(t, []float64{1.0, -1.0, 0.0}, 3) + zero := mustFromFloats(t, []float64{0.0, 0.0, 0.0}, 3) + zq, err := Div(num, zero) + if err != nil { + t.Fatalf("Div by zero: %v", err) + } + z := floatsOf(t, zq) + if !math.IsInf(z[0], 1) || !math.IsInf(z[1], -1) || !math.IsNaN(z[2]) { + t.Fatalf("IEEE: %v", z) + } +} + +func TestQuo(t *testing.T) { + a := mustFromInts(t, []int64{7, -7, 9}, 3) + b := mustFromInts(t, []int64{2, 2, 3}, 3) + + q, err := Quo(a, b) + if err != nil { + t.Fatalf("Quo: %v", err) + } + got := intsOf(t, q) + if got[0] != 3 || got[1] != -3 || got[2] != 3 { + t.Fatalf("Quo: %v", got) + } + + f := mustFromFloats(t, []float64{1}, 1) + if _, err := Quo(f, f); err == nil || !strings.Contains(err.Error(), "needs int arrays") { + t.Fatalf("Quo float: %v", err) + } + + z := mustFromInts(t, []int64{1, 0}, 2) + o := mustFromInts(t, []int64{1, 1}, 2) + if _, err := Quo(z, o); err != nil { + t.Fatalf("Quo zeros in dividend is fine: %v", err) + } + if _, err := Quo(o, z); err == nil || !strings.Contains(err.Error(), "division by zero") { + t.Fatalf("Quo zero divisor: %v", err) + } +} + +func TestShapeMismatchIsLoud(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3}, 3) + b := mustFromInts(t, []int64{1, 2}, 2) + _, err := Add(a, b) + if err == nil || !strings.Contains(err.Error(), "shape mismatch (3) vs (2)") { + t.Fatalf("shape mismatch: %v", err) + } + if _, err := Div(a, b); err == nil || !strings.Contains(err.Error(), "shape mismatch") { + t.Fatalf("Div shape mismatch: %v", err) + } + if _, err := Quo(a, b); err == nil || !strings.Contains(err.Error(), "shape mismatch") { + t.Fatalf("Quo shape mismatch: %v", err) + } +} + +func TestScalarOps(t *testing.T) { + a := mustFromInts(t, []int64{1, 2}, 2) + + if v, _ := IntAt(AddI(a, 5), 0); v != 6 { + t.Fatalf("AddI: %d", v) + } + if v, _ := IntAt(SubI(a, 1), 1); v != 1 { + t.Fatalf("SubI: %d", v) + } + if v, _ := IntAt(MulI(a, 3), 1); v != 6 { + t.Fatalf("MulI: %d", v) + } + + f := mustFromFloats(t, []float64{1.0, 2.0}, 2) + if v, _ := FloatAt(AddI(f, 5), 0); v != 6.0 { + t.Fatalf("AddI on float: %v", v) + } + if v, _ := FloatAt(AddF(a, 0.5), 0); v != 1.5 { + t.Fatalf("AddF: %v", v) + } + if v, _ := FloatAt(SubF(a, 0.5), 1); v != 1.5 { + t.Fatalf("SubF: %v", v) + } + if v, _ := FloatAt(MulF(a, 1.5), 1); v != 3.0 { + t.Fatalf("MulF: %v", v) + } + + // Scalar division is true division regardless of the scalar flavour. + div := DivI(a, 2) + if div.Dtype() != Float || floatsOf(t, div)[0] != 0.5 { + t.Fatalf("DivI: %s %v", div.Dtype(), floatsOf(t, div)) + } + if v, _ := FloatAt(DivF(a, 4), 1); v != 0.5 { + t.Fatalf("DivF: %v", floatsOf(t, DivF(a, 4))) + } + + q, err := QuoI(a, 2) + if err != nil { + t.Fatalf("QuoI: %v", err) + } + if v, _ := IntAt(q, 0); v != 0 { + t.Fatalf("QuoI: %d", v) + } + if _, err := QuoI(a, 0); err == nil || !strings.Contains(err.Error(), "division by zero") { + t.Fatalf("QuoI zero: %v", err) + } + if _, err := QuoI(f, 2); err == nil || !strings.Contains(err.Error(), "needs an int array") { + t.Fatalf("QuoI float: %v", err) + } +} + +// narrowBools builds a bool array for the narrow-dtype contract tests. +func narrowBools(t *testing.T, vals []bool, shape ...int) *Array { + t.Helper() + a, err := FromBools(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +// narrowInt8s builds an int8 array for the narrow-dtype contract tests. +func narrowInt8s(t *testing.T, vals []int8, shape ...int) *Array { + t.Helper() + a, err := FromInt8s(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +// narrowUint8s builds a uint8 array for the narrow-dtype contract tests. +func narrowUint8s(t *testing.T, vals []uint8, shape ...int) *Array { + t.Helper() + a, err := FromUint8s(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +// narrowInt16s builds an int16 array for the narrow-dtype contract tests. +func narrowInt16s(t *testing.T, vals []int16, shape ...int) *Array { + t.Helper() + a, err := FromInt16s(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +// narrowUint32s builds a uint32 array for the narrow-dtype contract tests. +func narrowUint32s(t *testing.T, vals []uint32, shape ...int) *Array { + t.Helper() + a, err := FromUint32s(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +// TestBoolArithmeticIsALoudError pins the contract that bool arrays +// carry no arithmetic: every element-wise arithmetic entry whose +// promoted dtype is bool answers the named refusal, while comparisons +// over bool keep the Int mask. +func TestBoolArithmeticIsALoudError(t *testing.T) { + x := narrowBools(t, []bool{true, true, false}, 3) + y := narrowBools(t, []bool{true, false, false}, 3) + entries := []struct { + name string + fn func() error + }{ + {"Add", func() error { _, err := Add(x, y); return err }}, + {"Sub", func() error { _, err := Sub(x, y); return err }}, + {"Mul", func() error { _, err := Mul(x, y); return err }}, + {"Div", func() error { _, err := Div(x, y); return err }}, + {"Pow", func() error { _, err := Pow(x, y); return err }}, + {"Minimum", func() error { _, err := Minimum(x, y); return err }}, + {"Maximum", func() error { _, err := Maximum(x, y); return err }}, + {"Dot", func() error { _, err := Dot(x, y); return err }}, + } + for _, tc := range entries { + err := tc.fn() + if err == nil || !strings.Contains(err.Error(), "bool arrays have no arithmetic") { + t.Errorf("%s on bool arrays: %v, want the named arithmetic refusal", tc.name, err) + } + } + // Comparisons are not arithmetic: the mask answers bool, true where + // the relation holds. + eq, err := Eq(x, y) + if err != nil { + t.Fatalf("Eq on bool: %v", err) + } + if eq.Dtype() != Bool { + t.Fatalf("Eq on bool answered dtype %s, want the bool mask", eq.Dtype()) + } + if want := []bool{true, false, true}; !slices.Equal(eq.RawBools(), want) { + t.Fatalf("Eq on bool = %v, want %v", eq.RawBools(), want) + } + // A mixed bool pair promotes into the other operand's dtype, where + // arithmetic is defined on the widened values. + i8 := narrowInt8s(t, []int8{2, 3, 4}, 3) + sum, err := Add(x, i8) + if err != nil { + t.Fatalf("Add bool with int8: %v", err) + } + if sum.Dtype() != Int8 { + t.Fatalf("Add bool with int8 answered %s, want int8", sum.Dtype()) + } + if want := []int8{3, 4, 4}; !slices.Equal(sum.RawInt8s(), want) { + t.Fatalf("Add bool with int8 = %v, want %v", sum.RawInt8s(), want) + } +} + +// TestNarrowIntegerElementwiseContract pins the narrow element-type +// loops: same-width dense pairs wrap natively in their own type, mixed +// signedness promotes through the containment table, integer true +// division answers float64 exactly as the int path does, and the scalar +// maps keep their kind with the implicit-store cast. +func TestNarrowIntegerElementwiseContract(t *testing.T) { + a8 := narrowInt8s(t, []int8{100, -128, 5}, 3) + b8 := narrowInt8s(t, []int8{100, 2, 3}, 3) + + s, err := Add(a8, b8) + if err != nil { + t.Fatalf("Add int8: %v", err) + } + if s.Dtype() != Int8 { + t.Fatalf("Add int8 answered %s, want int8", s.Dtype()) + } + if want := []int8{-56, -126, 8}; !slices.Equal(s.RawInt8s(), want) { + t.Fatalf("Add int8 wrap = %v, want %v", s.RawInt8s(), want) + } + m, err := Mul(a8, b8) + if err != nil { + t.Fatalf("Mul int8: %v", err) + } + if want := []int8{16, 0, 15}; !slices.Equal(m.RawInt8s(), want) { + t.Fatalf("Mul int8 wrap = %v, want %v", m.RawInt8s(), want) + } + mx, err := Maximum(a8, b8) + if err != nil { + t.Fatalf("Maximum int8: %v", err) + } + if want := []int8{100, 2, 5}; !slices.Equal(mx.RawInt8s(), want) { + t.Fatalf("Maximum int8 = %v, want %v", mx.RawInt8s(), want) + } + + // Mixed signedness: int8 with uint8 promotes to int16, where both + // value sets fit. + u8 := narrowUint8s(t, []uint8{200, 200, 200}, 3) + mix, err := Add(a8, u8) + if err != nil { + t.Fatalf("Add int8 with uint8: %v", err) + } + if mix.Dtype() != Int16 { + t.Fatalf("Add int8 with uint8 answered %s, want int16", mix.Dtype()) + } + if want := []int16{300, 72, 205}; !slices.Equal(mix.RawInt16s(), want) { + t.Fatalf("Add int8 with uint8 = %v, want %v", mix.RawInt16s(), want) + } + + // True division: an integer-class pair answers float64, the route + // the int pair has always taken. + d, err := Div(a8, b8) + if err != nil { + t.Fatalf("Div int8: %v", err) + } + if d.Dtype() != Float { + t.Fatalf("Div int8 answered %s, want float", d.Dtype()) + } + got := d.RawFloats() + if got[0] != 1 || got[1] != -64 || math.Abs(got[2]-5.0/3.0) > 1e-12 { + t.Fatalf("Div int8 = %v, want [1 -64 1.666...]", got) + } + + // Pow keeps the int contract's negative-exponent refusal on the + // narrow widths. + neg := narrowInt8s(t, []int8{-1}, 1) + pos := narrowInt8s(t, []int8{2}, 1) + if _, err := Pow(pos, neg); err == nil || !strings.Contains(err.Error(), "negative exponent") { + t.Fatalf("Pow int8 with a negative exponent: %v", err) + } + // The mixed pair whose promote() is Int: the exponent scan must + // refuse the negative narrow exponent instead of letting powInt's + // loop answer a silent 1. + pbase, perr := FromInts([]int64{2, 3}, 2) + if perr != nil { + t.Fatalf("FromInts: %v", perr) + } + if _, err := Pow(pbase, narrowInt8s(t, []int8{-3, 4}, 2)); err == nil || + !strings.Contains(err.Error(), "negative exponent") { + t.Fatalf("Pow int with a negative int8 exponent: %v", err) + } + if _, err := Pow(narrowUint32s(t, []uint32{2, 3}, 2), narrowInt16s(t, []int16{-1, 2}, 2)); err == nil || + !strings.Contains(err.Error(), "negative exponent") { + t.Fatalf("Pow uint32 with a negative int16 exponent: %v", err) + } + + // The scalar maps: AddI keeps the kind with the implicit-store cast, + // AddF widens the whole integer class to float64, and AddI on a bool + // array has no error channel, so it answers nil. + si := AddI(a8, 200) + if si == nil || si.Dtype() != Int8 { + t.Fatalf("AddI int8: %v %v", si, err) + } + // The int64 sum narrows on store: int8(300) = 44, -128+200 = 72, + // int8(205) = -51. + if want := []int8{44, 72, -51}; !slices.Equal(si.RawInt8s(), want) { + t.Fatalf("AddI int8 = %v, want %v", si.RawInt8s(), want) + } + sf := AddF(a8, 0.5) + if sf.Dtype() != Float { + t.Fatalf("AddF int8 answered %s, want float", sf.Dtype()) + } + if want := []float64{100.5, -127.5, 5.5}; !slices.Equal(sf.RawFloats(), want) { + t.Fatalf("AddF int8 = %v, want %v", sf.RawFloats(), want) + } + bl := narrowBools(t, []bool{true, false}, 2) + if got := AddI(bl, 1); got != nil { + t.Fatalf("AddI on a bool array = %v, want the nil refusal", got) + } +} diff --git a/internal/core/orthopoly.go b/internal/core/orthopoly.go new file mode 100644 index 0000000..c3782cc --- /dev/null +++ b/internal/core/orthopoly.go @@ -0,0 +1,312 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Classical orthogonal polynomials on the real line, each element-wise +// over real arrays with integer degree or order. The recurrences are +// the stable ones for every family, and the normalisations follow the +// standard physics conventions: Hermite physicists' Hₙ with leading +// coefficient 2ⁿ, Laguerre Lₙ (α = 0) and generalised Lₙ^α, Chebyshev +// of the first and second kind Tₙ and Uₙ on [−1, 1]. + +// Hermite returns the physicists' Hermite polynomial Hₙ(x) of each +// element, by the recurrence Hₙ₊₁ = 2x·Hₙ − 2n·Hₙ₋₁. +func Hermite(n int, x *Array) (*Array, error) { + if n < 0 { + return nil, errf("Hermite: degree must be ≥ 0, got %d", n) + } + if x.dt == Complex { + return nil, errf("Hermite: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv := x.floatAt(i) + h0, h1 := 1.0, 2*xv + for k := 1; k < n; k++ { + h0, h1 = h1, 2*xv*h1-2*float64(k)*h0 + } + if n == 0 { + h1 = 1 + } + out.floats[i] = h1 + } + }) + return out, nil +} + +// Laguerre returns the generalised Laguerre polynomial Lₙ^α(x) of each +// element, by the recurrence (n+1)Lₙ₊₁ = (2n+1+α−x)Lₙ − (n+α)Lₙ₋₁. +// The degree n must be ≥ 0. +func Laguerre(n int, alpha float64, x *Array) (*Array, error) { + if n < 0 { + return nil, errf("Laguerre: degree must be ≥ 0, got %d", n) + } + if x.dt == Complex { + return nil, errf("Laguerre: complex arrays are not supported") + } + if n == 0 { + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + for i := range x.Len() { + out.floats[i] = 1 + } + return out, nil + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv := x.floatAt(i) + l0, l1 := 1.0, 1.0+float64(alpha)-xv + for k := 1; k < n; k++ { + kf := float64(k) + l0, l1 = l1, ((2*kf+1+alpha-xv)*l1-(kf+alpha)*l0)/(kf+1) + } + out.floats[i] = l1 + } + }) + return out, nil +} + +// ChebyshevT returns the Chebyshev polynomial of the first kind +// Tₙ(x) = cos(n·arccos x) of each element, by the recurrence +// Tₙ₊₁ = 2x·Tₙ − Tₙ₋₁. +func ChebyshevT(n int, x *Array) (*Array, error) { + if n < 0 { + return nil, errf("ChebyshevT: degree must be ≥ 0, got %d", n) + } + if x.dt == Complex { + return nil, errf("ChebyshevT: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv := x.floatAt(i) + t0, t1 := 1.0, xv + for k := 1; k < n; k++ { + t0, t1 = t1, 2*xv*t1-t0 + } + if n == 0 { + t1 = 1 + } + out.floats[i] = t1 + } + }) + return out, nil +} + +// ChebyshevU returns the Chebyshev polynomial of the second kind +// Uₙ(x) of each element, by the recurrence Uₙ₊₁ = 2x·Uₙ − Uₙ₋₁ with +// U₀ = 1, U₁ = 2x. +func ChebyshevU(n int, x *Array) (*Array, error) { + if n < 0 { + return nil, errf("ChebyshevU: degree must be ≥ 0, got %d", n) + } + if x.dt == Complex { + return nil, errf("ChebyshevU: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv := x.floatAt(i) + u0, u1 := 1.0, 2*xv + for k := 1; k < n; k++ { + u0, u1 = u1, 2*xv*u1-u0 + } + if n == 0 { + u1 = 1 + } + out.floats[i] = u1 + } + }) + return out, nil +} + +// SphericalHarmonicReal returns the real spherical harmonic +// Y_l^m(θ, φ) built from the complex one, both branches under the +// standard convention Y_real = √2·(−1)^m·Re(Y_l^m) for m > 0 and +// Y_real = √2·(−1)^m·Im(Y_l^|m|) for m < 0, and for m = 0 the plain +// Y_l⁰. Values are real and element-wise over theta and phi, which +// must share a shape. +func SphericalHarmonicReal(l, m int, theta, phi *Array) (*Array, error) { + // m = 0 is purely real in the complex form, but the contract here + // is a Float array, so the value is extracted rather than handing + // the complex Y_l⁰ through unchanged. + yc, err := SphericalHarmonic(l, absInt(m), theta, phi) + if err != nil { + return nil, err + } + out := &Array{shape: append([]int{}, theta.shape...), dt: Float} + out.alloc(theta.Len()) + root2 := math.Sqrt2 + // The standard real forms carry the (−1)^m phase on both sides of + // zero, mirroring the Condon-Shortley factor the complex harmonics + // carry. + sign := 1.0 + if m%2 != 0 { + sign = -1 + } + for i := range theta.Len() { + z := yc.complexAt(i) + switch { + case m == 0: + out.floats[i] = real(z) + case m > 0: + out.floats[i] = sign * root2 * real(z) + default: + out.floats[i] = sign * root2 * imag(z) + } + } + return out, nil +} + +// Airy returns the Airy functions of the first and second kind at +// each element, Ai in the first return slot and Bi in the second. +// The power series about zero is used, which converges for every +// real x but is numerically dependable only within |x| ≤ 8; outside +// that window the functions return NaN rather than a silently wrong +// value. Positive arguments above 2 sum the series in extended +// precision, where the float64 reconstruction of Ai would cancel the +// growing solution and lose digits. The scaling convention is the +// standard Ai, Bi with Ai(0) = 1/(3^{2/3}Γ(2/3)), +// Bi(0) = 1/(3^{1/6}Γ(1/3)). +func Airy(a *Array) (ai, bi *Array, err error) { + if a.dt == Complex { + return nil, nil, errf("Airy: complex arrays are not supported") + } + ai = &Array{shape: append([]int{}, a.shape...), dt: Float} + ai.alloc(a.Len()) + bi = &Array{shape: append([]int{}, a.shape...), dt: Float} + bi.alloc(a.Len()) + engine.Parallel(a.Len(), func(s, e int) { + for i := s; i < e; i++ { + x := a.floatAt(i) + if math.Abs(x) > 8 { + ai.floats[i] = math.NaN() + bi.floats[i] = math.NaN() + continue + } + aiv, biv := airySeries(x) + ai.floats[i] = aiv + bi.floats[i] = biv + } + }) + return ai, bi, nil +} + +// airySeries evaluates the Maclaurin pair +// Ai(x) = c1·f(x) − c2·g(x), Bi(x) = √3·(c1·f(x) + c2·g(x)) +// with f = Σ 3ᵏ(1/3)ₖx^{3k}/(3k)!, g = Σ 3ᵏ(2/3)ₖx^{3k+1}/(3k+1)!, +// summed by the term-ratio recurrence to machine precision. For +// positive arguments above airyBigCut the Ai reconstruction runs in +// extended precision instead: Bi is a sum of like-signed terms and +// keeps the float64 series. +func airySeries(x float64) (aiv, biv float64) { + const ( + c1 = 0.35502805388781723926 + c2 = 0.25881940379280679840 + ) + tf, tg := 1.0, x + sumF, sumG := 1.0, x + for k := range 400 { + tf *= 3 * x * x * x * (float64(k) + 1.0/3) / float64((3*k+1)*(3*k+2)*(3*k+3)) + tg *= 3 * x * x * x * (float64(k) + 2.0/3) / float64((3*k+2)*(3*k+3)*(3*k+4)) + sumF += tf + sumG += tg + if math.Abs(tf) < 1e-18*math.Abs(sumF) && math.Abs(tg) < 1e-18*math.Abs(sumG) { + break + } + } + biv = math.Sqrt(3) * (c1*sumF + c2*sumG) + if x > airyBigCut { + aiv = airyBig(x) + return aiv, biv + } + aiv = c1*sumF - c2*sumG + return aiv, biv +} + +// airyBigCut is the positive argument above which Ai is reconstructed +// in extended precision. At the cut the float64 difference c1·f − c2·g +// has lost 2ζ/ln2 ≈ 5 bits to the cancellation, which caps its +// relative accuracy at about 1e-14; below the cut the float64 series +// still holds fourteen digits. +const airyBigCut = 2.0 + +// The Maclaurin constants of the extended-precision branch, Ai(0) and +// −Ai'(0) to 50 significant digits: at the working sizes the +// cancellation needs, the float64 roundings of airySeries would +// reintroduce the very error the branch removes. The regression test +// pins their float64 roundings against the float64 constants above. +const ( + airyC1Digits = "0.35502805388781723926006318600418317639797917419918" + airyC2Digits = "0.25881940379280679840518356018920396347909113835493" +) + +// airyBig evaluates Ai(x) for x > 0 in extended precision. The +// reconstruction c1·f − c2·g cancels e^{2ζ}, ζ = 2x^{3/2}/3: both +// products carry the growing solution e^{ζ} while Ai itself is the +// decaying e^{−ζ} one, so the float64 difference loses 2ζ/ln2 ≈ +// 1.92·x^{3/2} bits (4.5e-3 relative at x = 8). The same two series are +// therefore summed in math/big with a working size scaled to that loss, +// following the fresnelBig pattern; the term ratios here are the +// simplified forms x³/((3k+2)(3k+3)) and x³/((3k+3)(3k+4)) the float64 +// recurrence above evaluates with the (3k+1) and (3k+2) factors still +// in place. +func airyBig(x float64) float64 { + bits := uint(64 + int(1.9235*x*math.Sqrt(x)) + 96) + xf := new(big.Float).SetPrec(bits).SetFloat64(x) + x3 := new(big.Float).SetPrec(bits).Mul(xf, xf) + x3.Mul(x3, xf) + c1, _, err := big.ParseFloat(airyC1Digits, 10, bits, big.ToNearestEven) + if err != nil { + return math.NaN() + } + c2, _, err := big.ParseFloat(airyC2Digits, 10, bits, big.ToNearestEven) + if err != nil { + return math.NaN() + } + tf := new(big.Float).SetPrec(bits).SetInt64(1) + tg := new(big.Float).SetPrec(bits).Set(xf) + sumF := new(big.Float).SetPrec(bits).Set(tf) + sumG := new(big.Float).SetPrec(bits).Set(tg) + // The terms are all positive for x > 0, so a cut in absolute size + // relative to a sum that stays above 1 is a relative cut. + tiny := new(big.Float).SetPrec(bits).SetMantExp(big.NewFloat(1), -int(bits)+20) + for k := range 400 { + r1 := new(big.Float).SetPrec(bits).Quo(x3, big.NewFloat(float64((3*k+2)*(3*k+3)))) + r2 := new(big.Float).SetPrec(bits).Quo(x3, big.NewFloat(float64((3*k+3)*(3*k+4)))) + tf.Mul(tf, r1) + tg.Mul(tg, r2) + sumF.Add(sumF, tf) + sumG.Add(sumG, tg) + if tf.Cmp(tiny) < 0 && tg.Cmp(tiny) < 0 { + break + } + } + term := new(big.Float).SetPrec(bits).Mul(c2, sumG) + aiv := new(big.Float).SetPrec(bits).Mul(c1, sumF) + aiv.Sub(aiv, term) + out, _ := aiv.Float64() + return out +} + +// absInt returns |n|. +func absInt(n int) int { + if n < 0 { + return -n + } + return n +} diff --git a/internal/core/orthopoly_test.go b/internal/core/orthopoly_test.go new file mode 100644 index 0000000..0119a52 --- /dev/null +++ b/internal/core/orthopoly_test.go @@ -0,0 +1,201 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// TestOrthogonalPolynomials checks H, L, T and U against their closed +// forms for the first degrees. +func TestOrthogonalPolynomials(t *testing.T) { + x := mustFloats(t, []float64{-0.7, 0, 0.4, 1.1}) + cases := []struct { + name string + n int + want func(v float64) float64 + }{ + {"H0", 0, func(v float64) float64 { return 1 }}, + {"H1", 1, func(v float64) float64 { return 2 * v }}, + {"H2", 2, func(v float64) float64 { return 4*v*v - 2 }}, + {"H3", 3, func(v float64) float64 { return 8*v*v*v - 12*v }}, + {"T0", 0, func(v float64) float64 { return 1 }}, + {"T2", 2, func(v float64) float64 { return 2*v*v - 1 }}, + {"T3", 3, func(v float64) float64 { return 4*v*v*v - 3*v }}, + {"U0", 0, func(v float64) float64 { return 1 }}, + {"U1", 1, func(v float64) float64 { return 2 * v }}, + {"U2", 2, func(v float64) float64 { return 4*v*v - 1 }}, + } + for _, tc := range cases { + var got *Array + var err error + switch { + case tc.name[0] == 'H' && len(tc.name) == 2: + got, err = Hermite(tc.n, x) + case tc.name[0] == 'T': + got, err = ChebyshevT(tc.n, x) + default: + got, err = ChebyshevU(tc.n, x) + } + if err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + for i := range x.Len() { + want := tc.want(x.FloatAt(i)) + if math.Abs(got.FloatAt(i)-want) > 1e-12*(1+math.Abs(want)) { + t.Fatalf("%s(%v) = %v, want %v", tc.name, x.FloatAt(i), got.FloatAt(i), want) + } + } + } +} + +// TestLaguerrePolynomials checks L against closed forms and the +// orthogonality weight sanity at a sample point. +func TestLaguerrePolynomials(t *testing.T) { + x := mustFloats(t, []float64{-0.3, 0, 0.6, 2.1}) + closed := []struct { + n int + alpha float64 + want func(v float64) float64 + }{ + {0, 0, func(v float64) float64 { return 1 }}, + {1, 0, func(v float64) float64 { return 1 - v }}, + {2, 0, func(v float64) float64 { return 1 - 2*v + v*v/2 }}, + {2, 1, func(v float64) float64 { return 3 - 3*v + v*v/2 }}, + } + for _, tc := range closed { + l, err := Laguerre(tc.n, tc.alpha, x) + if err != nil { + t.Fatalf("Laguerre(%d, %v): %v", tc.n, tc.alpha, err) + } + for i := range x.Len() { + want := tc.want(x.FloatAt(i)) + if math.Abs(l.FloatAt(i)-want) > 1e-12*(1+math.Abs(want)) { + t.Fatalf("L_%d^%v(%v) = %v, want %v", tc.n, tc.alpha, x.FloatAt(i), l.FloatAt(i), want) + } + } + } + if _, err := Laguerre(-1, 0, x); err == nil { + t.Fatal("expected an error for a negative degree") + } +} + +// TestSphericalHarmonicReal checks the real harmonics against the +// complex ones under the standard convention: both branches carry the +// (−1)^m phase, Y_real = √2·(−1)^m·Re(Y_l^m) for m > 0 and +// √2·(−1)^m·Im(Y_l^|m|) for m < 0. +func TestSphericalHarmonicReal(t *testing.T) { + theta := mustFloats(t, []float64{0.4, 1.2, 2.5}) + phi := mustFloats(t, []float64{0.3, -0.8, 2.1}) + pos, err := SphericalHarmonicReal(2, 1, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonicReal(+1): %v", err) + } + pos2, err := SphericalHarmonicReal(3, 2, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonicReal(+2): %v", err) + } + neg, err := SphericalHarmonicReal(2, -1, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonicReal(−1): %v", err) + } + cx, err := SphericalHarmonic(2, 1, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonic: %v", err) + } + cx2, err := SphericalHarmonic(3, 2, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonic(3, 2): %v", err) + } + root2 := math.Sqrt2 + for i := range theta.Len() { + z := cx.ComplexAt(i) + // m = 1 is odd, so the standard real form negates the real + // part; m = 2, even, keeps it. + if math.Abs(pos.FloatAt(i)+root2*real(z)) > 1e-12 { + t.Fatalf("real m=1 [%d] = %v, want %v", i, pos.FloatAt(i), -root2*real(z)) + } + z2 := cx2.ComplexAt(i) + if math.Abs(pos2.FloatAt(i)-root2*real(z2)) > 1e-12 { + t.Fatalf("real m=2 [%d] = %v, want %v", i, pos2.FloatAt(i), root2*real(z2)) + } + if math.Abs(neg.FloatAt(i)+root2*imag(z)) > 1e-12 { + t.Fatalf("real m=−1 [%d] = %v, want %v", i, neg.FloatAt(i), -root2*imag(z)) + } + } + // Pinned standard values against the external Cartesian forms: + // Y_real(1, 1) = √(3/(4π))·x/r, so it is +√(3/(4π)) at (θ, φ) = + // (π/2, 0) and 0 at (π/2, π/2); Y_real(1, −1) = √(3/(4π))·y/r is + // +√(3/(4π)) at (π/2, π/2) and 0 at (π/2, 0); Y_real(2, 1) = + // √(15/(4π))·xz/r² at (π/4, 0) is √(15/(4π))·cos(π/4)·sin(π/4). + eq := mustFloats(t, []float64{math.Pi / 2}) + halfPi := mustFloats(t, []float64{math.Pi / 2}) + zero := mustFloats(t, []float64{0}) + y11, err := SphericalHarmonicReal(1, -1, eq, halfPi) + if err != nil { + t.Fatalf("SphericalHarmonicReal(1, −1): %v", err) + } + if want := math.Sqrt(3 / (4 * math.Pi)); math.Abs(y11.FloatAt(0)-want) > 1e-12 { + t.Fatalf("Y_real(1, −1)(π/2, π/2) = %.16g, want %.16g", y11.FloatAt(0), want) + } + y11z, err := SphericalHarmonicReal(1, -1, eq, zero) + if err != nil { + t.Fatalf("SphericalHarmonicReal(1, −1) at φ=0: %v", err) + } + if v := y11z.FloatAt(0); math.Abs(v) > 1e-12 { + t.Fatalf("Y_real(1, −1)(π/2, 0) = %v, want 0", v) + } + y11p, err := SphericalHarmonicReal(1, 1, eq, zero) + if err != nil { + t.Fatalf("SphericalHarmonicReal(1, 1): %v", err) + } + if want := math.Sqrt(3 / (4 * math.Pi)); math.Abs(y11p.FloatAt(0)-want) > 1e-12 { + t.Fatalf("Y_real(1, 1)(π/2, 0) = %.16g, want +%.16g", y11p.FloatAt(0), want) + } + y11q, err := SphericalHarmonicReal(1, 1, eq, halfPi) + if err != nil { + t.Fatalf("SphericalHarmonicReal(1, 1) at φ=π/2: %v", err) + } + if v := y11q.FloatAt(0); math.Abs(v) > 1e-12 { + t.Fatalf("Y_real(1, 1)(π/2, π/2) = %v, want 0", v) + } + qr := mustFloats(t, []float64{math.Pi / 4}) + y21, err := SphericalHarmonicReal(2, 1, qr, zero) + if err != nil { + t.Fatalf("SphericalHarmonicReal(2, 1): %v", err) + } + if want := math.Sqrt(15/(4*math.Pi)) * math.Cos(math.Pi/4) * math.Sin(math.Pi/4); math.Abs(y21.FloatAt(0)-want) > 1e-12 { + t.Fatalf("Y_real(2, 1)(π/4, 0) = %.16g, want %.16g", y21.FloatAt(0), want) + } +} + +// TestAiry checks the Airy pair against tabulated values at x = 0, +// ±1 and the NaN contract outside the series window. +func TestAiry(t *testing.T) { + x := mustFloats(t, []float64{0, 1, -1, 3.5}) + ai, bi, err := Airy(x) + if err != nil { + t.Fatalf("Airy: %v", err) + } + wantAI := []float64{0.3550280538878172, 0.13529241631288142, 0.5355608832928901, 0.002584098786988065} + wantBI := []float64{0.6149266274460007, 1.2074235949528713, 0.1039973894969447, 33.05550674355921} + for i := range 4 { + if math.Abs(ai.FloatAt(i)-wantAI[i]) > 1e-12*(1+math.Abs(wantAI[i])) { + t.Fatalf("Ai(%v) = %.16g, want %.16g", x.FloatAt(i), ai.FloatAt(i), wantAI[i]) + } + if math.Abs(bi.FloatAt(i)-wantBI[i]) > 5e-10*(1+math.Abs(wantBI[i])) { + t.Fatalf("Bi(%v) = %.16g, want %.16g", x.FloatAt(i), bi.FloatAt(i), wantBI[i]) + } + } + // Outside the series window the contract is NaN, not silence. + out := mustFloats(t, []float64{9}) + ai9, bi9, err := Airy(out) + if err != nil { + t.Fatalf("Airy(9): %v", err) + } + if !math.IsNaN(ai9.FloatAt(0)) || !math.IsNaN(bi9.FloatAt(0)) { + t.Fatalf("outside |x| ≤ 8 expected NaN, got %v and %v", ai9.FloatAt(0), bi9.FloatAt(0)) + } +} diff --git a/internal/core/overflow_refusal_pin_test.go b/internal/core/overflow_refusal_pin_test.go new file mode 100644 index 0000000..55a3ad5 --- /dev/null +++ b/internal/core/overflow_refusal_pin_test.go @@ -0,0 +1,142 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +// Regression pins: refusals that used to be panics or silent +// wrong shapes, and the einsum trace dtype contract. + +// TestGridComplexRefusal: complex axes reached FloatAt's nil-int branch +// and panicked; they are a named error now, as every sibling helper +// already answered them. +func TestGridComplexRefusal(t *testing.T) { + c := mustFromComplexes(t, []complex128{1, 2}, 2) + if _, _, err := Grid(c, c); err == nil || !strings.Contains(err.Error(), "Grid: complex") { + t.Fatalf("Grid on complex axes: err = %v", err) + } + f := mustFloats(t, []float64{1, 2}, 2) + if _, _, err := Grid(c, f); err == nil || !strings.Contains(err.Error(), "Grid: complex") { + t.Fatalf("Grid on mixed axes: err = %v", err) + } +} + +// TestTileOverflow: an unchecked as*r extent wrapped negative and the +// allocation panicked; the guard answers it like Repeat's. +func TestTileOverflow(t *testing.T) { + a := mustFromInts(t, []int64{1, 2}, 2) + if _, err := Tile(a, 1<<62); err == nil || !strings.Contains(err.Error(), "Tile:") { + t.Fatalf("Tile with a wrapping extent: err = %v", err) + } +} + +// TestOneHotOverflow: n*classes wrapped negative and the makeslice +// panicked. +func TestOneHotOverflow(t *testing.T) { + codes := mustFromInts(t, []int64{0, 0}, 2) + if _, err := OneHot(codes, math.MaxInt); err == nil || !strings.Contains(err.Error(), "OneHot:") { + t.Fatalf("OneHot with a wrapping row width: err = %v", err) + } +} + +// TestConcatDimOverflow: the joined length wrapped negative and a +// corrupt shape escaped to the caller. +func TestConcatDimOverflow(t *testing.T) { + a, err := Zeros(Int, 1<<62, 0) + if err != nil { + t.Fatal(err) + } + b, err := Zeros(Int, 1<<62+4, 0) + if err != nil { + t.Fatal(err) + } + if _, err := Concat(a, b, 0); err == nil || !strings.Contains(err.Error(), "Concat:") { + t.Fatalf("Concat with a wrapping joined length: err = %v", err) + } +} + +// TestKronOverflow: the outer-product shape bypassed validation and a +// wrapped dimension reached the allocation. +func TestKronOverflow(t *testing.T) { + a, err := Zeros(Int, 2, 0) + if err != nil { + t.Fatal(err) + } + b, err := Zeros(Int, 1<<62, 0) + if err != nil { + t.Fatal(err) + } + if _, err := Kron(a, b); err == nil || !strings.Contains(err.Error(), "Kron:") { + t.Fatalf("Kron with a wrapping shape product: err = %v", err) + } +} + +// TestSparseCOOShapeValidation: the stored shape is allocation-bearing +// metadata (SpMatMul allocates from it), so a hostile or wrapped shape +// is a named error instead of a later panic. +func TestSparseCOOShapeValidation(t *testing.T) { + idx := mustFromInts(t, []int64{0, 0, 0, 1}, 2, 2) + vals := mustFloats(t, []float64{1, 2}, 2) + if _, err := NewSparseCOO(idx, vals, []int{3, -1}); err == nil || !strings.Contains(err.Error(), "NewSparseCOO:") { + t.Fatalf("negative sparse shape: err = %v", err) + } + if _, err := NewSparseCOO(idx, vals, []int{3, 1 << 62}); err == nil || !strings.Contains(err.Error(), "NewSparseCOO:") { + t.Fatalf("wrapping sparse shape: err = %v", err) + } +} + +// TestEinsumTraceKeepsDtype: the trace fast path widened int operands +// to float64, flipping the result dtype against the general engine and +// rounding above 2^53, and it refused complex operands outright. +func TestEinsumTraceKeepsDtype(t *testing.T) { + m := mustFromInts(t, []int64{1 << 53, 1, 1, 1}, 2, 2) + out, err := Einsum("ii->", m) + if err != nil { + t.Fatal(err) + } + if out.Dtype() != Int { + t.Fatalf("trace of an int matrix is %s, want Int", out.Dtype()) + } + if v, _ := IntAt(out, 0); v != 1<<53+1 { + t.Fatalf("int trace = %d, want %d exactly", v, 1<<53+1) + } + c := mustFromComplexes(t, []complex128{1 + 2i, 3, 4, 5 - 1i}, 2, 2) + cout, err := Einsum("ii->", c) + if err != nil { + t.Fatal(err) + } + if cout.Dtype() != Complex { + t.Fatalf("trace of a complex matrix is %s, want Complex", cout.Dtype()) + } + if v, _ := ComplexAt(cout, 0); v != 6+1i { + t.Fatalf("complex trace = %v, want (6+1i)", v) + } + ns, err := Zeros(Float, 2, 3) + if err != nil { + t.Fatal(err) + } + if _, err := Einsum("ii->", ns); err == nil || !strings.Contains(err.Error(), "Einsum:") { + t.Fatalf("trace of a non-square operand: err = %v", err) + } +} + +// TestElementsIntStridedGather: the int64 branch read the payload raw +// while every other branch gathered through physIndex. +func TestElementsIntStridedGather(t *testing.T) { + strided := &Array{shape: []int{2, 2}, dt: Int, ints: []int64{10, 11, 12, 13, 14, 15}, strides: []int{3, 1}} + got, err := strided.Elements[int64]() + if err != nil { + t.Fatal(err) + } + want := []int64{10, 11, 13, 14} + for i, w := range want { + if got[i] != w { + t.Fatalf("Elements[int64] strided[%d] = %d, want %d", i, got[i], w) + } + } +} diff --git a/internal/core/pchip.go b/internal/core/pchip.go new file mode 100644 index 0000000..34d074f --- /dev/null +++ b/internal/core/pchip.go @@ -0,0 +1,117 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "math" + +// Monotone cubic interpolation. The Fritsch-Carlson tangents keep the +// curve inside the data's own range between knots: interpolating a +// monotone series cannot overshoot, which the plain cubic spline does +// around sharp turns and which is the reason this variant exists. + +// InterpolateMonotone evaluates the monotone piecewise cubic (PCHIP) +// through the samples at every query point. xs must be strictly +// increasing, ys real, and at least two samples long; queries outside +// [x0, xn] hold the boundary value, matching Interpolate's convention. +// At every knot the curve passes through the sample, and its slope +// there never exceeds twice the neighbouring secants, which is what +// keeps the interpolation inside the local data range. +func InterpolateMonotone(xs, ys, query *Array) (*Array, error) { + if xs.dt == Complex || ys.dt == Complex || query.dt == Complex { + return nil, errf("InterpolateMonotone: complex samples are not supported") + } + n := xs.Len() + if n != ys.Len() || n < 2 { + return nil, errf("InterpolateMonotone: xs/ys must share length ≥ 2") + } + for k := 1; k < n; k++ { + xk, xprev := xs.FloatAt(k), xs.FloatAt(k-1) + // NaN defeats the <= ordering test below (every comparison is + // false), so non-finiteness is refused by name first. + if math.IsNaN(xk) || math.IsInf(xk, 0) || math.IsNaN(xprev) || math.IsInf(xprev, 0) { + return nil, errf("InterpolateMonotone: knot %d is not finite", k) + } + if xk <= xprev { + return nil, errf("InterpolateMonotone: xs must be strictly increasing, knot %d repeats", k) + } + } + h := make([]float64, n-1) + delta := make([]float64, n-1) + for k := range n - 1 { + h[k] = xs.FloatAt(k+1) - xs.FloatAt(k) + delta[k] = (ys.FloatAt(k+1) - ys.FloatAt(k)) / h[k] + } + // Fritsch-Carlson tangents: zero where the data turns, the + // weighted harmonic mean of the neighbouring secants elsewhere, + // the clamped three-point estimate at the ends. + m := make([]float64, n) + switch { + case n == 2: + m[0], m[1] = delta[0], delta[0] + default: + m[0] = ((2*h[0]+h[1])*delta[0] - h[0]*delta[1]) / (h[0] + h[1]) + if oppositeSigns(m[0], delta[0]) { + m[0] = 0 + } else if math.Abs(m[0]) > 2*math.Abs(delta[0]) { + m[0] = 2 * delta[0] + } + m[n-1] = ((2*h[n-2]+h[n-3])*delta[n-2] - h[n-3]*delta[n-3]) / (h[n-2] + h[n-3]) + if oppositeSigns(m[n-1], delta[n-2]) { + m[n-1] = 0 + } else if math.Abs(m[n-1]) > 2*math.Abs(delta[n-2]) { + m[n-1] = 2 * delta[n-2] + } + for k := 1; k < n-1; k++ { + if delta[k-1]*delta[k] <= 0 { + m[k] = 0 + continue + } + w1 := 2*h[k] + h[k-1] + w2 := h[k] + 2*h[k-1] + m[k] = (w1 + w2) / (w1/delta[k-1] + w2/delta[k]) + } + } + out := &Array{shape: append([]int{}, query.Shape()...), dt: Float} + out.alloc(query.Len()) + for i := range query.Len() { + q := query.FloatAt(i) + // NaN compares false against both clamps below and would flow + // through the evaluation as NaN with no error, the same trap + // InterpolateGrid refuses: an undefined position is a loud + // error, since there is nothing sensible to clamp it to. + if math.IsNaN(q) { + return nil, errf("InterpolateMonotone: query %d is NaN, which cannot be clamped", i) + } + lo := 0 + if q <= xs.FloatAt(0) { + lo = 0 + } else if q >= xs.FloatAt(n-1) { + lo = n - 2 + } else { + for lo < n-2 && q > xs.FloatAt(lo+1) { + lo++ + } + } + t := (q - xs.FloatAt(lo)) / h[lo] + if t < 0 { + t = 0 + } + if t > 1 { + t = 1 + } + t2, t3 := t*t, t*t*t + y0, y1 := ys.FloatAt(lo), ys.FloatAt(lo+1) + v := (2*t3-3*t2+1)*y0 + (t3-2*t2+t)*h[lo]*m[lo] + + (-2*t3+3*t2)*y1 + (t3-t2)*h[lo]*m[lo+1] + out.SetFloatAt(i, v) + } + return out, nil +} + +// oppositeSigns reports whether two values carry strict opposite +// signs, the condition under which a tangent estimate is discarded as +// turning. +func oppositeSigns(a, b float64) bool { + return (a > 0 && b < 0) || (a < 0 && b > 0) +} diff --git a/internal/core/quasirandom.go b/internal/core/quasirandom.go new file mode 100644 index 0000000..cd06562 --- /dev/null +++ b/internal/core/quasirandom.go @@ -0,0 +1,242 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math/bits" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Quasi-random sequences: Halton points, the low- +// discrepancy workhorse of Monte Carlo that needs no tables. Each +// coordinate runs the radical inverse in its own prime base, which +// stratifies every b^k block of points evenly through the unit +// hypercube: the property random sampling only has in expectation. + +// HaltonPoints returns the first n Halton points of the given +// dimension as an (n, dim) float64 array, skipping the leading skip +// points (the early Halton coordinates correlate visibly in high +// dimensions; the standard cure is to drop them). dim must be at +// least 1 and at most 32: beyond that the available small primes +// run out and the stratification degrades; a larger request is an +// error, not a silently worse sequence. n+skip must stay below 2^32, +// the index budget both quasi-random constructors enforce; past it +// the point indices wrap the int arithmetic and every point collapses +// back to the origin instead of advancing. +func HaltonPoints(n, dim, skip int) (*Array, error) { + const name = "HaltonPoints" + if n < 0 { + return nil, errf("%s: n must be non-negative, got %d", name, n) + } + if dim < 1 || dim > 32 { + return nil, errf("%s: dim must lie in [1, 32], got %d", name, dim) + } + if skip < 0 { + return nil, errf("%s: skip must be non-negative, got %d", name, skip) + } + if uint64(skip)+uint64(n) >= 1<<32 { + return nil, errf("%s: n + skip must stay below 2^32, the index budget, got %d + %d", + name, n, skip) + } + bases := firstPrimes(dim) + out := &Array{shape: []int{n, dim}, dt: Float} + out.alloc(n * dim) + // The points are independent, so the walk splits over disjoint + // point ranges: every point writes its own slots and no value is + // accumulated, so the split cannot move a single coordinate. + engine.ParallelMin(n, copyMinPerWorker, func(s, e int) { + for p := s; p < e; p++ { + idx := p + skip + 1 // the point indexed 0 would be the origin + for d := range dim { + out.floats[p*dim+d] = radicalInverse(idx, bases[d]) + } + } + }) + return out, nil +} + +// radicalInverse reflects the base-b digits of i through the radix +// point: the van der Corput core of every Halton coordinate. +func radicalInverse(i int, b int) float64 { + f := 1.0 + r := 0.0 + for i > 0 { + f /= float64(b) + r += f * float64(i%b) + i /= b + } + return r +} + +// firstPrimes returns the first count primes by trial division. +func firstPrimes(count int) []int { + primes := make([]int, 0, count) + candidate := 2 + for len(primes) < count { + isPrime := true + for _, p := range primes { + if p*p > candidate { + break + } + if candidate%p == 0 { + isPrime = false + break + } + } + if isPrime { + primes = append(primes, candidate) + } + candidate++ + } + return primes +} + +// Sobol sequences: the digital low-discrepancy companion to +// Halton. Each coordinate runs its own linear recurrence over GF(2) +// driven by direction numbers derived from a primitive polynomial, so +// unlike Halton every dimension shares the same base-2 lattice and the +// first 2^k points are exactly stratified through every coordinate, +// not just on average. The parameters below are the Joe and Kuo (2008) +// initialisation table, the current standard, for the first 40 +// dimensions; dimension 1 is the plain Gray-coded van der Corput +// sequence and carries no polynomial. + +// sobolParams are the initialisation parameters of one dimension: the +// polynomial degree s, the primitive polynomial coefficient a whose +// bits select the recurrence taps, and the s odd initialisation +// integers m_i with 1 <= m_i < 2^i. +type sobolParams struct { + s uint32 + a uint32 + m []uint32 +} + +var sobolTable = [...]sobolParams{ + {0, 0, nil}, // 1: plain van der Corput + {1, 0, []uint32{1}}, // 2 + {2, 1, []uint32{1, 3}}, // 3 + {3, 1, []uint32{1, 3, 1}}, // 4 + {3, 2, []uint32{1, 1, 1}}, // 5 + {4, 1, []uint32{1, 1, 3, 3}}, // 6 + {4, 4, []uint32{1, 3, 5, 13}}, // 7 + {5, 2, []uint32{1, 1, 5, 5, 17}}, // 8 + {5, 4, []uint32{1, 1, 5, 5, 5}}, // 9 + {5, 7, []uint32{1, 1, 7, 11, 19}}, // 10 + {5, 11, []uint32{1, 1, 5, 1, 1}}, // 11 + {5, 13, []uint32{1, 1, 1, 3, 11}}, // 12 + {5, 14, []uint32{1, 3, 5, 5, 31}}, // 13 + {6, 1, []uint32{1, 3, 3, 9, 7, 49}}, // 14 + {6, 13, []uint32{1, 1, 1, 15, 21, 21}}, // 15 + {6, 16, []uint32{1, 3, 1, 13, 27, 49}}, // 16 + {6, 19, []uint32{1, 1, 1, 15, 7, 5}}, // 17 + {6, 22, []uint32{1, 3, 1, 15, 13, 25}}, // 18 + {6, 25, []uint32{1, 1, 5, 5, 19, 61}}, // 19 + {7, 1, []uint32{1, 3, 7, 11, 23, 15, 103}}, // 20 + {7, 4, []uint32{1, 3, 7, 13, 13, 15, 69}}, // 21 + {7, 7, []uint32{1, 1, 3, 13, 7, 35, 63}}, // 22 + {7, 8, []uint32{1, 3, 5, 9, 1, 25, 53}}, // 23 + {7, 14, []uint32{1, 3, 1, 13, 9, 35, 107}}, // 24 + {7, 19, []uint32{1, 3, 1, 5, 27, 61, 31}}, // 25 + {7, 21, []uint32{1, 1, 5, 11, 19, 41, 61}}, // 26 + {7, 28, []uint32{1, 3, 5, 3, 3, 13, 69}}, // 27 + {7, 31, []uint32{1, 1, 7, 13, 1, 19, 1}}, // 28 + {7, 32, []uint32{1, 3, 7, 5, 13, 19, 59}}, // 29 + {7, 37, []uint32{1, 1, 3, 9, 25, 29, 41}}, // 30 + {7, 41, []uint32{1, 3, 5, 13, 23, 1, 55}}, // 31 + {7, 42, []uint32{1, 3, 7, 3, 13, 59, 17}}, // 32 + {7, 50, []uint32{1, 3, 1, 3, 5, 53, 69}}, // 33 + {7, 55, []uint32{1, 1, 5, 5, 23, 33, 13}}, // 34 + {7, 56, []uint32{1, 1, 7, 7, 1, 61, 123}}, // 35 + {7, 59, []uint32{1, 1, 7, 9, 13, 61, 49}}, // 36 + {7, 62, []uint32{1, 3, 3, 5, 3, 55, 33}}, // 37 + {8, 14, []uint32{1, 3, 1, 15, 31, 13, 49, 245}}, // 38 + {8, 21, []uint32{1, 3, 5, 15, 31, 59, 63, 97}}, // 39 + {8, 22, []uint32{1, 3, 1, 11, 11, 11, 77, 249}}, // 40 +} + +// sobolDirections fills v with the 32-bit direction numbers of the +// given dimension. The first s come straight from the initialisation +// integers scaled into their binary place; the rest follow the +// recurrence v_i = v_{i-s} ^ (v_{i-s} >> s) ^ Σ a_k·v_{i-k} over the +// polynomial taps. +func sobolDirections(p sobolParams, v []uint32) { + s := int(p.s) + if s == 0 { + // Dimension 1: the direction numbers are the binary places + // themselves, which makes the sequence Gray-coded van der Corput. + for i := range v { + v[i] = 1 << (31 - i) + } + return + } + for i := range s { + v[i] = p.m[i] << (31 - i) + } + for i := s; i < len(v); i++ { + v[i] = v[i-s] ^ (v[i-s] >> uint(s)) + for k := 1; k < s; k++ { + if p.a>>(uint(s-1-k))&1 != 0 { + v[i] ^= v[i-k] + } + } + } +} + +// SobolPoints returns the first n Sobol points of the given dimension +// as an (n, dim) float64 array, skipping the leading skip points (the +// sequence starts at the origin index, which carries no information +// and is dropped, exactly as HaltonPoints drops it). dim must be at +// least 1 and at most 40: that is the width of the initialisation +// table, and a larger request is an error, not a silently worse +// sequence. n+skip must stay below 2^32, the period the 32-bit +// direction numbers give; beyond it the index arithmetic would wrap +// back to the origin instead of advancing. +func SobolPoints(n, dim, skip int) (*Array, error) { + const name = "SobolPoints" + if n < 0 { + return nil, errf("%s: n must be non-negative, got %d", name, n) + } + if dim < 1 || dim > len(sobolTable) { + return nil, errf("%s: dim must lie in [1, %d], got %d", name, len(sobolTable), dim) + } + if skip < 0 { + return nil, errf("%s: skip must be non-negative, got %d", name, skip) + } + if uint64(skip)+uint64(n) >= 1<<32 { + return nil, errf("%s: n + skip must stay below 2^32, the sequence period, got %d + %d", + name, n, skip) + } + const width = 32 + out := &Array{shape: []int{n, dim}, dt: Float} + out.alloc(n * dim) + // Each coordinate walks its own points. The worker seeds its range's + // Gray code from scratch, then advances one flip at a time: the Gray + // codes of consecutive indices differ in the lowest set bit of the + // newer index, so x picks up exactly the direction numbers the + // from-scratch walk would XOR, and the exclusive or combines them + // exactly whatever the order. + engine.ParallelMin(n, copyMinPerWorker, func(s, e int) { + for d := range dim { + var v [width]uint32 + sobolDirections(sobolTable[d], v[:]) + idx := s + skip + 1 // the point indexed 0 would be the origin + gray := uint32(idx) ^ uint32(idx>>1) + var x uint32 + for b := range width { + if gray&(1< (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strconv" + "strings" + "testing" +) + +// TestHaltonRadicalInverse pins the van der Corput construction against +// the digit-reflection definition. +func TestHaltonRadicalInverse(t *testing.T) { + cases := []struct { + i, b int + want float64 + }{ + {1, 2, 0.5}, + {2, 2, 0.25}, + {3, 2, 0.75}, + {4, 2, 0.125}, + {5, 2, 0.625}, + {1, 3, 1.0 / 3}, + {2, 3, 2.0 / 3}, + {3, 3, 1.0 / 9}, + {7, 3, 5.0 / 9}, // 21₃ reflected is 12₃ = 1/3 + 2/9 + {5, 3, 7.0 / 9}, // 12₃ reflected is 21₃ = 2/3 + 1/9 + {1, 5, 0.2}, + {4, 5, 0.8}, + } + for _, c := range cases { + if got := radicalInverse(c.i, c.b); math.Abs(got-c.want) > 1e-15 { + t.Fatalf("radicalInverse(%d, %d) = %g, want %g", c.i, c.b, got, c.want) + } + } +} + +// TestHaltonStratification pins the defining property: the first b^k +// points integrate a smooth function far better than random sampling +// would, and ∫x dx through the mean sits at 1/2 to high accuracy. +func TestHaltonStratification(t *testing.T) { + const n = 4096 + pts, err := HaltonPoints(n, 2, 0) + if err != nil { + t.Fatalf("HaltonPoints: %v", err) + } + mean := 0.0 + for i := range n { + mean += pts.FloatAt(i * 2) + } + mean /= float64(n) + // The base-2 marginal is a permutation of {k/8192}: its mean sits + // at 1/2 to within one point's contribution. + if math.Abs(mean-0.5) > 1.5/float64(n) { + t.Fatalf("mean of base-2 marginal = %.8f, want 1/2 within one point", mean) + } +} + +// TestHaltonQuadratureAccuracy pins the Monte Carlo payoff: a Halton +// estimate of a smooth 2-D integral beats the random-sampling error +// scale by orders of magnitude. +func TestHaltonQuadratureAccuracy(t *testing.T) { + const n = 1024 + pts, err := HaltonPoints(n, 2, 100) + if err != nil { + t.Fatalf("HaltonPoints: %v", err) + } + est := 0.0 + for i := range n { + x := pts.FloatAt(i * 2) + y := pts.FloatAt(i*2 + 1) + est += math.Exp(-x-y*y) / float64(n) + } + // Reference: ∫₀¹∫₀¹ e^{-x-y²} = (1−e^{-1})·sqrt(π)/2·erf(1). The + // Halton error decays like O(log n / n); a thousand points sit + // around 1e-4, where random sampling would scatter at 1/sqrt(n) + // ~ 3e-2. + want := (1 - math.Exp(-1)) * math.Sqrt(math.Pi) / 2 * math.Erf(1) + if math.Abs(est-want) > 5e-4 { + t.Fatalf("Halton integral = %.8f, exact %.8f", est, want) + } +} + +// TestHaltonErrors pins the input gates. +func TestHaltonErrors(t *testing.T) { + if _, err := HaltonPoints(10, 0, 0); err == nil { + t.Error("dim 0 accepted") + } + if _, err := HaltonPoints(10, 40, 0); err == nil { + t.Error("dim 40 accepted") + } + if _, err := HaltonPoints(10, 2, -1); err == nil { + t.Error("negative skip accepted") + } + // Zero points is a valid empty request. + empty, err := HaltonPoints(0, 2, 0) + if err != nil || empty.Len() != 0 { + t.Fatalf("n = 0 must give an empty array, got err %v", err) + } +} + +// TestSobolCanonicalPoints pins the two dimensions every Sobol table +// agrees on: dimension 1 is Gray-coded van der Corput and the first +// four 2-D points are the canonical square-covering quartet. +func TestSobolCanonicalPoints(t *testing.T) { + want1 := []float64{0.5, 0.75, 0.25, 0.375, 0.875, 0.625, 0.125, 0.1875} + p1, err := SobolPoints(len(want1), 1, 0) + if err != nil { + t.Fatalf("SobolPoints: %v", err) + } + for i, want := range want1 { + if got := p1.FloatAt(i); got != want { + t.Fatalf("dim 1 point %d = %g, want %g", i+1, got, want) + } + } + + pts, err := SobolPoints(4, 2, 0) + if err != nil { + t.Fatalf("SobolPoints: %v", err) + } + want2 := [4][2]float64{{0.5, 0.5}, {0.75, 0.25}, {0.25, 0.75}, {0.375, 0.375}} + for p, want := range want2 { + for d := range 2 { + if got := pts.FloatAt(p*2 + d); got != want[d] { + t.Fatalf("point %d dim %d = %g, want %g", p+1, d+1, got, want[d]) + } + } + } +} + +// TestSobolStratification pins the defining digital property across +// the whole table: in every dimension the 2^m points X_0..X_{2^m-1} +// (the origin included, since it carries the zero cell) hit each of +// the 2^m equal subintervals exactly once, which holds only if the +// direction numbers are linearly independent over GF(2). +func TestSobolStratification(t *testing.T) { + for _, b := range []int{4, 8} { + count := 1 << b + for dim := 1; dim <= len(sobolTable); dim++ { + // X_1..X_{count-1}; X_0 is the origin and fills cell 0. + pts, err := SobolPoints(count-1, dim, 0) + if err != nil { + t.Fatalf("SobolPoints dim %d: %v", dim, err) + } + for d := range dim { + seen := make([]bool, count) + seen[0] = true + for p := range count - 1 { + cell := int(pts.FloatAt(p*dim+d) * float64(count)) + if cell == count { + cell = count - 1 // a value of 1.0 would be out of range + } + if seen[cell] { + t.Fatalf("dim %d, coordinate %d: subinterval %d hit twice in the first %d points", dim, d, cell, count) + } + seen[cell] = true + } + } + } + } +} + +// TestSobolSkip: skipping must be equivalent to generating more points +// and discarding the leading ones, for every coordinate. +func TestSobolSkip(t *testing.T) { + full, err := SobolPoints(12, 3, 0) + if err != nil { + t.Fatalf("SobolPoints: %v", err) + } + tail, err := SobolPoints(4, 3, 8) + if err != nil { + t.Fatalf("SobolPoints: %v", err) + } + for p := range 4 { + for d := range 3 { + got := tail.FloatAt(p*3 + d) + want := full.FloatAt((p+8)*3 + d) + if got != want { + t.Fatalf("skipped point %d dim %d = %g, want %g", p, d+1, got, want) + } + } + } +} + +// TestSobolQuadratureAccuracy: the same integral TestHaltonQuadrature +// pins, at the same sample count and bound. Sobol carries no +// early-point correlation to skip away, so the sample runs from the +// start, where the digital structure is densest. +func TestSobolQuadratureAccuracy(t *testing.T) { + const n = 1024 + pts, err := SobolPoints(n, 2, 0) + if err != nil { + t.Fatalf("SobolPoints: %v", err) + } + est := 0.0 + for i := range n { + x := pts.FloatAt(i * 2) + y := pts.FloatAt(i*2 + 1) + est += math.Exp(-x-y*y) / float64(n) + } + want := (1 - math.Exp(-1)) * math.Sqrt(math.Pi) / 2 * math.Erf(1) + if math.Abs(est-want) > 5e-4 { + t.Fatalf("Sobol integral = %.8f, exact %.8f", est, want) + } +} + +// TestSobolErrors pins the input gates. +func TestSobolErrors(t *testing.T) { + if _, err := SobolPoints(10, 0, 0); err == nil { + t.Error("dim 0 accepted") + } + if _, err := SobolPoints(10, 41, 0); err == nil { + t.Error("dim 41 accepted") + } + if _, err := SobolPoints(-1, 2, 0); err == nil { + t.Error("negative n accepted") + } + if _, err := SobolPoints(10, 2, -1); err == nil { + t.Error("negative skip accepted") + } + empty, err := SobolPoints(0, 2, 0) + if err != nil || empty.Len() != 0 { + t.Fatalf("n = 0 must give an empty array, got err %v", err) + } + // The period guard needs indices a 32-bit int cannot express. + if strconv.IntSize >= 64 { + if _, err := SobolPoints(2, 1, (1<<32)-1); err == nil || + !strings.Contains(err.Error(), "sequence period") { + t.Errorf("period wrap accepted: %v", err) + } + // The last legal index is still fine. + if _, err := SobolPoints(1, 1, (1<<32)-2); err != nil { + t.Errorf("last legal index refused: %v", err) + } + } +} diff --git a/internal/core/random.go b/internal/core/random.go new file mode 100644 index 0000000..10353ca --- /dev/null +++ b/internal/core/random.go @@ -0,0 +1,224 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/bits" +) + +// Reproducible randomness. The generator runs xoshiro256++ +// seeded through splitmix64: implemented here because math/rand/v2 does +// not promise stable output across Go versions, and reproducibility is +// the point of a seed: the same seed yields bit-identical arrays on every +// Go release. The quality suits simulation, not cryptography. The splitmix +// machinery is exported for callers who seed their own streams, and +// Substream hands out the provably distinct members of one seed's stream +// family. + +// Generator produces deterministic random arrays from a seed. +type Generator struct { + s [4]uint64 +} + +// NewGenerator seeds a fresh generator. Any int64 seed is valid. +func NewGenerator(seed int64) *Generator { + return generatorFrom(uint64(seed)) +} + +// Splitmix64 advances the splitmix64 stream one step: from the given +// state it returns the advanced state and the mixed output. The mixer +// is a bijection on uint64, xor shifts and odd multipliers both +// inverting, which is the property the substream construction rests +// on. Any state is valid, the zero included. +func Splitmix64(state uint64) (uint64, uint64) { + state += 0x9E3779B97F4A7C15 + return state, mix64(state) +} + +// Substream seeds the index-th member of the stream family one seed +// carries. The start state mixes the index through splitmix64 before +// combining it with the seed and mixing again, and both mixings are +// bijections, so distinct indices give provably distinct initial +// states. That is the property a seed + index stride lacks: its +// streams are one stream read at different offsets, which measures a +// rearrangement where a family of independent streams was wanted. +// Indices start at zero. +func Substream(seed int64, index int) (*Generator, error) { + if index < 0 { + return nil, errf("Substream: the index must be zero or greater, got %d", index) + } + return generatorFrom(mix64(uint64(seed) ^ mix64(uint64(index)))), nil +} + +// generatorFrom seeds a generator by walking the splitmix64 stream +// four steps from z. +func generatorFrom(z uint64) *Generator { + g := &Generator{} + for i := range g.s { + var v uint64 + z, v = Splitmix64(z) + g.s[i] = v + } + // xoshiro never leaves the all-zero state, so nudge it away. + if g.s == [4]uint64{} { + g.s = [4]uint64{1, 2, 3, 4} + } + return g +} + +// mix64 is the splitmix64 finaliser, a bijection on uint64. +func mix64(z uint64) uint64 { + z = (z ^ (z >> 30)) * 0xBF58476D1CE4E5B9 + z = (z ^ (z >> 27)) * 0x94D049BB133111EB + return z ^ (z >> 31) +} + +// Floats returns n uniform floats in [0, 1) with 53-bit resolution. +func Floats(g *Generator, n int) (*Array, error) { + if n < 0 { + return nil, errf("Floats: n must be zero or greater, got %d", n) + } + floats := make([]float64, n) + for i := range floats { + floats[i] = g.unit() + } + return &Array{shape: []int{n}, dt: Float, floats: floats}, nil +} + +// Ints returns n uniform ints in [min, max), drawn without modulo bias +// via Lemire's multiply-shift rejection. +func Ints(g *Generator, n int, min, max int64) (*Array, error) { + if n < 0 { + return nil, errf("Ints: n must be zero or greater, got %d", n) + } + if min >= max { + return nil, errf("Ints: min must be less than max, got %d and %d", min, max) + } + span := uint64(max) - uint64(min) + ints := make([]int64, n) + for i := range ints { + ints[i] = min + int64(g.bounded(span)) + } + return &Array{shape: []int{n}, dt: Int, ints: ints}, nil +} + +// bounded returns an unbiased uniform value in [0, n) for n > 0, via +// Lemire's multiply-shift rejection. The 128-bit product x*n splits into +// the candidate hi and residue lo; the rejected band has exactly +// t = 2^64 mod n residues, so a draw is accepted once lo >= t (lo >= n +// already implies that, deferring the modulo to the rare retry path). +func (g *Generator) bounded(n uint64) uint64 { + hi, lo := bits.Mul64(g.next(), n) + if lo < n { + t := -n % n // 2^64 mod n + for lo < t { + hi, lo = bits.Mul64(g.next(), n) + } + } + return hi +} + +// Normal returns n Gaussian draws of the given mean and standard +// deviation, paired Box-Muller transforms of the xoshiro stream, +// bit-stable across Go releases like every other draw. A +// negative or NaN std is an error. +func Normal(g *Generator, n int, mean, std float64) (*Array, error) { + if n < 0 { + return nil, errf("Normal: n must be zero or greater, got %d", n) + } + if std < 0 || math.IsNaN(std) { + return nil, errf("Normal: std must be zero or greater, got %v", std) + } + floats := make([]float64, n) + for i := range n { + floats[i] = mean + std*g.normalUnit() + } + return &Array{shape: []int{n}, dt: Float, floats: floats}, nil +} + +// normalUnit draws one standard normal value via the polar Box-Muller +// method: rejection until a point lands in the unit disk, then the +// spare-free transform. +func (g *Generator) normalUnit() float64 { + for { + u := 2*g.unit() - 1 + v := 2*g.unit() - 1 + if s := u*u + v*v; s < 1 && s > 0 { + return u * math.Sqrt(-2*math.Log(s)/s) + } + } +} + +// unit draws one uniform float in [0, 1) with 53-bit resolution. +func (g *Generator) unit() float64 { + return float64(g.next()>>11) / (1 << 53) +} + +// Permutation returns a uniform permutation of 0..n-1 (Fisher-Yates +// over unbiased bounded draws). Renamed from `Permute` to avoid +// confusion with `TransposeAxes` (which used to be called Permute +// before it became the axis-permutation variant). +func Permutation(g *Generator, n int) (*Array, error) { + if n < 0 { + return nil, errf("Permutation: n must be zero or greater, got %d", n) + } + ints := make([]int64, n) + for i := range ints { + ints[i] = int64(i) + } + for i := n - 1; i > 0; i-- { + j := g.bounded(uint64(i) + 1) + ints[i], ints[j] = ints[j], ints[i] + } + return &Array{shape: []int{n}, dt: Int, ints: ints}, nil +} + +// Shuffle returns a shuffled copy of a; the receiver itself is never +// touched. +func Shuffle(g *Generator, a *Array) *Array { + perm, _ := Permutation(g, a.Len()) + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(a.Len()) + for i := range a.Len() { + out.setFrom(i, a, int(perm.ints[i])) + } + return out +} + +// next advances xoshiro256++ and returns the next 64-bit value. +func (g *Generator) next() uint64 { + result := bits.RotateLeft64(g.s[0]+g.s[3], 23) + g.s[0] + t := g.s[1] << 17 + g.s[2] ^= g.s[0] + g.s[3] ^= g.s[1] + g.s[1] ^= g.s[2] + g.s[0] ^= g.s[3] + g.s[2] ^= t + g.s[3] = bits.RotateLeft64(g.s[3], 45) + return result +} + +// Float32s returns n uniform float32 values in [0, 1) with 24-bit +// resolution. +func Float32s(g *Generator, n int) (*Array, error) { + if n < 0 { + return nil, errf("Float32s: n must be zero or greater, got %d", n) + } + floats32 := make([]float32, n) + for i := range floats32 { + floats32[i] = float32(g.next()>>40) / (1 << 24) + } + return &Array{shape: []int{n}, dt: Float32, floats32: floats32}, nil +} + +// Unit draws one uniform float in [0, 1) with 53-bit resolution. +func (g *Generator) Unit() float64 { return g.unit() } + +// NormalUnit draws one standard normal value via the polar +// Box-Muller method. +func (g *Generator) NormalUnit() float64 { return g.normalUnit() } + +// Next advances the generator and returns the raw 64-bit value. +func (g *Generator) Next() uint64 { return g.next() } diff --git a/internal/core/random_test.go b/internal/core/random_test.go new file mode 100644 index 0000000..b59a43d --- /dev/null +++ b/internal/core/random_test.go @@ -0,0 +1,149 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +func TestGeneratorDeterminism(t *testing.T) { + g1 := NewGenerator(42) + g2 := NewGenerator(42) + f1, err := Floats(g1, 8) + if err != nil { + t.Fatalf("Floats: %v", err) + } + f2, err := Floats(g2, 8) + if err != nil { + t.Fatalf("Floats: %v", err) + } + if !Equal(f1, f2) { + t.Fatalf("the same seed must give identical arrays") + } + + g3 := NewGenerator(43) + f3, _ := Floats(g3, 8) + if Equal(f1, f3) { + t.Fatalf("different seeds must give different arrays") + } +} + +func TestGeneratorFloats(t *testing.T) { + g := NewGenerator(7) + f, err := Floats(g, 1000) + if err != nil { + t.Fatalf("Floats: %v", err) + } + if f.Dtype() != Float || f.Shape()[0] != 1000 { + t.Fatalf("Floats shape: %s", f) + } + for i := range f.Len() { + v, _ := FloatAt(f, i) + if v < 0 || v >= 1 { + t.Fatalf("Floats out of [0,1): %v", v) + } + } + // A thousand draws should cover both halves of the interval. + m, _ := Mean(f) + if m < 0.4 || m > 0.6 { + t.Fatalf("Floats mean suggests bias: %v", m) + } + + if _, err := Floats(g, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") { + t.Fatalf("Floats negative: %v", err) + } +} + +func TestGeneratorInts(t *testing.T) { + g := NewGenerator(9) + v, err := Ints(g, 600, 1, 4) + if err != nil { + t.Fatalf("Ints: %v", err) + } + if v.Dtype() != Int { + t.Fatalf("Ints dtype: %s", v.Dtype()) + } + seen := map[int64]bool{} + for i := range v.Len() { + x, _ := IntAt(v, i) + if x < 1 || x >= 4 { + t.Fatalf("Ints out of [1,4): %d", x) + } + seen[x] = true + } + // Six hundred draws over three values must hit all of them. + if len(seen) != 3 { + t.Fatalf("Ints coverage: %v", seen) + } + + if _, err := Ints(g, 1, 5, 5); err == nil || !strings.Contains(err.Error(), "min must be less than max") { + t.Fatalf("Ints range: %v", err) + } + if _, err := Ints(g, -1, 0, 1); err == nil || !strings.Contains(err.Error(), "zero or greater") { + t.Fatalf("Ints negative: %v", err) + } +} + +// TestGeneratorBoundedUniform pins the two defects of the old rejection +// rule: it rejected the whole [t, n) residue band, starving the upper +// buckets (for spans above 2^63 some values never appeared), and a full +// [MinInt64, MaxInt64) span accepted only lo == 2^64-1, hanging the draw. +func TestGeneratorBoundedUniform(t *testing.T) { + g := NewGenerator(11) + const draws = 30_000 + counts := map[int64]int{} + for range draws { + v, err := Ints(g, 1, 0, 3) + if err != nil { + t.Fatalf("Ints: %v", err) + } + x, _ := IntAt(v, 0) + counts[x]++ + } + // The counts are binomial around draws/3; the old biased rule kept + // bucket 2 a full draws/3 * (1/3) short, far outside this band. + want := draws / 3 + tol := want / 10 + for b := range int64(3) { + if diff := counts[b] - want; diff > tol || diff < -tol { + t.Fatalf("bucket %d drawn %d times, want %d +/- %d", b, counts[b], want, tol) + } + } + + // Full-span draws must terminate promptly and spread over the range. + g = NewGenerator(12) + seen := map[int64]bool{} + for range 64 { + v, err := Ints(g, 1, math.MinInt64, math.MaxInt64) + if err != nil { + t.Fatalf("full-span Ints: %v", err) + } + x, _ := IntAt(v, 0) + seen[x] = true + } + if len(seen) < 2 { + t.Fatalf("full-span draws collapsed to %d distinct value(s)", len(seen)) + } +} + +func TestGeneratorZeroSeed(t *testing.T) { + // Seed zero must not degenerate into the all-zero state. + g := NewGenerator(0) + f, err := Floats(g, 4) + if err != nil { + t.Fatalf("Floats from seed 0: %v", err) + } + nonzero := false + for i := range f.Len() { + v, _ := FloatAt(f, i) + if v != 0 { + nonzero = true + } + } + if !nonzero { + t.Fatalf("seed 0 produced only zeros") + } +} diff --git a/internal/core/randomsub_test.go b/internal/core/randomsub_test.go new file mode 100644 index 0000000..d2b0edf --- /dev/null +++ b/internal/core/randomsub_test.go @@ -0,0 +1,130 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// TestGeneratorStreamIsPinned pins the raw stream of one seed: the +// splitmix64 seeding walk and the xoshiro core must keep producing +// exactly these bits, because every recorded oracle digest downstream +// of the generator hangs on them. +func TestGeneratorStreamIsPinned(t *testing.T) { + g := NewGenerator(1) + for _, want := range []uint64{ + 0xcfc5d07f6f03c29b, 0xbf424132963fe08d, 0x19a37d5757aaf520, 0xbf08119f05cd56d6, + } { + if got := g.Next(); got != want { + t.Fatalf("Next() = %#016x, want %#016x", got, want) + } + } + n, err := Normal(NewGenerator(7), 3, 0, 1) + if err != nil { + t.Fatal(err) + } + for i, want := range []float64{1.674036445441065, 0.53789816819896552, 1.2079282540944534} { + if n.FloatAt(i) != want { + t.Fatalf("Normal draw %d = %.17g, want %.17g", i, n.FloatAt(i), want) + } + } +} + +// TestSplitmix64 pins the scalar mixer against an independent +// transcription of the finaliser and pins the state discipline: the +// returned state is the input plus the golden constant, the output is +// the mixed state, and chaining reproduces the walk. +func TestSplitmix64(t *testing.T) { + mix := func(z uint64) uint64 { + z = (z ^ (z >> 30)) * 0xBF58476D1CE4E5B9 + z = (z ^ (z >> 27)) * 0x94D049BB133111EB + return z ^ (z >> 31) + } + const golden = 0x9E3779B97F4A7C15 + state, v := Splitmix64(0) + if state != golden || v != mix(golden) { + t.Fatalf("Splitmix64(0) = (%#016x, %#016x), want (%#016x, %#016x)", state, v, uint64(golden), mix(golden)) + } + for _, s := range []uint64{0, 1, 0xdeadbeef, ^uint64(0)} { + state, v := Splitmix64(s) + if state != s+golden { + t.Fatalf("Splitmix64(%#016x) advanced the state to %#016x, want %#016x", s, state, s+golden) + } + if v != mix(state) { + t.Fatalf("Splitmix64(%#016x) output %#016x, want %#016x", s, v, mix(state)) + } + // Chaining from the returned state repeats the definition. + next, nv := Splitmix64(state) + if nv != mix(next) || next != state+golden { + t.Fatal("the chained step disagrees with the one-step definition") + } + } +} + +// mustSub builds a substream, failing the test on a bad index. +func mustSub(t *testing.T, seed int64, index int) *Generator { + t.Helper() + g, err := Substream(seed, index) + if err != nil { + t.Fatalf("Substream(%d, %d): %v", seed, index, err) + } + return g +} + +// draw returns n uniform floats from g. +func draw(t *testing.T, g *Generator, n int) []float64 { + t.Helper() + a, err := Floats(g, n) + if err != nil { + t.Fatal(err) + } + return a.RawFloats()[:n] +} + +// sameBits reports whether two draw prefixes are bit for bit equal. +func sameBits(a, b []float64) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if math.Float64bits(a[i]) != math.Float64bits(b[i]) { + return false + } + } + return true +} + +// TestSubstream pins the stream family: the same seed and index +// reproduce the same draws, distinct indices give distinct draws (so +// no member is a copy of another), the family differs from the plain +// seeded generator, and a negative index is an error. +func TestSubstream(t *testing.T) { + const draws = 16 + base := draw(t, mustSub(t, 42, 0), draws) + for i := 1; i < 8; i++ { + if sameBits(draw(t, mustSub(t, 42, i), draws), base) { + t.Fatalf("substream %d drew the same prefix as substream 0", i) + } + } + // Determinism: the same index redraws the same bits. + if !sameBits(draw(t, mustSub(t, 42, 3), draws), draw(t, mustSub(t, 42, 3), draws)) { + t.Fatal("the same substream drew different bits on a second construction") + } + // The family hangs off the seed, not beside the plain generator: + // substream 0 must differ from NewGenerator(42) and from the + // neighbouring seed's plain stream. + for _, seed := range []int64{42, 43} { + if sameBits(draw(t, NewGenerator(seed), draws), base) { + t.Fatalf("substream 0 coincides with NewGenerator(%d)", seed) + } + } + // A large index is as legal as a small one. + if _, err := Substream(42, 1<<40); err != nil { + t.Fatalf("Substream at index 2^40: %v", err) + } + if _, err := Substream(42, -1); err == nil { + t.Fatal("a negative index: want an error") + } +} diff --git a/internal/core/reduce.go b/internal/core/reduce.go new file mode 100644 index 0000000..2782061 --- /dev/null +++ b/internal/core/reduce.go @@ -0,0 +1,1494 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "fmt" + "math" + "strconv" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Scalar boxes a single numeric result whose dtype is known only at +// runtime: Sum of an int array is an int, of a float array a +// float, of a complex array a complex. Int, Float and Complex unpack it; +// IsFloat and IsComplex say which. +type Scalar struct { + isFloat bool + isComplex bool + i int64 + f float64 + c complex128 +} + +// Int returns the scalar as an int64. It converts a float or complex +// scalar by truncation, mirroring Go's conversions; use IsFloat and +// IsComplex when the distinction matters. +func (s Scalar) Int() int64 { + if s.isFloat { + return int64(s.f) + } + if s.isComplex { + return int64(real(s.c)) + } + return s.i +} + +// Float returns the scalar as a float64, converting int and complex +// scalars (a complex scalar contributes its real part). +func (s Scalar) Float() float64 { + if s.isFloat { + return s.f + } + if s.isComplex { + return real(s.c) + } + return float64(s.i) +} + +// Complex returns the scalar as a complex128, converting real scalars. +func (s Scalar) Complex() complex128 { + if s.isComplex { + return s.c + } + if s.isFloat { + return complex(s.f, 0) + } + return complex(float64(s.i), 0) +} + +// IsFloat reports whether the scalar came from a float computation. +func (s Scalar) IsFloat() bool { return s.isFloat } + +// IsComplex reports whether the scalar came from a complex computation. +func (s Scalar) IsComplex() bool { return s.isComplex } + +// String renders the scalar with its dtype, as in "int 7" or +// "complex (4-2i)". The complex form is the one %v prints, matching +// Array.String: a single sign between the parts, never "+-". +func (s Scalar) String() string { + switch { + case s.isComplex: + return "complex " + fmt.Sprintf("%v", s.c) + case s.isFloat: + return "float " + strconv.FormatFloat(s.f, 'g', -1, 64) + default: + return "int " + strconv.FormatInt(s.i, 10) + } +} + +// Sum returns the sum of all elements. An int sum wraps on overflow like +// Go's int64 arithmetic; the narrow integer widths widen +// exactly and accumulate in int64, answered as an Int scalar, and a bool +// sum counts its true elements into the same Int scalar; float16 and +// float32 sums accumulate in float64 and answer a float scalar; +// an empty array sums to zero. +// +// The fold is partitioned by length alone, never by the worker count, so +// the same input answers the same scalar on any machine and under any +// worker setting: the range is cut into fixed chunks, each chunk +// accumulates into four interleaved partials, and the chunk results +// combine through a balanced pairwise tree over the chunk indices +// (treeSum). The tree is the shape the distributed reduction shares: a +// range cut at chunk boundaries into shards answers the same bits as +// this fold, because shard partials and array chunks are entries of one +// tree combined by one function. The integer sum is exact under any +// order. The floating-point fold rounds differently from the single +// chain it replaced, not always in the same direction: the interleaved +// partials shorten each dependency chain and the tree adds a logarithmic +// number of combination roundings, which measured within a small factor +// of the chain at every length tried and better than it once the chain +// is long enough for its own roundings to accumulate. Only the array's +// own elements take part: a rebased view's payload may run past its +// element count, and those invisible tail slots never contribute. +func Sum(a *Array) Scalar { + n := a.Len() + switch a.dt { + case Int: + return Scalar{i: intFoldSum(a.ints[:n])} + case Bool: + // The bool sum counts the true elements, answered as the Int + // scalar every integer-class reduction answers. + return Scalar{i: boolFoldCount(a.bools[:n])} + case Int8: + return Scalar{i: intFoldSumNarrow(a.i8s[:n])} + case Uint8: + return Scalar{i: intFoldSumNarrow(a.u8s[:n])} + case Int16: + return Scalar{i: intFoldSumNarrow(a.i16s[:n])} + case Uint16: + return Scalar{i: intFoldSumNarrow(a.u16s[:n])} + case Int32: + return Scalar{i: intFoldSumNarrow(a.i32s[:n])} + case Uint32: + return Scalar{i: intFoldSumNarrow(a.u32s[:n])} + case Float16: + // The half sum accumulates in float64 and answers a float + // scalar, exactly as the float32 sum does: every + // widening is exact, so the fold sees the same addends an + // accessor walk would hand over. + return Scalar{isFloat: true, f: floatFoldSumF16(a.halves[:n])} + case Float32: + return Scalar{isFloat: true, f: floatFoldSumF32(a.floats32[:n])} + case Float: + return Scalar{isFloat: true, f: floatFoldSum(a.floats[:n])} + default: + return Scalar{isComplex: true, c: complexFoldSum(a.complexes[:n])} + } +} + +// foldChunk is the element count one partial fold carries. It is a +// constant of the arithmetic: the partition follows from the length +// alone, never from the worker count, so the same input answers the +// same scalar on any machine and under any worker setting. +const foldChunk = 1 << 16 + +// foldParts is the number of chunks a length is cut into, bounded so the +// partial table stays small for very long arrays. +func foldParts(n int) int { + if n <= foldChunk { + return 1 + } + parts := min((n+foldChunk-1)/foldChunk, 1<<12) + return parts +} + +// FoldParts reports how many blocks the canonical reduction partition +// cuts a range of n elements into. The partition is a function of the +// length alone, so it is the same on every machine and under any worker +// setting; the spmd package cuts distributed data on these boundaries, +// which is what makes a sharded reduction compose into the single-array +// fold's exact bits. +func FoldParts(n int) int { return foldParts(n) } + +// FoldBoundary reports the index where block c of the canonical +// partition of n elements begins; block c spans [FoldBoundary(n, c), +// FoldBoundary(n, c+1)). Block 0 begins at 0 and FoldBoundary(n, +// FoldParts(n)) is n. +func FoldBoundary(n, c int) int { return c * n / foldParts(n) } + +// TreeSum combines reduction partials through the balanced pairwise +// tree the folds combine their chunk partials with: the node over a +// range splits at its midpoint and adds the two halves' nodes, left +// before right. The value of any contiguous range of partials is one +// node of the tree, so partials gathered from a sharded cut combine +// into the single-array fold's exact bits; the spmd package feeds it +// the block values the shards folded. +func TreeSum[N int64 | float64 | complex128](vals []N) N { return treeSum(vals) } + +// FoldRange is one canonical block's fold over a float64 payload: the +// four interleaved chains combined as a balanced pair. Shards fold +// their own blocks with it, so a block's value is the same bits +// wherever its elements live. +func FoldRange(src []float64) float64 { return foldRange(src) } + +// FoldRangeF32 is FoldRange over a float32 payload; every widening is +// exact. +func FoldRangeF32(src []float32) float64 { return foldRangeF32(src) } + +// FoldRangeF16 is FoldRange over a half payload's raw bit patterns; +// every widening is exact. +func FoldRangeF16(src []uint16) float64 { return foldRangeF16(src) } + +// FoldRangeC128 is FoldRange over a complex payload. +func FoldRangeC128(src []complex128) complex128 { return foldRangeC128(src) } + +// ExtremeRange is the serial extremum rule on one canonical block of a +// float64 payload: seed past the leading NaNs, keep the first strictly +// better value, and report whether the block holds a candidate at all. +func ExtremeRange(src []float64, greater bool) (float64, bool) { return floatExtreme(src, greater) } + +// ExtremeRangeF32 is ExtremeRange over a float32 payload; every +// widening is exact. +func ExtremeRangeF32(src []float32, greater bool) (float64, bool) { + return float32Extreme(src, greater) +} + +// ExtremeRangeF16 is ExtremeRange over a half payload's raw bit +// patterns; every widening is exact. +func ExtremeRangeF16(src []uint16, greater bool) (float64, bool) { return halfExtreme(src, greater) } + +// CombineExtrema combines the blocks' extrema with the serial walk's +// strict comparison in index order: a tie keeps the earlier block's +// value, a block with no candidate contributes nothing, and an array +// with no candidate anywhere answers the last block's fallback, which +// is the whole array's last element. The extrema the spmd shards fold +// combine through it, so a sharded extremum is the single-array +// extremum bit for bit. +func CombineExtrema(vals []float64, oks []bool, greater bool) float64 { + return combineExtreme(vals, oks, greater) +} + +// FloatScalar boxes a float64 as the scalar the float reductions +// answer. +func FloatScalar(f float64) Scalar { return Scalar{isFloat: true, f: f} } + +// IntScalar boxes an int64 as the scalar the integer-class reductions +// answer. +func IntScalar(i int64) Scalar { return Scalar{i: i} } + +// ComplexScalar boxes a complex128 as the scalar the complex +// reductions answer. +func ComplexScalar(c complex128) Scalar { return Scalar{isComplex: true, c: c} } + +// TreeProd combines partial products through the balanced midpoint +// tree, multiplying the left half before the right: partials gathered +// from a sharded cut combine into the single-array product's exact +// bits. The float32 partials multiply natively in float32, the way the +// float32 product fold keeps. +func TreeProd[N int64 | float64 | float32](vals []N) N { return treeProd(vals) } + +// TreeProdHalf combines half-precision partial products: the values +// carry exact half bits in float64 and every combine narrows through +// half, the per-step rounding the half product fold keeps. +func TreeProdHalf(vals []float64) float64 { return treeProdHalf(vals) } + +// FoldProd is one canonical block's product over a float64 payload. +// Shards fold their own blocks with it, so a block's product is the +// same bits wherever its elements live. +func FoldProd(src []float64) float64 { return foldProdRange(src) } + +// FoldProdF32 is FoldProd over a float32 payload, multiplied natively +// in float32. +func FoldProdF32(src []float32) float32 { return foldProdRangeF32(src) } + +// FoldProdF16 is FoldProd over a half payload's raw bit patterns, with +// the per-step half rounding the line product keeps; the answer is an +// exact half value carried in float64. +func FoldProdF16(src []uint16) float64 { return foldProdRangeF16(src) } + +// FoldProdI64 is FoldProd over an int64 payload; the wrapping product +// is exact under any grouping. +func FoldProdI64(src []int64) int64 { return foldProdRangeI64(src) } + +// FoldNormPower is one canonical block's sum of |v|^p over a float64 +// payload, with the per-element arithmetic the norm fold keeps. +func FoldNormPower(src []float64, p float64) float64 { return foldNormPowerRange(src, p) } + +// FoldNormPowerF32 is FoldNormPower over a float32 payload; every +// widening is exact. +func FoldNormPowerF32(src []float32, p float64) float64 { return foldNormPowerRangeF32(src, p) } + +// FoldNormPowerF16 is FoldNormPower over a half payload's raw bit +// patterns; every widening is exact. +func FoldNormPowerF16(src []uint16, p float64) float64 { return foldNormPowerRangeF16(src, p) } + +// FoldNormPowerI64 is FoldNormPower over an int64 payload. +func FoldNormPowerI64(src []int64, p float64) float64 { return foldNormPowerRangeI64(src, p) } + +// NormRoot closes a power sum into the norm: Sqrt for p = 2, the sum +// for p = 1, Pow of the sum for every other exponent. +func NormRoot(sum, p float64) float64 { return normRoot(sum, p) } + +// FoldDot is one canonical block's dot product over float64 payloads +// of equal length. +func FoldDot(x, y []float64) float64 { return foldDotRange(x, y) } + +// FoldDotF32 is FoldDot over float32 payloads; every product is exact +// in float64. +func FoldDotF32(x, y []float32) float64 { return foldDotRangeF32(x, y) } + +// FoldDotF16 is FoldDot over half payloads' raw bit patterns; every +// product is exact in float64. +func FoldDotF16(x, y []uint16) float64 { return foldDotRangeF16(x, y) } + +// FoldDotC128 is FoldDot over complex payloads. +func FoldDotC128(x, y []complex128) complex128 { return foldDotRangeC128(x, y) } + +// FoldDotI64 is FoldDot over int64 payloads: the wrapping products and +// the wrapping sums are exact under any grouping. +func FoldDotI64(x, y []int64) int64 { + var s int64 + for i := range x { + s += x[i] * y[i] + } + return s +} + +// foldRange runs one chunk's fold over a float payload: four interleaved +// chains, so the adds of a long array overlap instead of queueing on one +// adder, combined as ((s0+s1)+(s2+s3)), the pairing a balanced tree +// gives. Each chain carries a quarter of the chunk, so the dependency +// chain is four times shorter and the combination costs three roundings; +// against the single chain it replaced the total error measured within a +// small factor either way, so this form is chosen for the throughput and +// not for a claimed accuracy win. +func foldRange(src []float64) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= len(src); i += 4 { + s0 += src[i] + s1 += src[i+1] + s2 += src[i+2] + s3 += src[i+3] + } + for ; i < len(src); i++ { + s0 += src[i] + } + return (s0 + s1) + (s2 + s3) +} + +// treeSum combines chunk partials with a balanced pairwise tree over +// their indices: the node over a range splits at its midpoint and adds +// the two halves' nodes, left before right. The shape depends on nothing +// but the partial count, and the value of any contiguous range of +// partials is one node of the tree. That property is what lets a +// reduction cut at chunk boundaries into shards reproduce this fold +// exactly, for any number of shards: the spmd package gathers the shard +// partials and combines them with this same function. +func treeSum[N int64 | float64 | complex128](vals []N) N { + return treeSumRange(vals, 0, len(vals)) +} + +// treeSumRange is one node of the treeSum tree over the partials [l, r). +func treeSumRange[N int64 | float64 | complex128](vals []N, l, r int) N { + if r-l == 1 { + return vals[l] + } + m := l + (r-l)/2 + return treeSumRange(vals, l, m) + treeSumRange(vals, m, r) +} + +// treeProd combines partial products through the balanced midpoint +// tree treeSum uses, multiplying the left half before the right. The +// shape depends on nothing but the partial count, so partial products +// gathered from a sharded cut combine into the single-array product's +// exact bits; the spmd package feeds it the block values the shards +// folded. +func treeProd[N int64 | float64 | float32](vals []N) N { + return treeProdRange(vals, 0, len(vals)) +} + +func treeProdRange[N int64 | float64 | float32](vals []N, l, r int) N { + if r-l == 1 { + return vals[l] + } + m := l + (r-l)/2 + return treeProdRange(vals, l, m) * treeProdRange(vals, m, r) +} + +// treeProdHalf is treeProd for a half-precision product: the partials +// carry exact half values and every combine narrows through half, the +// per-step rounding the half product fold keeps. +func treeProdHalf(vals []float64) float64 { + return treeProdHalfRange(vals, 0, len(vals)) +} + +func treeProdHalfRange(vals []float64, l, r int) float64 { + if r-l == 1 { + return vals[l] + } + m := l + (r-l)/2 + return halfRound(treeProdHalfRange(vals, l, m) * treeProdHalfRange(vals, m, r)) +} + +// halfRound narrows through half and back, the per-step rounding the +// half precision folds keep. +func halfRound(f float64) float64 { + return HalfToFloat64(HalfFromFloat64(f)) +} + +// foldProdRangeI64 is one canonical block's product over an int64 +// payload: wrapping multiplication is associative, so the grouping +// cannot change the value. +func foldProdRangeI64(src []int64) int64 { + m := int64(1) + for _, v := range src { + m *= v + } + return m +} + +// foldProdRange is one canonical block's product over a float64 +// payload: a single chain from one, the association the product fold +// keeps inside a block. +func foldProdRange(src []float64) float64 { + m := 1.0 + for _, v := range src { + m *= v + } + return m +} + +// foldProdRangeF32 is one canonical block's product over a float32 +// payload, multiplied natively in float32. +func foldProdRangeF32(src []float32) float32 { + m := float32(1) + for _, v := range src { + m *= v + } + return m +} + +// foldProdRangeF16 is one canonical block's product over a half +// payload: the running product narrows through half before every +// combine, exactly the per-step rounding the line fold keeps, and the +// answer is an exact half value carried in float64. +func foldProdRangeF16(src []uint16) float64 { + m := 1.0 + for _, v := range src { + m = halfRound(m * HalfToFloat64(v)) + } + return m +} + +// foldNormPowerRange is one canonical block's sum of |v|^p over a +// float64 payload, with the same per-element arithmetic the norm fold +// keeps: the bare absolute for p = 1, the squared absolute for p = 2, +// and Pow of the absolute for every other exponent. +func foldNormPowerRange(src []float64, p float64) float64 { + switch { + case p == 1: + var acc float64 + for _, v := range src { + acc += math.Abs(v) + } + return acc + case p == 2: + var acc float64 + for _, v := range src { + w := math.Abs(v) + acc += w * w + } + return acc + default: + var acc float64 + for _, v := range src { + acc += math.Pow(math.Abs(v), p) + } + return acc + } +} + +// foldNormPowerRangeF32 is foldNormPowerRange over a float32 payload; +// every widening is exact. +func foldNormPowerRangeF32(src []float32, p float64) float64 { + switch { + case p == 1: + var acc float64 + for _, v := range src { + acc += math.Abs(float64(v)) + } + return acc + case p == 2: + var acc float64 + for _, v := range src { + w := math.Abs(float64(v)) + acc += w * w + } + return acc + default: + var acc float64 + for _, v := range src { + acc += math.Pow(math.Abs(float64(v)), p) + } + return acc + } +} + +// foldNormPowerRangeF16 is foldNormPowerRange over a half payload's +// raw bit patterns; every widening is exact. +func foldNormPowerRangeF16(src []uint16, p float64) float64 { + switch { + case p == 1: + var acc float64 + for _, v := range src { + acc += math.Abs(HalfToFloat64(v)) + } + return acc + case p == 2: + var acc float64 + for _, v := range src { + w := math.Abs(HalfToFloat64(v)) + acc += w * w + } + return acc + default: + var acc float64 + for _, v := range src { + acc += math.Pow(math.Abs(HalfToFloat64(v)), p) + } + return acc + } +} + +// foldNormPowerRangeI64 is foldNormPowerRange over an int64 payload. +func foldNormPowerRangeI64(src []int64, p float64) float64 { + switch { + case p == 1: + var acc float64 + for _, v := range src { + acc += math.Abs(float64(v)) + } + return acc + case p == 2: + var acc float64 + for _, v := range src { + w := math.Abs(float64(v)) + acc += w * w + } + return acc + default: + var acc float64 + for _, v := range src { + acc += math.Pow(math.Abs(float64(v)), p) + } + return acc + } +} + +// normRoot closes a power sum into the norm: Sqrt for p = 2, the sum +// itself for p = 1, and Pow of the sum for every other exponent, the +// closing the norm fold keeps. +func normRoot(sum, p float64) float64 { + switch { + case p == 2: + return math.Sqrt(sum) + case p == 1: + return sum + default: + return math.Pow(sum, 1/p) + } +} + +// foldDotRange is one canonical block's dot product over float64 +// payloads of equal length: four interleaved product chains, combined +// as a balanced pair. +func foldDotRange(x, y []float64) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= len(x); i += 4 { + s0 += x[i] * y[i] + s1 += x[i+1] * y[i+1] + s2 += x[i+2] * y[i+2] + s3 += x[i+3] * y[i+3] + } + for ; i < len(x); i++ { + s0 += x[i] * y[i] + } + return (s0 + s1) + (s2 + s3) +} + +// foldDotRangeF32 is foldDotRange over float32 payloads; every product +// is exact in float64. +func foldDotRangeF32(x, y []float32) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= len(x); i += 4 { + s0 += float64(x[i]) * float64(y[i]) + s1 += float64(x[i+1]) * float64(y[i+1]) + s2 += float64(x[i+2]) * float64(y[i+2]) + s3 += float64(x[i+3]) * float64(y[i+3]) + } + for ; i < len(x); i++ { + s0 += float64(x[i]) * float64(y[i]) + } + return (s0 + s1) + (s2 + s3) +} + +// foldDotRangeF16 is foldDotRange over half payloads' raw bit +// patterns; every product is exact in float64. +func foldDotRangeF16(x, y []uint16) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= len(x); i += 4 { + s0 += HalfToFloat64(x[i]) * HalfToFloat64(y[i]) + s1 += HalfToFloat64(x[i+1]) * HalfToFloat64(y[i+1]) + s2 += HalfToFloat64(x[i+2]) * HalfToFloat64(y[i+2]) + s3 += HalfToFloat64(x[i+3]) * HalfToFloat64(y[i+3]) + } + for ; i < len(x); i++ { + s0 += HalfToFloat64(x[i]) * HalfToFloat64(y[i]) + } + return (s0 + s1) + (s2 + s3) +} + +// foldDotRangeC128 is foldDotRange over complex payloads. +func foldDotRangeC128(x, y []complex128) complex128 { + var s0, s1, s2, s3 complex128 + i := 0 + for ; i+4 <= len(x); i += 4 { + s0 += x[i] * y[i] + s1 += x[i+1] * y[i+1] + s2 += x[i+2] * y[i+2] + s3 += x[i+3] * y[i+3] + } + for ; i < len(x); i++ { + s0 += x[i] * y[i] + } + return (s0 + s1) + (s2 + s3) +} + +// floatFoldSum sums a float64 payload over the fixed partition. +func floatFoldSum(src []float64) float64 { + parts := foldParts(len(src)) + if parts == 1 { + return foldRange(src) + } + partials := make([]float64, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + partials[c] = foldRange(src[c*len(src)/parts : (c+1)*len(src)/parts]) + } + }) + return treeSum(partials) +} + +// foldRangeF32 is one block's fold over a float32 payload: four +// interleaved float64 chains, every widening exact. +func foldRangeF32(part []float32) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= len(part); i += 4 { + s0 += float64(part[i]) + s1 += float64(part[i+1]) + s2 += float64(part[i+2]) + s3 += float64(part[i+3]) + } + for ; i < len(part); i++ { + s0 += float64(part[i]) + } + return (s0 + s1) + (s2 + s3) +} + +// foldRangeF16 is one block's fold over a half payload's raw bit +// patterns, widening each element exactly as it is read. +func foldRangeF16(part []uint16) float64 { + var s0, s1, s2, s3 float64 + i := 0 + for ; i+4 <= len(part); i += 4 { + s0 += HalfToFloat64(part[i]) + s1 += HalfToFloat64(part[i+1]) + s2 += HalfToFloat64(part[i+2]) + s3 += HalfToFloat64(part[i+3]) + } + for ; i < len(part); i++ { + s0 += HalfToFloat64(part[i]) + } + return (s0 + s1) + (s2 + s3) +} + +// floatFoldSumF32 sums a float32 payload: every widening to float64 is +// exact, so the fold sees the values an accessor walk would hand over. +func floatFoldSumF32(src []float32) float64 { + parts := foldParts(len(src)) + return sumOverParts(len(src), parts, func(c int) float64 { + return foldRangeF32(src[c*len(src)/parts : (c+1)*len(src)/parts]) + }) +} + +// floatFoldSumF16 sums a half-precision payload, widening each element +// exactly as it is read. +func floatFoldSumF16(src []uint16) float64 { + parts := foldParts(len(src)) + return sumOverParts(len(src), parts, func(c int) float64 { + return foldRangeF16(src[c*len(src)/parts : (c+1)*len(src)/parts]) + }) +} + +// sumOverParts runs the per-chunk fold and combines the results through +// the balanced partial tree: the only order-dependent step, and its shape +// depends on nothing but the chunk count. +func sumOverParts(n, parts int, fold func(c int) float64) float64 { + if parts == 1 { + return fold(0) + } + partials := make([]float64, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + partials[c] = fold(c) + } + }) + return treeSum(partials) +} + +// intFoldSum sums an int64 payload: wrapping addition is associative, so +// the partition cannot change the value and a plain chunked scan needs no +// interleaved partials. +func intFoldSum(src []int64) int64 { + parts := foldParts(len(src)) + if parts == 1 { + var s int64 + for _, v := range src { + s += v + } + return s + } + partials := make([]int64, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + var s int64 + for _, v := range src[c*len(src)/parts : (c+1)*len(src)/parts] { + s += v + } + partials[c] = s + } + }) + var total int64 + for _, v := range partials { + total += v + } + return total +} + +// intFoldSumNarrow sums a narrow integer payload into int64 through the +// fixed partition intFoldSum uses; every widening is exact, so the sum +// stays machine-independent and the accumulation order changes nothing. +func intFoldSumNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T) int64 { + parts := foldParts(len(src)) + if parts == 1 { + var s int64 + for _, v := range src { + s += int64(v) + } + return s + } + partials := make([]int64, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + var s int64 + for _, v := range src[c*len(src)/parts : (c+1)*len(src)/parts] { + s += int64(v) + } + partials[c] = s + } + }) + var total int64 + for _, v := range partials { + total += v + } + return total +} + +// boolFoldCount counts the true elements of a bool payload through the +// same fixed partition intFoldSum uses. +func boolFoldCount(src []bool) int64 { + parts := foldParts(len(src)) + if parts == 1 { + var s int64 + for _, v := range src { + if v { + s++ + } + } + return s + } + partials := make([]int64, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + var s int64 + for _, v := range src[c*len(src)/parts : (c+1)*len(src)/parts] { + if v { + s++ + } + } + partials[c] = s + } + }) + var total int64 + for _, v := range partials { + total += v + } + return total +} + +// foldRangeC128 is one block's fold over a complex payload: four +// interleaved chains, combined as a balanced pair. +func foldRangeC128(part []complex128) complex128 { + var s0, s1, s2, s3 complex128 + i := 0 + for ; i+4 <= len(part); i += 4 { + s0 += part[i] + s1 += part[i+1] + s2 += part[i+2] + s3 += part[i+3] + } + for ; i < len(part); i++ { + s0 += part[i] + } + return (s0 + s1) + (s2 + s3) +} + +// complexFoldSum sums a complex payload the way the float fold sums a +// real one: fixed chunks, four interleaved partials, partials combined +// through the balanced partial tree. +func complexFoldSum(src []complex128) complex128 { + parts := foldParts(len(src)) + if parts == 1 { + return foldRangeC128(src) + } + partials := make([]complex128, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + partials[c] = foldRangeC128(src[c*len(src)/parts : (c+1)*len(src)/parts]) + } + }) + return treeSum(partials) +} + +// Min returns the smallest element; an empty array, or a complex array +// (no ordering), is an error. +func Min(a *Array) (Scalar, error) { + return a.reduceOrder("Min", false) +} + +// Max returns the largest element; an empty array, or a complex array +// (no ordering), is an error. +func Max(a *Array) (Scalar, error) { + return a.reduceOrder("Max", true) +} + +// Mean returns the arithmetic mean of all elements as a float64, +// computed in float64 even for int arrays (true division); an empty or +// complex array is an error. +// +// The sum runs through Sum, so every ordinary input rounds exactly as +// the standalone Sum does. The one exception is the input whose largest +// magnitude could carry the running total past float64's range: there +// the plain sum overflows to an Inf the final division cannot undo +// (Mean of two MaxFloat64s answered +Inf). Such an input takes a scaled +// two-pass form instead: the elements are summed divided by the largest +// magnitude and the result rescaled, which stays finite. The scale is +// chosen by the pre-pass below, so any input below the overflow +// threshold, and every int array whose whole range sits far below +// float64's ceiling, keeps the plain path and its digest. +func Mean(a *Array) (float64, error) { + if a.dt == Complex { + return 0, errf("Mean: complex arrays have no float mean") + } + if a.Len() == 0 { + return 0, errf("Mean: an empty array has no mean") + } + n := a.Len() + fn := float64(n) + // Every widening in this pre-pass is exact, and NaN fails the + // comparison, so a NaN payload never triggers the scaled path and + // reaches the caller through Sum's fold as before. The walk reads the + // payload directly, which is the value the accessor would hand over. + maxAbs := a.maxAbs(n) + if maxAbs > math.MaxFloat64/fn && !math.IsInf(maxAbs, 0) { + var acc float64 + for i := range n { + acc += a.floatAt(i) / maxAbs + } + // The rescale divides by n first: multiplying acc by maxAbs + // before the division could overflow again, which is the very + // failure this path exists to avoid. + return acc / fn * maxAbs, nil + } + return Sum(a).Float() / fn, nil +} + +// maxAbs returns the largest magnitude among the first n elements, the +// seed pre-pass the scaled mean path needs. The dtype dispatch sits +// outside the walk and float16 and float32 widen exactly, so every +// value is the one an accessor read would hand over, at two +// instructions per element instead of a switch on the dtype. +// +// A magnitude maximum is associative and commutative, and it has no sign +// to lose: |+0| and |−0| are the same zero, so unlike the signed extrema +// the answer cannot depend on the visit order. The walk therefore takes +// the same fixed partition the sums use and combines the partial maxima +// with max, which makes it both parallel and deterministic. +func (a *Array) maxAbs(n int) float64 { + if a.strides != nil { + parts := foldParts(n) + if parts == 1 { + return a.maxAbsSerial(n) + } + partials := make([]float64, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + lo, hi := c*n/parts, (c+1)*n/parts + m := 0.0 + for i := lo; i < hi; i++ { + if w := math.Abs(a.floatAt(i)); w > m { + m = w + } + } + partials[c] = m + } + }) + m := 0.0 + for _, v := range partials { + m = math.Max(m, v) + } + return m + } + switch a.dt { + case Int: + return foldMaxAbsParts(n, func(lo, hi int) float64 { + m := 0.0 + for _, v := range a.ints[lo:hi] { + if w := math.Abs(float64(v)); w > m { + m = w + } + } + return m + }) + case Float16: + return foldMaxAbsParts(n, func(lo, hi int) float64 { + m := 0.0 + for _, v := range a.halves[lo:hi] { + if w := math.Abs(HalfToFloat64(v)); w > m { + m = w + } + } + return m + }) + case Float32: + return foldMaxAbsParts(n, func(lo, hi int) float64 { + m := 0.0 + for _, v := range a.floats32[lo:hi] { + if w := math.Abs(float64(v)); w > m { + m = w + } + } + return m + }) + case Float: + return foldMaxAbsParts(n, func(lo, hi int) float64 { + m := 0.0 + for _, v := range a.floats[lo:hi] { + if w := math.Abs(v); w > m { + m = w + } + } + return m + }) + case Bool: + return foldMaxAbsParts(n, func(lo, hi int) float64 { + for _, v := range a.bools[lo:hi] { + if v { + return 1 + } + } + return 0 + }) + case Int8: + return maxAbsNarrow(a.i8s, n) + case Uint8: + return maxAbsNarrow(a.u8s, n) + case Int16: + return maxAbsNarrow(a.i16s, n) + case Uint16: + return maxAbsNarrow(a.u16s, n) + case Int32: + return maxAbsNarrow(a.i32s, n) + case Uint32: + return maxAbsNarrow(a.u32s, n) + } + // Complex never reaches here: Mean rejects it before the guard. + return 0 +} + +// maxAbsNarrow is maxAbs's chunk walk for a narrow integer payload: +// every widening is exact, so the magnitude sees the accessor value. +func maxAbsNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int) float64 { + return foldMaxAbsParts(n, func(lo, hi int) float64 { + m := 0.0 + for _, v := range src[lo:hi] { + if w := math.Abs(float64(v)); w > m { + m = w + } + } + return m + }) +} + +// foldMaxAbsParts runs the magnitude maximum over the fixed partition of +// n elements. The chunk closure walks a payload slice, never a function +// per element, and the partial maxima combine with max, so the partition +// changes nothing. +func foldMaxAbsParts(n int, chunk func(lo, hi int) float64) float64 { + parts := foldParts(n) + if parts == 1 { + return chunk(0, n) + } + partials := make([]float64, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + partials[c] = chunk(c*n/parts, (c+1)*n/parts) + } + }) + m := 0.0 + for _, v := range partials { + m = math.Max(m, v) + } + return m +} + +// maxAbsSerial is the strided fallback for a length below the partition +// floor. +func (a *Array) maxAbsSerial(n int) float64 { + var m float64 + for i := range n { + if v := math.Abs(a.floatAt(i)); v > m { + m = v + } + } + return m +} + +// Dot returns the dot product of two 1-D arrays of equal length. +// Integer-class pairs produce an int scalar accumulated in int64 (the +// products wrap like the int64 kernel's); bool pairs are refused: bool +// carries no arithmetic. Float32 operands accumulate in float64 and +// produce a float scalar; any float64 operand promotes the +// result to float, any complex operand to complex. +func Dot(a, b *Array) (Scalar, error) { + if a.NDim() != 1 || b.NDim() != 1 { + return Scalar{}, errf("Dot: needs 1-D arrays, got shapes %s and %s", + shapeText(a.shape), shapeText(b.shape)) + } + if a.Len() != b.Len() { + return Scalar{}, errf("Dot: length mismatch %d vs %d", a.Len(), b.Len()) + } + n := a.Len() + dt := promote(a.dt, b.dt) + if dt == Bool { + return Scalar{}, errf("Dot: bool arrays have no arithmetic") + } + switch { + case dt == Int && a.dt == Int && b.dt == Int: + return Scalar{i: intFoldDot(a.ints[:n], b.ints[:n])}, nil + case intClass(dt): + // Every integer-class pair widens exactly into int64; the + // products wrap exactly as the int64 kernel's do, and wrapping + // addition is associative, so the fixed partition changes + // nothing. + if a.dt == b.dt && a.dt != Bool && a.isContiguous() && b.isContiguous() { + switch a.dt { + case Int: + return Scalar{i: intFoldDot(a.ints[:n], b.ints[:n])}, nil + case Int8: + return Scalar{i: narrowFoldDot(a.i8s[:n], b.i8s[:n])}, nil + case Uint8: + return Scalar{i: narrowFoldDot(a.u8s[:n], b.u8s[:n])}, nil + case Int16: + return Scalar{i: narrowFoldDot(a.i16s[:n], b.i16s[:n])}, nil + case Uint16: + return Scalar{i: narrowFoldDot(a.u16s[:n], b.u16s[:n])}, nil + case Int32: + return Scalar{i: narrowFoldDot(a.i32s[:n], b.i32s[:n])}, nil + case Uint32: + return Scalar{i: narrowFoldDot(a.u32s[:n], b.u32s[:n])}, nil + } + } + parts := foldParts(n) + return Scalar{i: foldPartsOver(n, parts, func(c int) int64 { + lo, hi := c*n/parts, (c+1)*n/parts + var s int64 + for i := lo; i < hi; i++ { + s += a.intAt(i) * b.intAt(i) + } + return s + })}, nil + } + switch dt { + case Float16: + var s float64 + if a.dt == Float16 && b.dt == Float16 { + // Both payloads hold half bit patterns at their flat index, + // so the kernel streams the raw slices; every product is + // exact in float64 either way. + as, bs := a.halves[:n], b.halves[:n] + s = foldDotF16(as, bs) + } else { + for i := range n { + s += a.floatAt(i) * b.floatAt(i) + } + } + return Scalar{isFloat: true, f: s}, nil + case Float32: + var s float64 + if a.dt == Float32 && b.dt == Float32 { + // Both payloads hold float32 at their flat index, so the + // kernel streams the raw slices; every product is exact in + // float64 either way, so the values match accessor reads. + as, bs := a.floats32[:n], b.floats32[:n] + s = foldDotF32(as, bs) + } else { + for i := range n { + s += a.floatAt(i) * b.floatAt(i) + } + } + return Scalar{isFloat: true, f: s}, nil + case Float: + var s float64 + if a.dt == Float && b.dt == Float { + as, bs := a.floats[:n], b.floats[:n] + s = foldDotF64(as, bs) + } else { + for i := range n { + s += a.floatAt(i) * b.floatAt(i) + } + } + return Scalar{isFloat: true, f: s}, nil + default: + var s complex128 + if a.dt == Complex && b.dt == Complex { + as, bs := a.complexes[:n], b.complexes[:n] + s = foldDotC128(as, bs) + } else { + for i := range n { + s += a.complexAt(i) * b.complexAt(i) + } + } + return Scalar{isComplex: true, c: s}, nil + } +} + +// reduceOrder walks the array keeping the smallest (greater=false) or +// largest (greater=true) element per dtype. An empty array is an error, +// and so is a complex array: it has no ordering. The comparisons are +// inlined per dtype: a closure per element would dominate the walk. +// reduceOrder walks the array keeping the smallest (greater=false) or +// largest (greater=true) element per dtype. An empty array is an error, +// and so is a complex array: it has no ordering. +// +// The walk is partitioned by length alone, never by the worker count, so +// the same input answers the same element on any machine: the range is +// cut into fixed chunks, each chunk applies the serial rule to its own +// range, and the chunks combine in index order with the same strict +// comparison. A tie therefore keeps the earlier chunk's element and, in a +// chunk, the earlier element's, which is what the serial walk kept; the +// rules for a NaN (never a candidate) and an all-NaN array (the last +// element, where the seed walk ends) are preserved exactly, zeros +// included. +func (a *Array) reduceOrder(name string, greater bool) (Scalar, error) { + if a.dt == Complex { + return Scalar{}, errf("%s: complex arrays have no ordering", name) + } + if a.Len() == 0 { + return Scalar{}, errf("%s: an empty array has no %s", name, name) + } + n := a.Len() + switch a.dt { + case Int: + return Scalar{i: foldExtreme(n, greater, func(lo, hi int) (int64, bool) { + is := a.ints[lo:hi] + best := is[0] + if greater { + for _, v := range is[1:] { + if v > best { + best = v + } + } + } else { + for _, v := range is[1:] { + if v < best { + best = v + } + } + } + return best, true + })}, nil + case Bool: + return Scalar{i: foldExtremeBool(a.bools, n, greater)}, nil + case Int8: + return Scalar{i: foldExtremeNarrow(a.i8s, n, greater)}, nil + case Uint8: + return Scalar{i: foldExtremeNarrow(a.u8s, n, greater)}, nil + case Int16: + return Scalar{i: foldExtremeNarrow(a.i16s, n, greater)}, nil + case Uint16: + return Scalar{i: foldExtremeNarrow(a.u16s, n, greater)}, nil + case Int32: + return Scalar{i: foldExtremeNarrow(a.i32s, n, greater)}, nil + case Uint32: + return Scalar{i: foldExtremeNarrow(a.u32s, n, greater)}, nil + } + // The float folds read the payload slices directly: float16 and + // float32 widen exactly, so the raw walk sees the same values the + // accessor would hand over, in the same order. + var f float64 + switch a.dt { + case Float16: + f = foldExtreme(n, greater, func(lo, hi int) (float64, bool) { + return halfExtreme(a.halves[lo:hi], greater) + }) + case Float32: + f = foldExtreme(n, greater, func(lo, hi int) (float64, bool) { + return float32Extreme(a.floats32[lo:hi], greater) + }) + default: + f = foldExtreme(n, greater, func(lo, hi int) (float64, bool) { + return floatExtreme(a.floats[lo:hi], greater) + }) + } + return Scalar{isFloat: true, f: f}, nil +} + +// floatExtreme is the serial fold's rule applied to one chunk of a float64 +// payload: skip the leading NaNs to seed, then keep the first strictly +// better value. A chunk holding no non-NaN reports its last element, where +// the serial seed walk would have stopped, and says there is no candidate. +func floatExtreme(src []float64, greater bool) (float64, bool) { + best, k := src[0], 1 + for math.IsNaN(best) && k < len(src) { + best = src[k] + k++ + } + if math.IsNaN(best) { + return src[len(src)-1], false + } + if greater { + for _, v := range src[k:] { + if v > best { + best = v + } + } + } else { + for _, v := range src[k:] { + if v < best { + best = v + } + } + } + return best, true +} + +// float32Extreme is floatExtreme over a float32 payload; every widening is +// exact, so the comparisons see the values an accessor read would. +func float32Extreme(src []float32, greater bool) (float64, bool) { + best, k := float64(src[0]), 1 + for math.IsNaN(best) && k < len(src) { + best = float64(src[k]) + k++ + } + if math.IsNaN(best) { + return float64(src[len(src)-1]), false + } + if greater { + for _, v := range src[k:] { + if w := float64(v); w > best { + best = w + } + } + } else { + for _, v := range src[k:] { + if w := float64(v); w < best { + best = w + } + } + } + return best, true +} + +// halfExtreme is floatExtreme over a half-precision payload. +func halfExtreme(src []uint16, greater bool) (float64, bool) { + best, k := HalfToFloat64(src[0]), 1 + for math.IsNaN(best) && k < len(src) { + best = HalfToFloat64(src[k]) + k++ + } + if math.IsNaN(best) { + return HalfToFloat64(src[len(src)-1]), false + } + if greater { + for _, v := range src[k:] { + if w := HalfToFloat64(v); w > best { + best = w + } + } + } else { + for _, v := range src[k:] { + if w := HalfToFloat64(v); w < best { + best = w + } + } + } + return best, true +} + +// foldExtreme runs a chunk fold over the fixed partition of n elements +// and combines the partials with the serial walk's rule, so the answer +// is the serial walk's element whatever the worker count. A chunk with +// no candidate contributes nothing; an array with no candidate at all +// is all NaN and answers the last element, which is what the final +// chunk stored. +func foldExtreme[N int64 | float64](n int, greater bool, chunk func(lo, hi int) (N, bool)) N { + parts := foldParts(n) + vals := make([]N, parts) + oks := make([]bool, parts) + if parts == 1 { + vals[0], oks[0] = chunk(0, n) + } else { + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + vals[c], oks[c] = chunk(c*n/parts, (c+1)*n/parts) + } + }) + } + return combineExtreme(vals, oks, greater) +} + +// combineExtreme is foldExtreme's combination: the partials compare in +// index order with the strict rule the serial walk used, so a tie keeps +// the earlier block's value; a block with no candidate contributes +// nothing; a whole with no candidate answers the last block's stored +// fallback. +func combineExtreme[N int64 | float64](vals []N, oks []bool, greater bool) N { + best, have := vals[0], false + for c := range len(vals) { + if !oks[c] { + continue + } + if !have || (greater && vals[c] > best) || (!greater && vals[c] < best) { + best, have = vals[c], true + } + } + if have { + return best + } + return vals[len(vals)-1] +} + +// foldExtremeNarrow applies reduceOrder's chunk rule to a narrow integer +// payload: the comparison runs in the payload's own type, so no value +// ever meets a float64 rounding, and the winner widens exactly into the +// int64 the fold combines. +func foldExtremeNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int, greater bool) int64 { + return foldExtreme(n, greater, func(lo, hi int) (int64, bool) { + is := src[lo:hi] + best := is[0] + if greater { + for _, v := range is[1:] { + if v > best { + best = v + } + } + } else { + for _, v := range is[1:] { + if v < best { + best = v + } + } + } + return int64(best), true + }) +} + +// foldExtremeBool is foldExtremeNarrow for a bool payload: false below +// true, widened to the 0/1 the Int scalar carries. Go orders no bool +// with < or >, so the strict improvement writes its own logic. +func foldExtremeBool(src []bool, n int, greater bool) int64 { + return foldExtreme(n, greater, func(lo, hi int) (int64, bool) { + is := src[lo:hi] + best := is[0] + if greater { + for _, v := range is[1:] { + if v && !best { + best = v + } + } + } else { + for _, v := range is[1:] { + if !v && best { + best = v + } + } + } + if best { + return 1, true + } + return 0, true + }) +} + +// foldPartsOver runs a per-chunk fold over the fixed partition of n +// elements and combines the results through the balanced partial tree: +// the only ordering step, and its shape depends on the chunk count +// alone. The chunk boundaries c·n/parts are a function of the length and +// the part count, so the total is the same on any machine and under any +// worker setting. +func foldPartsOver[N int64 | float64 | complex128](n, parts int, fold func(c int) N) N { + if parts == 1 { + return fold(0) + } + partials := make([]N, parts) + engine.Parallel(parts, func(cs, ce int) { + for c := cs; c < ce; c++ { + partials[c] = fold(c) + } + }) + return treeSum(partials) +} + +// intFoldDot is the integer dot product: wrapping addition is +// associative, so the partition cannot change the value. +func intFoldDot(x, y []int64) int64 { + n := len(x) + parts := foldParts(n) + return foldPartsOver(n, parts, func(c int) int64 { + lo, hi := c*n/parts, (c+1)*n/parts + var s int64 + for i := lo; i < hi; i++ { + s += x[i] * y[i] + } + return s + }) +} + +// narrowFoldDot is the same-dtype narrow integer dot product: the +// products run in int64 over the exact widenings, exactly what the +// accessor fold's intAt reads produce, so the partition and the value +// are the accessor fold's own. +func narrowFoldDot[T int8 | uint8 | int16 | uint16 | int32 | uint32](x, y []T) int64 { + n := len(x) + parts := foldParts(n) + return foldPartsOver(n, parts, func(c int) int64 { + lo, hi := c*n/parts, (c+1)*n/parts + var s int64 + for i := lo; i < hi; i++ { + s += int64(x[i]) * int64(y[i]) + } + return s + }) +} + +// foldDotF64 is the float64 dot product: four interleaved product chains +// per chunk, combined as a balanced pair. +func foldDotF64(x, y []float64) float64 { + n := len(x) + parts := foldParts(n) + return foldPartsOver(n, parts, func(c int) float64 { + lo, hi := c*n/parts, (c+1)*n/parts + return foldDotRange(x[lo:hi], y[lo:hi]) + }) +} + +// foldDotF32 is the float32 dot product: every product is exact in +// float64, so the accumulation sees the values an accessor read would. +func foldDotF32(x, y []float32) float64 { + n := len(x) + parts := foldParts(n) + return foldPartsOver(n, parts, func(c int) float64 { + lo, hi := c*n/parts, (c+1)*n/parts + return foldDotRangeF32(x[lo:hi], y[lo:hi]) + }) +} + +// foldDotF16 is the half-precision dot product, widening each element +// exactly as it is read. +func foldDotF16(x, y []uint16) float64 { + n := len(x) + parts := foldParts(n) + return foldPartsOver(n, parts, func(c int) float64 { + lo, hi := c*n/parts, (c+1)*n/parts + return foldDotRangeF16(x[lo:hi], y[lo:hi]) + }) +} + +// foldDotC128 is the complex dot product with the same partition. +func foldDotC128(x, y []complex128) complex128 { + n := len(x) + parts := foldParts(n) + return foldPartsOver(n, parts, func(c int) complex128 { + lo, hi := c*n/parts, (c+1)*n/parts + return foldDotRangeC128(x[lo:hi], y[lo:hi]) + }) +} diff --git a/internal/core/reduce_accuracy_test.go b/internal/core/reduce_accuracy_test.go new file mode 100644 index 0000000..7c0298e --- /dev/null +++ b/internal/core/reduce_accuracy_test.go @@ -0,0 +1,357 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The sum and the dot product are partitioned by length alone, so their +// result must not depend on how many workers the engine runs, and the +// interleaved partials must round no worse than the single chain they +// replaced. Both claims are pinned here against a machine-independent +// reference: big.Float at 200 bits. + +func foldFixture(n int) []float64 { + s := uint64(20260921) + v := make([]float64, n) + for i := range v { + s = s*6364136223846793005 + 1442695040888963407 + // Alternating signs and a wide magnitude spread, so cancellation + // is the dominant error source rather than a rounding curiosity. + mag := math.Pow(10, float64(int((s>>50)%13)-6)) + v[i] = mag * float64(int64((s>>20)%2001)-1000) / 1000 + } + return v +} + +func bigSum(v []float64) *big.Float { + acc := new(big.Float).SetPrec(200) + for _, x := range v { + acc.Add(acc, new(big.Float).SetPrec(200).SetFloat64(x)) + } + return acc +} + +// chainSum is the shape the fold replaced: one accumulator, one chain. +func chainSum(v []float64) float64 { + var s float64 + for _, x := range v { + s += x + } + return s +} + +func chainDot(x, y []float64) float64 { + var s float64 + for i := range x { + s += x[i] * y[i] + } + return s +} + +func relErr(got float64, want *big.Float) float64 { + w, _ := want.Float64() + if w == 0 { + return math.Abs(got) + } + return math.Abs(got-w) / math.Abs(w) +} + +func TestSumIndependentOfWorkerCount(t *testing.T) { + v := foldFixture(1 << 20) + a, err := FromFloats(v, 1<<20) + if err != nil { + t.Fatal(err) + } + prev := engine.SetNumWorkers(1) + defer engine.SetNumWorkers(prev) + one := Sum(a).Float() + for _, w := range []int{2, 4, 8, 32} { + engine.SetNumWorkers(w) + if got := Sum(a).Float(); got != one { + t.Fatalf("Sum with %d workers = %v, with 1 worker = %v: the fold must not depend on the worker count", w, got, one) + } + } + engine.SetNumWorkers(prev) + // The same for the dot product and the mean. + b, err := FromFloats(v, 1<<20) + if err != nil { + t.Fatal(err) + } + engine.SetNumWorkers(1) + dOne, err := Dot(a, b) + if err != nil { + t.Fatal(err) + } + mOne, err := Mean(a) + if err != nil { + t.Fatal(err) + } + for _, w := range []int{3, 7, 16} { + engine.SetNumWorkers(w) + d, err := Dot(a, b) + if err != nil { + t.Fatal(err) + } + m, err := Mean(a) + if err != nil { + t.Fatal(err) + } + if d.Float() != dOne.Float() || m != mOne { + t.Fatalf("worker count %d changed the result: dot %v vs %v, mean %v vs %v", w, d.Float(), dOne.Float(), m, mOne) + } + } +} + +func TestSumDotAccuracyAgainstBigFloat(t *testing.T) { + for _, n := range []int{1 << 12, 1 << 16, 1 << 20} { + v := foldFixture(n) + a, err := FromFloats(v, n) + if err != nil { + t.Fatal(err) + } + want := bigSum(v) + got := Sum(a).Float() + old := chainSum(v) + newErr, oldErr := relErr(got, want), relErr(old, want) + t.Logf("n=%d: fold %.3e, one chain %.3e", n, newErr, oldErr) + // The two forms round differently, not always in the same + // direction: the interleaved partials shorten each dependency + // chain, the combination adds three roundings, and on some + // lengths the chain wins by luck. What is pinned is the order: + // the fold stays within a small factor of the chain and beats it + // where the chain is long enough for its own roundings to + // accumulate. + if newErr > 4*oldErr && newErr > 1e-16 { + t.Errorf("n=%d: the fold rounds an order worse than the chain: %.3e against %.3e", n, newErr, oldErr) + } + // The dot product of the fixture with itself: the same claim. + d, err := Dot(a, a) + if err != nil { + t.Fatal(err) + } + acc := new(big.Float).SetPrec(200) + for _, x := range v { + bx := new(big.Float).SetPrec(200).SetFloat64(x) + acc.Add(acc, new(big.Float).SetPrec(200).Mul(bx, bx)) + } + dOld := chainDot(v, v) + if dErr, dOldErr := relErr(d.Float(), acc), relErr(dOld, acc); dErr > 4*dOldErr { + t.Errorf("n=%d: the dot fold rounds an order worse than the chain: %.3e against %.3e", n, dErr, dOldErr) + } + } +} + +// TestFoldComposesAcrossBlockCuts pins the property the distributed +// reductions stand on: partials folded per block of the canonical +// partition combine, through the same tree the whole fold uses, to the +// whole fold's exact bits, whatever contiguous cuts of the block range +// produced them. +func TestFoldComposesAcrossBlockCuts(t *testing.T) { + for _, n := range []int{foldChunk + 1, 3*foldChunk + 17, 17 * foldChunk} { + v := foldFixture(n) + a, err := FromFloats(v, n) + if err != nil { + t.Fatal(err) + } + whole := Sum(a).Float() + parts := foldParts(n) + cuts := [][]int{{0, parts}, {0, 1, parts}, {0, parts / 3, (2 * parts) / 3, parts}} + if parts >= 13 { + cuts = append(cuts, []int{0, 2, 3, 5, 8, 13, parts}) + } + for _, cut := range cuts { + partials := make([]float64, 0, parts) + for j := 0; j+1 < len(cut); j++ { + lo, hi := cut[j], cut[j+1] + for c := lo; c < hi; c++ { + partials = append(partials, foldRange(v[c*n/parts:(c+1)*n/parts])) + } + } + if got := treeSum(partials); math.Float64bits(got) != math.Float64bits(whole) { + t.Fatalf("n=%d cut %v: sharded fold %v (%b) against whole fold %v (%b)", + n, cut, got, math.Float64bits(got), whole, math.Float64bits(whole)) + } + } + } +} + +// serialExtreme is the walk the partitioned fold replaced, kept here as +// the reference: seed past the leading NaNs, then keep the first strictly +// better element. +func serialExtreme(src []float64, greater bool) float64 { + best, k := src[0], 1 + for math.IsNaN(best) && k < len(src) { + best = src[k] + k++ + } + for _, v := range src[k:] { + if (greater && v > best) || (!greater && v < best) { + best = v + } + } + return best +} + +func TestExtremesMatchSerialWalk(t *testing.T) { + nan := math.NaN() + shapes := [][]float64{ + {1, 2, 3}, + {nan, nan, 5, 1}, + {5, nan, nan}, + {nan, nan, nan}, + {0, math.Copysign(0, -1)}, + {math.Copysign(0, -1), 0}, + {0, math.Copysign(0, -1), 0}, + {math.Inf(1), math.Inf(-1), 1}, + {-1, -1, -1}, + } + // A large array with NaNs and zeros pinned to the chunk boundaries. + big := make([]float64, 3*foldChunk+17) + for i := range big { + big[i] = float64((i*37)%101) - 50 + } + for _, idx := range []int{0, foldChunk - 1, foldChunk, 2 * foldChunk, len(big) - 1} { + big[idx] = nan + } + big[foldChunk+3] = 0 + big[foldChunk+4] = math.Copysign(0, -1) + shapes = append(shapes, big) + + prev := engine.SetNumWorkers(1) + defer engine.SetNumWorkers(prev) + for si, vals := range shapes { + a, err := FromFloats(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + for _, w := range []int{1, 2, 3, 8, 32} { + engine.SetNumWorkers(w) + for _, greater := range []bool{true, false} { + var got Scalar + var err error + if greater { + got, err = Max(a) + } else { + got, err = Min(a) + } + if err != nil { + t.Fatalf("shape %d: %v", si, err) + } + want := serialExtreme(vals, greater) + if math.Float64bits(got.Float()) != math.Float64bits(want) { + t.Fatalf("shape %d workers %d greater=%v: %v (%b), serial %v (%b)", + si, w, greater, got.Float(), math.Float64bits(got.Float()), want, math.Float64bits(want)) + } + } + } + } + // The integer path, which has no NaN or zero subtlety but must still + // be partition-independent. + ivals := make([]int64, foldChunk+5) + for i := range ivals { + ivals[i] = int64((i*13)%97) - 48 + } + ia, err := FromInts(ivals, len(ivals)) + if err != nil { + t.Fatal(err) + } + engine.SetNumWorkers(1) + imin, _ := Min(ia) + imax, _ := Max(ia) + for _, w := range []int{2, 5, 16} { + engine.SetNumWorkers(w) + lo, _ := Min(ia) + hi, _ := Max(ia) + if lo.Int() != imin.Int() || hi.Int() != imax.Int() { + t.Fatalf("integer extremes at %d workers: %v/%v against %v/%v", w, lo.Int(), hi.Int(), imin.Int(), imax.Int()) + } + } +} + +// prodFixture keeps every factor a relative hair away from one, so a +// million-fold product stays finite and the error the roundings make +// is measurable against the exact referent. +func prodFixture(n int) []float64 { + s := uint64(20260922) + v := make([]float64, n) + for i := range v { + s = s*6364136223846793005 + 1442695040888963407 + v[i] = 1 + float64(int64((s>>40)%2001)-1000)/1e6 + } + return v +} + +func bigProd(v []float64) *big.Float { + acc := new(big.Float).SetPrec(200) + acc.SetFloat64(1) + for _, x := range v { + acc.Mul(acc, new(big.Float).SetPrec(200).SetFloat64(x)) + } + return acc +} + +func chainProd(v []float64) float64 { + m := 1.0 + for _, x := range v { + m *= x + } + return m +} + +func TestProdNormAccuracyAgainstBigFloat(t *testing.T) { + const n = 1 << 20 + // The product: the block tree against the single chain. + v := prodFixture(n) + a, err := FromFloats(v, n) + if err != nil { + t.Fatal(err) + } + want := bigProd(v) + got, err := Prod(a, 0, false) + if err != nil { + t.Fatal(err) + } + old := chainProd(v) + if newErr, oldErr := relErr(got.FloatAt(0), want), relErr(old, want); newErr > 4*oldErr && newErr > 1e-16 { + t.Errorf("n=%d: the product tree rounds an order worse than the chain: %.3e against %.3e", n, newErr, oldErr) + } else { + t.Logf("n=%d: product tree %.3e, one chain %.3e", n, newErr, oldErr) + } + // The two norm, whose power sum folds through the same tree: the + // squared factors keep the fixture near one, so the referent is + // meaningful. + sq := make([]float64, n) + for i := range sq { + sq[i] = v[i] * v[i] + } + acc := new(big.Float).SetPrec(200) + for _, x := range sq { + acc.Add(acc, new(big.Float).SetPrec(200).SetFloat64(x)) + } + nrm, err := Norm(a, 2, 0, false) + if err != nil { + t.Fatal(err) + } + // The norm closes through Sqrt; compare the power sums, where the + // rounding the tree moves lives. + chainSum := 0.0 + for _, x := range sq { + chainSum += x + } + // The norm answers Sqrt of its power sum; recover the sum to keep + // the comparison on the folded quantity. + gotSum := nrm.FloatAt(0) * nrm.FloatAt(0) + if newErr, oldErr := relErr(gotSum, acc), relErr(chainSum, acc); newErr > 4*oldErr && newErr > 1e-16 { + t.Errorf("n=%d: the norm's power sum rounds an order worse than the chain: %.3e against %.3e", n, newErr, oldErr) + } else { + t.Logf("n=%d: norm power sum %.3e, one chain %.3e", n, newErr, oldErr) + } + _ = want +} diff --git a/internal/core/reduce_test.go b/internal/core/reduce_test.go new file mode 100644 index 0000000..d3f4f1d --- /dev/null +++ b/internal/core/reduce_test.go @@ -0,0 +1,330 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +func TestSum(t *testing.T) { + i := mustFromInts(t, []int64{1, 2, 3}, 3) + s := Sum(i) + if s.IsFloat() || s.Int() != 6 { + t.Fatalf("Sum int: %s", s) + } + + f := mustFromFloats(t, []float64{0.5, 1.5}, 2) + fs := Sum(f) + if !fs.IsFloat() || fs.Float() != 2.0 { + t.Fatalf("Sum float: %s", fs) + } + + empty := mustFromInts(t, nil, 0) + if Sum(empty).Int() != 0 { + t.Fatalf("Sum of empty must be zero") + } +} + +func TestMinMeanMax(t *testing.T) { + a := mustFromInts(t, []int64{3, 1, 2}, 3) + mn, err := Min(a) + if err != nil || mn.Int() != 1 { + t.Fatalf("Min: %s %v", mn, err) + } + mx, err := Max(a) + if err != nil || mx.Int() != 3 { + t.Fatalf("Max: %s %v", mx, err) + } + + f := mustFromFloats(t, []float64{2.5, -1.5}, 2) + fmn, _ := Min(f) + if !fmn.IsFloat() || fmn.Float() != -1.5 { + t.Fatalf("Min float: %s", fmn) + } + + mean, err := Mean(a) + if err != nil || mean != 2.0 { + t.Fatalf("Mean int: %v %v", mean, err) + } + fm, err := Mean(f) + if err != nil || fm != 0.5 { + t.Fatalf("Mean float: %v %v", fm, err) + } + + empty := mustFromInts(t, nil, 0) + if _, err := Min(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("Min empty: %v", err) + } + if _, err := Max(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("Max empty: %v", err) + } + if _, err := Mean(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("Mean empty: %v", err) + } +} + +// TestProdEmptyAxisIsUnitProduct pins the empty product along a reduced +// axis: a line of zero elements multiplies to the multiplicative +// identity in the array's own dtype, the same value the global +// one-dimensional product answers for an empty array. The zeroed +// allocation must never surface as the answer. +func TestProdEmptyAxisIsUnitProduct(t *testing.T) { + f := mustFromFloats(t, nil, 2, 0) + got, err := Prod(f, 1, false) + if err != nil { + t.Fatalf("Prod float over an empty axis: %v", err) + } + if got.Dtype() != Float || got.Len() != 2 { + t.Fatalf("Prod float over an empty axis: dtype %s shape %v", got.Dtype(), got.Shape()) + } + for i := range 2 { + if v := got.RawFloats()[i]; v != 1 { + t.Errorf("Prod float over an empty axis [%d] = %v, want 1", i, v) + } + } + + i := mustFromInts(t, nil, 2, 0) + gotI, err := Prod(i, 1, false) + if err != nil { + t.Fatalf("Prod int over an empty axis: %v", err) + } + if gotI.Dtype() != Int { + t.Fatalf("Prod int over an empty axis: dtype %s", gotI.Dtype()) + } + for k := range 2 { + if v := gotI.RawInts()[k]; v != 1 { + t.Errorf("Prod int over an empty axis [%d] = %v, want 1", k, v) + } + } + + h := mustFromFloat16s(t, nil, 2, 0) + gotH, err := Prod(h, 1, false) + if err != nil { + t.Fatalf("Prod float16 over an empty axis: %v", err) + } + for k := range 2 { + if bits := gotH.RawHalves()[k]; bits != halfOne { + t.Errorf("Prod float16 over an empty axis [%d] = %#04x, want the half one %#04x", k, bits, halfOne) + } + } + + // keepDim keeps the unit answer under the reinserted size-1 axis. + kept, err := Prod(f, 1, true) + if err != nil { + t.Fatalf("Prod keepDim over an empty axis: %v", err) + } + if kept.Shape()[0] != 2 || kept.Shape()[1] != 1 || kept.RawFloats()[1] != 1 { + t.Fatalf("Prod keepDim over an empty axis: shape %v values %v", kept.Shape(), kept.RawFloats()) + } + + // The global one-dimensional empty product agrees: both routes to a + // product of zero factors answer the unit. + global, err := Prod(mustFromFloats(t, nil, 0), 0, false) + if err != nil { + t.Fatalf("Prod of the empty array: %v", err) + } + if v := global.RawFloats()[0]; v != 1 { + t.Errorf("Prod of the empty array = %v, want 1", v) + } +} + +func TestScalarBox(t *testing.T) { + i := Scalar{i: 7} + if i.IsFloat() || i.Int() != 7 || i.Float() != 7.0 || i.String() != "int 7" { + t.Fatalf("int scalar: %s", i) + } + f := Scalar{isFloat: true, f: 2.5} + if !f.IsFloat() || f.Float() != 2.5 || f.Int() != 2 || f.String() != "float 2.5" { + t.Fatalf("float scalar: %s", f) + } +} + +func TestDot(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3}, 3) + b := mustFromInts(t, []int64{4, 5, 6}, 3) + d, err := Dot(a, b) + if err != nil { + t.Fatalf("Dot: %v", err) + } + if d.IsFloat() || d.Int() != 32 { + t.Fatalf("Dot int: %s", d) + } + + f := mustFromFloats(t, []float64{0.5, 0.5}, 2) + fd, err := Dot(f, f) + if err != nil || !fd.IsFloat() || fd.Float() != 0.5 { + t.Fatalf("Dot float: %s %v", fd, err) + } + + // Mixed dtypes promote to float. + md, err := Dot(a, mustFromFloats(t, []float64{1, 1, 1}, 3)) + if err != nil || !md.IsFloat() || md.Float() != 6.0 { + t.Fatalf("Dot mixed: %s %v", md, err) + } + + if _, err := Dot(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "length mismatch") { + t.Fatalf("Dot length: %v", err) + } + m2 := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + if _, err := Dot(m2, m2); err == nil || !strings.Contains(err.Error(), "needs 1-D arrays") { + t.Fatalf("Dot 2-D: %v", err) + } +} + +// TestNormIntArms pins the magnitude walks over an int payload, whose +// values are taken in int64 and widened exactly: p = 1 adds the +// magnitudes, p = 2 squares them, and p = Inf takes the largest. +func TestNormIntArms(t *testing.T) { + a := mustFromInts(t, []int64{3, -4}, 2) + for _, tc := range []struct { + p float64 + want float64 + }{ + {1, 7}, // |3| + |-4| + {2, 5}, // √(3² + 4²), exact + {math.Inf(1), 4}, // the largest magnitude + } { + got, err := Norm(a, tc.p, 0, false) + if err != nil { + t.Fatalf("Norm int p=%v: %v", tc.p, err) + } + if v := got.RawFloats()[0]; v != tc.want { + t.Errorf("Norm int p=%v over [3 -4]: %v, want %v", tc.p, v, tc.want) + } + } + // A second sample keeps the sum of squares off one Pythagorean + // triple: 2² + 3² + 6² = 49. + b := mustFromInts(t, []int64{2, -3, 6}, 3) + l2, err := Norm(b, 2, 0, false) + if err != nil { + t.Fatalf("Norm int p=2: %v", err) + } + if v := l2.RawFloats()[0]; v != 7 { + t.Errorf("Norm int p=2 over [2 -3 6]: %v, want 7", v) + } +} + +// TestNormFloat32Arms pins the same walks over a float32 payload: every +// value widens to float64, and a negative component must contribute its +// magnitude rather than its sign. +func TestNormFloat32Arms(t *testing.T) { + a := mustFromFloat32s(t, []float32{-3, 4}, 2) + for _, tc := range []struct { + p float64 + want float64 + }{ + {1, 7}, + {2, 5}, + {math.Inf(1), 4}, + } { + got, err := Norm(a, tc.p, 0, false) + if err != nil { + t.Fatalf("Norm float32 p=%v: %v", tc.p, err) + } + if v := got.RawFloats()[0]; v != tc.want { + t.Errorf("Norm float32 p=%v over [-3 4]: %v, want %v", tc.p, v, tc.want) + } + } + b := mustFromFloat32s(t, []float32{-1.5, 2.5, -4}, 3) + l1, err := Norm(b, 1, 0, false) + if err != nil { + t.Fatalf("Norm float32 p=1: %v", err) + } + if v := l1.RawFloats()[0]; v != 8 { + t.Errorf("Norm float32 p=1 over [-1.5 2.5 -4]: %v, want 8", v) + } +} + +// TestMeanOverIntSample pins the mean of an int sample through the +// magnitude pre-pass. The scaled branch cannot engage for an int +// payload: maxAbs is at most 2^63 while the guard asks for a magnitude +// above MaxFloat64/n, which no element count reaches, so the values +// below are the plain sum divided by the count. +func TestMeanOverIntSample(t *testing.T) { + a := mustFromInts(t, []int64{-8, -4, 4, 8}, 4) + got, err := Mean(a) + if err != nil { + t.Fatalf("Mean int: %v", err) + } + if got != 0 { + t.Errorf("Mean int over [-8 -4 4 8]: %v, want 0", got) + } + // Both ends of the int64 range: the magnitudes are taken as float64 + // (exact), the sum wraps to -1 as the int64 fold does, and the mean + // is exact in float64. + b := mustFromInts(t, []int64{math.MinInt64, math.MaxInt64}, 2) + got, err = Mean(b) + if err != nil { + t.Fatalf("Mean int range: %v", err) + } + if got != -0.5 { + t.Errorf("Mean int over [MinInt64 MaxInt64]: %v, want -0.5", got) + } +} + +// TestNarrowReductionsContract pins the integer-class reduction rules: +// a bool sum counts its true elements, narrow sums widen exactly into an +// Int scalar, Min/Max compare natively in each payload type and answer +// Int scalars, Mean runs the float64 contract over every non-complex +// dtype, and Dot refuses the bool pair by name. +func TestNarrowReductionsContract(t *testing.T) { + bl := narrowBools(t, []bool{true, false, true, true}, 4) + u8 := narrowUint8s(t, []uint8{200, 200, 200}, 3) + i8 := narrowInt8s(t, []int8{-100, 50}, 2) + + s := Sum(bl) + if s.IsFloat() || s.IsComplex() || s.Int() != 3 { + t.Fatalf("Sum bool = %v, want the Int scalar 3", s) + } + if s := Sum(u8); s.Int() != 600 { + t.Fatalf("Sum uint8 = %v, want 600", s) + } + if s := Sum(i8); s.Int() != -50 { + t.Fatalf("Sum int8 = %v, want -50", s) + } + + mn, err := Min(bl) + if err != nil || mn.Int() != 0 { + t.Fatalf("Min bool = %v %v, want the Int scalar 0", mn, err) + } + mx, err := Max(bl) + if err != nil || mx.Int() != 1 { + t.Fatalf("Max bool = %v %v, want the Int scalar 1", mx, err) + } + if mn, err = Min(u8); err != nil || mn.Int() != 200 { + t.Fatalf("Min uint8 = %v %v, want 200", mn, err) + } + if mx, err = Max(i8); err != nil || mx.Int() != 50 { + t.Fatalf("Max int8 = %v %v, want 50", mx, err) + } + + // Native comparison per payload type: the top of the uint32 range + // keeps its exact int64 image. + u32 := narrowUint32s(t, []uint32{math.MaxUint32, math.MaxUint32 - 1}, 2) + if mx, err = Max(u32); err != nil || mx.Int() != int64(math.MaxUint32) { + t.Fatalf("Max uint32 = %v %v, want %d", mx, err, uint64(math.MaxUint32)) + } + + // Mean runs the float64 contract over every non-complex dtype. + m, err := Mean(i8) + if err != nil || m != -25 { + t.Fatalf("Mean int8 = %v %v, want -25", m, err) + } + if m, err = Mean(bl); err != nil || m != 0.75 { + t.Fatalf("Mean bool = %v %v, want 0.75", m, err) + } + + // Dot: an integer-class pair answers an Int scalar accumulated in + // int64; a bool pair is the arithmetic refusal. + d, err := Dot(i8, i8) + if err != nil || d.IsFloat() || d.Int() != 12500 { + t.Fatalf("Dot int8 = %v %v, want the Int scalar 12500", d, err) + } + if _, err := Dot(bl, bl); err == nil || + !strings.Contains(err.Error(), "bool arrays have no arithmetic") { + t.Fatalf("Dot bool: %v", err) + } +} diff --git a/internal/core/reduction2.go b/internal/core/reduction2.go new file mode 100644 index 0000000..1dd664b --- /dev/null +++ b/internal/core/reduction2.go @@ -0,0 +1,879 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "math" + +// Extended reductions: prefix scans (CumSum, CumProd) along one +// dimension, product along one dimension (Prod), and the Lp norm along +// one dimension (Norm). The cumulative scans return the same dtype as +// the input; int scans wrap on overflow (consistent with the rest of +// the library). Norm always returns float64 (the result +// even for int inputs: Lp norms need a real-valued magnitude). + +// CumSum returns the cumulative sum along dim; the result has the same +// shape as a. +func CumSum(a *Array, dim int) (*Array, error) { + return a.scanDim(dim, "CumSum", false) +} + +// CumProd returns the cumulative product along dim. +func CumProd(a *Array, dim int) (*Array, error) { + return a.scanDim(dim, "CumProd", true) +} + +// Prod returns the product along dim; keepDim preserves the reduced +// dimension as size 1. +func Prod(a *Array, dim int, keepDim bool) (*Array, error) { + // The one-dimensional global product is the shape the sharded + // reductions mirror: it folds the canonical blocks where they lie + // and combines them through the same balanced tree, so a sharded + // product and this product are one computation. Lines of higher + // ranks keep the sequential walk below. + if a.NDim() == 1 && dim == 0 && a.isContiguous() { + out, err := prodGlobal1D(a) + if err != nil { + return nil, err + } + if keepDim { + return keepReducedDim(out, a.shape, dim), nil + } + return out, nil + } + out, err := a.reduceDimProd(dim, "Prod") + if err != nil { + return nil, err + } + if keepDim { + return keepReducedDim(out, a.shape, dim), nil + } + return out, nil +} + +// prodGlobal1D folds the whole one-dimensional array through the +// canonical partition: each block multiplies into one partial and the +// partials combine through the balanced tree, the exact shape the spmd +// shards reproduce. The integer product is exact under any grouping; +// the floating products round differently from the single chain at +// lengths past one block, measured against the exact referent in the +// accuracy test's bound. +func prodGlobal1D(a *Array) (*Array, error) { + if narrowRefused(a.dt) { + return nil, errf("Prod: dtype %s is not supported; convert with Astype", a.dt) + } + if a.dt == Complex { + return nil, errf("Prod: complex arrays have no real-valued product") + } + n := a.Len() + out := &Array{shape: []int{1}, dt: a.dt} + out.alloc(1) + parts := foldParts(n) + switch a.dt { + case Int: + ints := make([]int64, parts) + for c := range parts { + ints[c] = foldProdRangeI64(a.ints[c*n/parts : (c+1)*n/parts]) + } + out.ints[0] = treeProd(ints) + case Float16: + // Each block value is an exact half carried in float64, and the + // tree combines narrow through half, the per-step rounding the + // line product keeps. + vals := make([]float64, parts) + for c := range parts { + vals[c] = foldProdRangeF16(a.halves[c*n/parts : (c+1)*n/parts]) + } + out.halves[0] = HalfFromFloat64(treeProdHalf(vals)) + case Float32: + // The partials multiply natively in float32, the way the line + // product keeps. + vals := make([]float32, parts) + for c := range parts { + vals[c] = foldProdRangeF32(a.floats32[c*n/parts : (c+1)*n/parts]) + } + out.floats32[0] = treeProd(vals) + case Float: + vals := make([]float64, parts) + for c := range parts { + vals[c] = foldProdRange(a.floats[c*n/parts : (c+1)*n/parts]) + } + out.floats[0] = treeProd(vals) + default: + return nil, errf("Prod: dtype %s is not supported; convert with Astype", a.dt) + } + return out, nil +} + +// Norm returns the Lp norm along dim: (sum |x|^p)^(1/p). p must be a +// positive number; p == math.Inf returns the max-abs, where NaN +// elements never win and a line whose every element is NaN reports +// NaN, exactly as MinAxis and MaxAxis fold their extrema. NaN is +// rejected like any other invalid p, because every comparison against +// NaN is false and it would otherwise slip through the positivity +// test. The result is always float64, dtype promoted. +// +// Lines are independent and each line writes only its own destination +// slot, so the walk splits across workers with the same ascending +// addend order per line, and the fan-out is capped so a small norm +// never pays a spawn bill. The general exponent calls math.Pow per +// element, so its floor is far lower than the plain folds'. +func Norm(a *Array, p float64, dim int, keepDim bool) (*Array, error) { + if math.IsNaN(p) || (p <= 0 && !math.IsInf(p, 1)) { + return nil, errf("Norm: p must be positive, got %v", p) + } + if a.dt == Complex { + return nil, errf("Norm: complex arrays have no real-valued norm") + } + if narrowRefused(a.dt) { + // The narrow dtypes stay refused for the Norm family. + return nil, errf("Norm: dtype %s is not supported; convert with Astype", a.dt) + } + if dim < 0 || dim >= a.NDim() { + return nil, errf("Norm: dimension %d out of range for shape %s", dim, shapeText(a.shape)) + } + // The one-dimensional global norm for a finite p is the shape the + // sharded reductions mirror: the power sums fold the canonical + // blocks where they lie and combine through the same balanced tree, + // so a sharded norm and this norm are one computation. The infinity + // norm is a maximum, not a sum, and keeps the line walk below. + if a.NDim() == 1 && dim == 0 && a.isContiguous() && !math.IsInf(p, 1) { + out, err := normGlobal1D(a, p) + if err != nil { + return nil, err + } + if keepDim { + return keepReducedDim(out, a.shape, dim), nil + } + return out, nil + } + outShape := reduceShape(a.shape, dim) + out := &Array{shape: outShape, dt: Float} + total := 1 + for _, d := range outShape { + total *= d + } + out.alloc(total) + stride := 1 + for k := dim + 1; k < a.NDim(); k++ { + stride *= a.shape[k] + } + // The walk goes line by line: a line is the run of a.shape[dim] + // elements that share every surviving coordinate. Each line is + // gathered and folded through the canonical partition the global + // norm uses: fixed blocks, the block partials combined through the + // balanced tree. The partition follows from the line length alone, + // so a single-line norm answers normGlobal1D's exact bits whatever + // the shape or the worker split. + line := a.shape[dim] + perLine := stride * line + lines := 0 + if perLine > 0 { + lines = a.Len() / perLine + } + pInf := math.IsInf(p, 1) + if pInf && line == 0 { + // An empty reduced dimension has no maximum: MinAxis and + // MaxAxis report NaN for it, and the infinity norm agrees + // rather than reporting the untouched zero. + for i := range out.floats { + out.floats[i] = math.NaN() + } + return out, nil + } + if pInf { + // The infinity norm is a maximum, not a sum: the extremum walk + // below keeps the line order, which no grouping can move. + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + switch a.dt { + case Int: + src := a.ints + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + dst := b*stride + s + acc := out.floats[dst] + for off := range line { + if v := math.Abs(float64(src[base+off*stride+s])); v > acc { + acc = v + } + } + out.floats[dst] = acc + } + } + case Float16: + src := a.halves + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + dst := b*stride + s + i, m := 0, 0.0 + for i < line { + if v := math.Abs(HalfToFloat64(src[base+i*stride+s])); v == v { + m = v + break + } + i++ + } + if i == line { + out.floats[dst] = math.NaN() + continue + } + i++ + for off := i; off < line; off++ { + if v := math.Abs(HalfToFloat64(src[base+off*stride+s])); v > m { + m = v + } + } + out.floats[dst] = m + } + } + case Float32: + src := a.floats32 + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + dst := b*stride + s + i, m := 0, 0.0 + for i < line { + if v := math.Abs(float64(src[base+i*stride+s])); v == v { + m = v + break + } + i++ + } + if i == line { + out.floats[dst] = math.NaN() + continue + } + i++ + for off := i; off < line; off++ { + if v := math.Abs(float64(src[base+off*stride+s])); v > m { + m = v + } + } + out.floats[dst] = m + } + } + default: + src := a.floats + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + dst := b*stride + s + i, m := 0, 0.0 + for i < line { + if v := math.Abs(src[base+i*stride+s]); v == v { + m = v + break + } + i++ + } + if i == line { + out.floats[dst] = math.NaN() + continue + } + i++ + for off := i; off < line; off++ { + if v := math.Abs(src[base+off*stride+s]); v > m { + m = v + } + } + out.floats[dst] = m + } + } + } + }) + if keepDim { + return keepReducedDim(out, a.shape, dim), nil + } + return out, nil + } + pOne := p == 1 + pTwo := p == 2 + // One canonical power-sum fold serves every finite exponent: the + // per-element arithmetic the mode picks is the one foldNormPower + // keeps, the blocks are the canonical ones and treeSum combines the + // partials. The common exponents multiply instead of calling Pow, + // bit-identically to the general path. The dtype and the mode are + // picked once per worker segment, so the fold reads the payload + // where the elements live with no per-element call and no gather + // scratch, and the partial table is shared by the segment's lines. + const ( + normP1 = iota + normP2 + normPGeneral + ) + mode := normPGeneral + switch { + case pOne: + mode = normP1 + case pTwo: + mode = normP2 + } + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + parts := foldParts(line) + partials := make([]float64, parts) + var foldRange func(base, s, lo, hi int) float64 + switch a.dt { + case Int: + src := a.ints + switch mode { + case normP1: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Abs(float64(src[base+off*stride+s])) + } + return acc + } + case normP2: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + v := math.Abs(float64(src[base+off*stride+s])) + acc += v * v + } + return acc + } + default: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Pow(math.Abs(float64(src[base+off*stride+s])), p) + } + return acc + } + } + case Float16: + src := a.halves + switch mode { + case normP1: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Abs(HalfToFloat64(src[base+off*stride+s])) + } + return acc + } + case normP2: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + v := math.Abs(HalfToFloat64(src[base+off*stride+s])) + acc += v * v + } + return acc + } + default: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Pow(math.Abs(HalfToFloat64(src[base+off*stride+s])), p) + } + return acc + } + } + case Float32: + src := a.floats32 + switch mode { + case normP1: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Abs(float64(src[base+off*stride+s])) + } + return acc + } + case normP2: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + v := math.Abs(float64(src[base+off*stride+s])) + acc += v * v + } + return acc + } + default: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Pow(math.Abs(float64(src[base+off*stride+s])), p) + } + return acc + } + } + default: + src := a.floats + switch mode { + case normP1: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Abs(src[base+off*stride+s]) + } + return acc + } + case normP2: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + v := math.Abs(src[base+off*stride+s]) + acc += v * v + } + return acc + } + default: + foldRange = func(base, s, lo, hi int) float64 { + var acc float64 + for off := lo; off < hi; off++ { + acc += math.Pow(math.Abs(src[base+off*stride+s]), p) + } + return acc + } + } + } + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + if parts == 1 { + out.floats[b*stride+s] = foldRange(base, s, 0, line) + continue + } + for c := range parts { + partials[c] = foldRange(base, s, c*line/parts, (c+1)*line/parts) + } + out.floats[b*stride+s] = treeSum(partials) + } + } + }) + switch { + case pTwo: + for k := range out.floats { + out.floats[k] = math.Sqrt(out.floats[k]) + } + case !pOne: + for k := range out.floats { + out.floats[k] = math.Pow(out.floats[k], 1/p) + } + } + if keepDim { + return keepReducedDim(out, a.shape, dim), nil + } + return out, nil +} + +// scanDim walks a row-major, maintaining a per-line running value for +// dim: mul selects the product (CumProd) over the sum (CumSum). The +// combine happens in registers per line slot; each output element still +// receives the same combine sequence the off-then-s walk produced, so +// results are bit-identical, with no per-element closure call. Whole +// lines go to whole workers and each line carries its own running value, +// so the split moves no bit; the fan-out is capped so a small scan stays +// on the calling goroutine. The float64 sum carry is Neumaier +// compensated, the accuracy the compensated-scan test measured against +// the exact referent; the product and every other dtype keep the plain +// chain, and the compensation runs per line in element order, so the +// split still moves no bit. +func (a *Array) scanDim(dim int, name string, mul bool) (*Array, error) { + if dim < 0 || dim >= a.NDim() { + return nil, errf("%s: dimension %d out of range for shape %s", name, dim, shapeText(a.shape)) + } + if narrowRefused(a.dt) { + // The reduction surface offers Sum, Mean, Min, Max and + // the arg extremes; the scans refuse the narrow dtypes. + return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt) + } + if a.shape[dim] == 0 { + return nil, errf("%s: dimension %d of shape %s is empty", name, dim, shapeText(a.shape)) + } + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(a.Len()) + // A zero trailing dimension leaves no payload to scan: the result + // is the empty array, and the per-line division below would + // otherwise divide zero by zero. + if a.Len() == 0 { + return out, nil + } + stride := 1 + for k := dim + 1; k < a.NDim(); k++ { + stride *= a.shape[k] + } + line := a.shape[dim] + perLine := stride * line + lines := a.Len() / perLine + switch a.dt { + case Int: + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for blk := ls * perLine; blk < le*perLine; blk += perLine { + for s := range stride { + pos := blk + s + acc := a.ints[pos] + out.ints[pos] = acc + for off := 1; off < line; off++ { + pos += stride + if mul { + acc *= a.ints[pos] + } else { + acc += a.ints[pos] + } + out.ints[pos] = acc + } + } + } + }) + case Float16: + // Compute in float64 for accuracy; each step narrows the carry + // to half before combining, exactly as the float32 path narrows + // to float32. + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for blk := ls * perLine; blk < le*perLine; blk += perLine { + for s := range stride { + pos := blk + s + acc := HalfToFloat64(a.halves[pos]) + out.halves[pos] = HalfFromFloat64(acc) + for off := 1; off < line; off++ { + pos += stride + if mul { + acc = HalfToFloat64(HalfFromFloat64(acc)) * HalfToFloat64(a.halves[pos]) + } else { + acc = HalfToFloat64(HalfFromFloat64(acc)) + HalfToFloat64(a.halves[pos]) + } + out.halves[pos] = HalfFromFloat64(acc) + } + } + } + }) + case Float32: + // Compute in float64 for accuracy, round back. Each step + // narrows the carry to float32 before combining, exactly as the + // previous element's stored value fed the next one. + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for blk := ls * perLine; blk < le*perLine; blk += perLine { + for s := range stride { + pos := blk + s + acc := float64(a.floats32[pos]) + out.floats32[pos] = float32(acc) + for off := 1; off < line; off++ { + pos += stride + if mul { + acc = float64(float32(acc)) * float64(a.floats32[pos]) + } else { + acc = float64(float32(acc)) + float64(a.floats32[pos]) + } + out.floats32[pos] = float32(acc) + } + } + } + }) + case Float: + if mul { + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for blk := ls * perLine; blk < le*perLine; blk += perLine { + for s := range stride { + pos := blk + s + acc := a.floats[pos] + out.floats[pos] = acc + for off := 1; off < line; off++ { + pos += stride + acc *= a.floats[pos] + out.floats[pos] = acc + } + } + } + }) + break + } + // The float64 sum carry is compensated: a running high part + // plus a correction, the output the corrected running value. + // The plain chain's accumulation error was measured against a + // 512-bit big.Float referent at 9.4e12 absolute on a 2^20 + // random sample and a full loss of every small addend on the + // cancellation pattern [1, 1e100, 1, -1e100], where the + // compensated walk answers exactly; on benign data it answers + // to the last ulp. A non-finite partial freezes the + // correction, so an overflow sticks to infinity and a NaN + // poisons the tail exactly as the plain chain's would. + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for blk := ls * perLine; blk < le*perLine; blk += perLine { + for s := range stride { + pos := blk + s + sum := a.floats[pos] + var comp float64 + out.floats[pos] = sum + for off := 1; off < line; off++ { + pos += stride + x := a.floats[pos] + t := sum + x + if !math.IsInf(t, 0) && t == t { + if math.Abs(sum) >= math.Abs(x) { + comp += (sum - t) + x + } else { + comp += (x - t) + sum + } + } + sum = t + out.floats[pos] = sum + comp + } + } + } + }) + default: + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + for blk := ls * perLine; blk < le*perLine; blk += perLine { + for s := range stride { + pos := blk + s + acc := a.complexes[pos] + out.complexes[pos] = acc + for off := 1; off < line; off++ { + pos += stride + if mul { + acc *= a.complexes[pos] + } else { + acc += a.complexes[pos] + } + out.complexes[pos] = acc + } + } + } + }) + } + return out, nil +} + +// reduceDimProd multiplies along dim, same pattern as the axis reductions +// but a product. Used by Prod. +func (a *Array) reduceDimProd(dim int, name string) (*Array, error) { + if dim < 0 || dim >= a.NDim() { + return nil, errf("%s: dimension %d out of range for shape %s", name, dim, shapeText(a.shape)) + } + if narrowRefused(a.dt) { + // The narrow dtypes are refused for Prod; the refusal names + // the dtype and the conversion. + return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt) + } + if a.dt == Complex { + return nil, errf("%s: complex arrays have no real-valued product", name) + } + outShape := reduceShape(a.shape, dim) + out := &Array{shape: outShape, dt: a.dt} + total := 1 + for _, d := range outShape { + total *= d + } + out.alloc(total) + stride := 1 + for k := dim + 1; k < a.NDim(); k++ { + stride *= a.shape[k] + } + // Line by line like Norm: each line is gathered into the payload's + // own scratch and multiplied through the canonical partition the + // global product uses, one block partial per block and the balanced + // tree over them. The partition follows from the line length alone, + // so a single-line product answers prodGlobal1D's exact bits + // whatever the shape, the stride or the worker split, and the + // payload's own arithmetic (the half and float32 narrowings) is + // preserved. + line := a.shape[dim] + perLine := stride * line + lines := 0 + if perLine > 0 { + lines = a.Len() / perLine + } + if line == 0 { + // An empty reduced dimension is the empty product: every line's + // answer is the multiplicative identity in the array's own dtype, + // exactly the value prodGlobal1D answers for an empty + // one-dimensional array. The zeroed allocation must not surface. + for i := range total { + switch a.dt { + case Int: + out.ints[i] = 1 + case Float16: + out.halves[i] = halfOne + case Float32: + out.floats32[i] = 1 + default: + out.floats[i] = 1 + } + } + return out, nil + } + switch a.dt { + case Int: + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + var scratch []int64 + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + if stride == 1 { + // A contiguous line is a slice of the payload: + // the fold reads it where it lies, no gather. + out.ints[b*stride+s] = prodLine(a.ints[base+s:base+s+line], foldProdRangeI64) + continue + } + if scratch == nil { + scratch = make([]int64, line) + } + for off := range line { + scratch[off] = a.ints[base+off*stride+s] + } + out.ints[b*stride+s] = prodLine(scratch, foldProdRangeI64) + } + } + }) + case Float16: + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + var scratch []uint16 + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + if stride == 1 { + out.halves[b*stride+s] = HalfFromFloat64(prodLineHalf(a.halves[base+s : base+s+line])) + continue + } + if scratch == nil { + scratch = make([]uint16, line) + } + for off := range line { + scratch[off] = a.halves[base+off*stride+s] + } + out.halves[b*stride+s] = HalfFromFloat64(prodLineHalf(scratch)) + } + } + }) + case Float32: + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + var scratch []float32 + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + if stride == 1 { + out.floats32[b*stride+s] = prodLine(a.floats32[base+s:base+s+line], foldProdRangeF32) + continue + } + if scratch == nil { + scratch = make([]float32, line) + } + for off := range line { + scratch[off] = a.floats32[base+off*stride+s] + } + out.floats32[b*stride+s] = prodLine(scratch, foldProdRangeF32) + } + } + }) + default: + splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) { + var scratch []float64 + for b := ls; b < le; b++ { + base := b * perLine + for s := range stride { + if stride == 1 { + out.floats[b*stride+s] = prodLine(a.floats[base+s:base+s+line], foldProdRange) + continue + } + if scratch == nil { + scratch = make([]float64, line) + } + for off := range line { + scratch[off] = a.floats[base+off*stride+s] + } + out.floats[b*stride+s] = prodLine(scratch, foldProdRange) + } + } + }) + } + return out, nil +} + +// prodLine multiplies one gathered line through the canonical partition: +// one block partial per block, the balanced tree over them, the shape +// prodGlobal1D multiplies the whole array with. +func prodLine[T int64 | float64 | float32](line []T, block func([]T) T) T { + parts := foldParts(len(line)) + if parts == 1 { + return block(line) + } + partials := make([]T, parts) + for c := range parts { + partials[c] = block(line[c*len(line)/parts : (c+1)*len(line)/parts]) + } + return treeProd(partials) +} + +// prodLineHalf is prodLine for a half payload: the block partials carry +// exact half values in float64 and the tree narrows through half, the +// rounding foldProdRangeF16 and treeProdHalf keep. +func prodLineHalf(line []uint16) float64 { + parts := foldParts(len(line)) + if parts == 1 { + return foldProdRangeF16(line) + } + partials := make([]float64, parts) + for c := range parts { + partials[c] = foldProdRangeF16(line[c*len(line)/parts : (c+1)*len(line)/parts]) + } + return treeProdHalf(partials) +} + +// normGlobal1D folds the whole one-dimensional array's power sums +// through the canonical partition: each block sums |v|^p with the +// norm fold's own per-element arithmetic and the partials combine +// through the balanced tree, the exact shape the spmd shards +// reproduce. The root closes through the same Sqrt and Pow the line +// norm keeps. +func normGlobal1D(a *Array, p float64) (*Array, error) { + n := a.Len() + parts := foldParts(n) + sums := make([]float64, parts) + switch a.dt { + case Int: + for c := range parts { + sums[c] = foldNormPowerRangeI64(a.ints[c*n/parts:(c+1)*n/parts], p) + } + case Float16: + for c := range parts { + sums[c] = foldNormPowerRangeF16(a.halves[c*n/parts:(c+1)*n/parts], p) + } + case Float32: + for c := range parts { + sums[c] = foldNormPowerRangeF32(a.floats32[c*n/parts:(c+1)*n/parts], p) + } + default: + for c := range parts { + sums[c] = foldNormPowerRange(a.floats[c*n/parts:(c+1)*n/parts], p) + } + } + out := &Array{shape: []int{1}, dt: Float} + out.alloc(1) + out.floats[0] = normRoot(treeSum(sums), p) + return out, nil +} + +// reduceShape returns the shape with dim dropped. Used by the axis-style +// helpers so they all agree. +func reduceShape(shape []int, dim int) []int { + out := make([]int, 0, len(shape)-1) + out = append(out, shape[:dim]...) + out = append(out, shape[dim+1:]...) + if len(out) == 0 { + out = []int{1} + } + return out +} + +// keepReducedDim returns the array with the reduced dimension reinserted +// as size 1, used by keepDim=true on Prod and Norm. +func keepReducedDim(a *Array, original []int, dim int) *Array { + sh := make([]int, 0, len(original)) + sh = append(sh, original[:dim]...) + sh = append(sh, 1) + sh = append(sh, original[dim+1:]...) + // Every payload field the dtype may carry rides along: a narrow + // result must never lose its elements to a five-slice copy. + return &Array{shape: sh, dt: a.dt, + ints: a.ints, halves: a.halves, floats32: a.floats32, + floats: a.floats, complexes: a.complexes, + bools: a.bools, i8s: a.i8s, u8s: a.u8s, + i16s: a.i16s, u16s: a.u16s, i32s: a.i32s, u32s: a.u32s} +} diff --git a/internal/core/reduction_route_test.go b/internal/core/reduction_route_test.go new file mode 100644 index 0000000..b13c90f --- /dev/null +++ b/internal/core/reduction_route_test.go @@ -0,0 +1,80 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// TestEinsumReductionKeepsFloat32WithEngine pins the routing of the +// single-operand axis reduction for a float32 operand: the fold declines +// it, because the fold widens, sums in float64 and narrows once, while +// the engine narrows into the float32 slot once per addend. The operand +// carries 1, 1, 1e8 and -1e8 across the two summed axes, so the engine's +// per-addend narrowing lets 1e8 swallow the running 2 and answers 0, +// where a fold's narrow-once answer would be 2. +func TestEinsumReductionKeepsFloat32WithEngine(t *testing.T) { + a := mustFromFloat32s(t, []float32{1, 1, 1e8, -1e8}, 1, 2, 2) + out, err := Einsum("ijk->i", a) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + if out.Dtype() != Float32 { + t.Fatalf("Einsum answered dtype %s, want float32", out.Dtype()) + } + if got := out.FloatAt(0); got != 0 { + t.Fatalf("float32 reduction answered %v, want exactly 0; a narrow-once fold answers 2", got) + } +} + +// TestEinsumReductionFloat32ExactSum pins the plain finite case of the +// same float32 routing: the engine's per-addend walk answers the exact +// small-integer sum. +func TestEinsumReductionFloat32ExactSum(t *testing.T) { + a := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 1, 2, 2) + out, err := Einsum("ijk->i", a) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + if got := out.FloatAt(0); got != 10 { + t.Fatalf("float32 reduction answered %v, want 10", got) + } +} + +// TestEinsumReductionKeepsComplexWithEngine pins the routing of the +// single-operand axis reduction for a complex operand: the fold declines +// it, because the engine's per-slot product seed multiplies the first +// read by 1+0i, the identity for finite values but not for a purely +// imaginary infinity, which meets 0 times Inf inside the complex product +// and turns the real part NaN. A plain addition would carry the +// infinity through untouched with the real part 2. +func TestEinsumReductionKeepsComplexWithEngine(t *testing.T) { + a := mustFromComplexes(t, []complex128{1, 1, complex(0, math.Inf(1)), 0}, 1, 2, 2) + out, err := Einsum("ijk->i", a) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + got := out.ComplexAt(0) + if !math.IsNaN(real(got)) { + t.Fatalf("real part %v, want NaN from the engine's 1+0i product seed", real(got)) + } + if imag(got) != math.Inf(1) { + t.Fatalf("imaginary part %v, want +Inf carried through", imag(got)) + } +} + +// TestEinsumReductionComplexExactSum pins the plain finite case of the +// same complex routing: the seed is the multiplicative identity there, +// and the exact small-integer sum comes back whole. +func TestEinsumReductionComplexExactSum(t *testing.T) { + a := mustFromComplexes(t, []complex128{1 + 1i, 2 + 2i, 3 + 3i, 4 + 4i}, 1, 2, 2) + out, err := Einsum("ijk->i", a) + if err != nil { + t.Fatalf("Einsum: %v", err) + } + if got := out.ComplexAt(0); got != 10+10i { + t.Fatalf("complex reduction answered %v, want (10+10i)", got) + } +} diff --git a/internal/core/reduction_shape_consistency_test.go b/internal/core/reduction_shape_consistency_test.go new file mode 100644 index 0000000..e503f2d --- /dev/null +++ b/internal/core/reduction_shape_consistency_test.go @@ -0,0 +1,197 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + "testing" +) + +// The two-dimensional Prod and Norm fold each line through the canonical +// partition the one-dimensional global reductions use, so the shape +// alone cannot move a bit: reducing a packed single-line matrix agrees +// exactly with the global answer, reducing strided columns agrees with +// the packed columns, and against big.Float the line fold holds the +// tree's accuracy where the chain it replaced drifts. + +func TestProdShapeConsistency(t *testing.T) { + const n = 1 << 16 + v := prodFixture(n) + one, err := FromFloats(v, n) + if err != nil { + t.Fatal(err) + } + packed, err := FromFloats(v, 1, n) + if err != nil { + t.Fatal(err) + } + p1, err := Prod(one, 0, false) + if err != nil { + t.Fatal(err) + } + p2, err := Prod(packed, 1, false) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(p1.FloatAt(0)) != math.Float64bits(p2.FloatAt(0)) { + t.Fatalf("the packed single-line product %.17g disagrees with the global product %.17g", + p2.FloatAt(0), p1.FloatAt(0)) + } + // Strided columns against packed columns. + const rows = 8 + grid, err := FromFloats(v[:rows*(1<<13)], rows, 1<<13) + if err != nil { + t.Fatal(err) + } + cols, err := Prod(grid, 0, false) + if err != nil { + t.Fatal(err) + } + for c := range 4 { + column := make([]float64, rows) + for r := range rows { + column[r] = v[r*(1<<13)+c] + } + packedCol, err := FromFloats(column, 1, rows) + if err != nil { + t.Fatal(err) + } + want, err := Prod(packedCol, 1, false) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(cols.FloatAt(c)) != math.Float64bits(want.FloatAt(0)) { + t.Fatalf("strided column %d disagrees with the packed product", c) + } + } +} + +// TestNormShapeConsistency holds every exponent class across the shapes: +// a packed single-line norm is the global norm's own bits, strided +// columns are the packed columns, and the accuracy is pinned against +// big.Float where the chain the fold replaced visibly drifts. +func TestNormShapeConsistency(t *testing.T) { + const n = 1 << 16 + v := foldFixture(n) + one, err := FromFloats(v, n) + if err != nil { + t.Fatal(err) + } + packed, err := FromFloats(v, 1, n) + if err != nil { + t.Fatal(err) + } + for _, p := range []float64{1, 2, 3.5} { + g1, err := Norm(one, p, 0, false) + if err != nil { + t.Fatal(err) + } + g2, err := Norm(packed, p, 1, false) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(g1.FloatAt(0)) != math.Float64bits(g2.FloatAt(0)) { + t.Fatalf("p=%v: the packed single-line norm %.17g disagrees with the global norm %.17g", + p, g2.FloatAt(0), g1.FloatAt(0)) + } + // The power-sum accuracy against big.Float: the tree's blocks + // keep the small addends the chain's running total swallows. + var want *big.Float + switch p { + case 1: + want = bigSumAbs(v) + case 2: + sq := make([]float64, n) + for i, x := range v { + sq[i] = x * x + } + want = bigSum(sq) + default: + pw := make([]float64, n) + for i, x := range v { + pw[i] = math.Pow(math.Abs(x), p) + } + want = bigSum(pw) + } + // Recover the power sum from the norm the way the accuracy test + // of the global path does, so the comparison stays on the folded + // quantity. + gotSum := g2.FloatAt(0) + switch p { + case 2: + gotSum *= gotSum + case 1: + default: + gotSum = math.Pow(gotSum, p) + } + // The square-and-root closing roughly doubles the folded error, + // so the scaled bound admits it. + if d := relErr(gotSum, want); d > 4e-14 { + t.Fatalf("p=%v: the norm's power sum rounds %.3e from big.Float", p, d) + } + } + const rows = 8 + grid, err := FromFloats(v[:rows*(1<<13)], rows, 1<<13) + if err != nil { + t.Fatal(err) + } + cols, err := Norm(grid, 2, 0, false) + if err != nil { + t.Fatal(err) + } + for c := range 4 { + column := make([]float64, rows) + for r := range rows { + column[r] = v[r*(1<<13)+c] + } + packedCol, err := FromFloats(column, 1, rows) + if err != nil { + t.Fatal(err) + } + want, err := Norm(packedCol, 2, 1, false) + if err != nil { + t.Fatal(err) + } + if math.Float64bits(cols.FloatAt(c)) != math.Float64bits(want.FloatAt(0)) { + t.Fatalf("strided column %d disagrees with the packed norm", c) + } + } +} + +func bigSumAbs(v []float64) *big.Float { + acc := new(big.Float).SetPrec(200) + for _, x := range v { + acc.Add(acc, new(big.Float).SetPrec(200).SetFloat64(math.Abs(x))) + } + return acc +} + +// TestProdNormLineAccuracyAgainstBigFloat reports the accuracy the +// canonical line fold buys over the chain on the counterpoint shapes. +func TestProdNormLineAccuracyAgainstBigFloat(t *testing.T) { + const n = 1 << 16 + v := prodFixture(n) + packed, err := FromFloats(v, 1, n) + if err != nil { + t.Fatal(err) + } + chain := 1.0 + for _, x := range v { + chain *= x + } + got, err := Prod(packed, 1, false) + if err != nil { + t.Fatal(err) + } + acc := new(big.Float).SetPrec(200).SetFloat64(1) + for _, x := range v { + acc.Mul(acc, new(big.Float).SetPrec(200).SetFloat64(x)) + } + newErr, oldErr := relErr(got.FloatAt(0), acc), relErr(chain, acc) + t.Logf("product line n=%d: tree %.3e, one chain %.3e", n, newErr, oldErr) + if newErr > 4*oldErr && newErr > 1e-15 { + t.Fatalf("the product line rounds an order worse than the chain: %.3e against %.3e", newErr, oldErr) + } +} diff --git a/internal/core/reference_walks_test.go b/internal/core/reference_walks_test.go new file mode 100644 index 0000000..81d62fc --- /dev/null +++ b/internal/core/reference_walks_test.go @@ -0,0 +1,965 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +// Reference walks: the bisections, the odometer packing and the +// trapezoid sums the fixed-iteration searches, the count-then-fill and +// the halved products replace. They pin the exact semantics the fast +// forms must reproduce bit for bit: bin clamping with NaN to the +// outermost bin, the rightmost insertion rule, the left-segment knot +// rule, the row-major coordinate order and the exact accumulation. + +func refBinFloat(ev []float64, v float64) int { + if v < ev[0] { + return 0 + } + lo, hi := 0, len(ev)-1 + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if ev[mid] <= v { + lo = mid + 1 + } else { + hi = mid + } + } + if lo == 0 { + return len(ev) - 2 + } + return lo - 1 +} + +func refBinInt(ev []int64, v int64) int { + if v < ev[0] { + return 0 + } + lo, hi := 0, len(ev)-1 + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if ev[mid] <= v { + lo = mid + 1 + } else { + hi = mid + } + } + return lo - 1 +} + +func refUpperFloat(h []float64, q float64) int { + lo, hi := 0, len(h) + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if h[mid] <= q { + lo = mid + 1 + } else { + hi = mid + } + } + return lo +} + +func refUpperInt(h []int64, q int64) int { + lo, hi := 0, len(h) + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if h[mid] <= q { + lo = mid + 1 + } else { + hi = mid + } + } + return lo +} + +func refInterpSegment(xv []float64, q float64) int { + n := len(xv) + if q <= xv[0] { + return 0 + } + if q >= xv[n-1] { + return n - 2 + } + lo, hi := 1, n + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if xv[mid] < q { + lo = mid + 1 + } else { + hi = mid + } + } + return lo - 1 +} + +// refArgwhere collects coordinates the way the single odometer walk did. +func refArgwhere(a *Array) []int64 { + var rows []int64 + coord := make([]int, a.NDim()) + for i := range a.Len() { + if !isZero(a, i) { + for d := range a.NDim() { + rows = append(rows, int64(coord[d])) + } + } + advanceOdometer(coord, a.shape) + } + return rows +} + +// refGather is the per-element accessor walk Gather replaced. +func refGather(src *Array, dim int, index *Array) (*Array, error) { + out := &Array{shape: index.Shape(), dt: src.dt} + out.alloc(index.Len()) + dst := make([]int, index.NDim()) + for i := range index.Len() { + idx := int(index.ints[index.physIndex(i)]) + if idx < 0 || idx >= src.shape[dim] { + return nil, errf("Gather: index %d out of range for dimension %d of size %d at position %d", idx, dim, src.shape[dim], i) + } + off := 0 + for d := range dst { + c := dst[d] + if d == dim { + c = idx + } + off = off*src.shape[d] + c + } + out.setFrom(i, src, off) + advanceOdometer(dst, index.shape) + } + return out, nil +} + +// refTake is the validated per-element copy Take replaced. +func refTake(a *Array, indices *Array) (*Array, error) { + out := &Array{shape: []int{indices.Len()}, dt: a.dt} + out.alloc(indices.Len()) + src := a + if !src.isContiguous() { + src = src.materialise() + } + k := intPayload(indices) + for i, v := range k { + if v < 0 || v >= int64(a.Len()) { + return nil, errf("Take: index %d out of range for flat size %d at position %d", v, a.Len(), i) + } + } + switch out.dt { + case Int: + for i, v := range k { + out.ints[i] = src.ints[v] + } + case Float16: + for i, v := range k { + out.halves[i] = src.halves[v] + } + case Float32: + for i, v := range k { + out.floats32[i] = src.floats32[v] + } + case Float: + for i, v := range k { + out.floats[i] = src.floats[v] + } + default: + for i, v := range k { + out.complexes[i] = src.complexes[v] + } + } + return out, nil +} + +// probeSearchValues returns the adversarial probe values every search +// equivalence test walks: exact edges, the two zeros, neighbours of the +// edges, out-of-range ends, both NaN signs, infinities, subnormals and +// a deterministic random spread. +func probeSearchValues(edges []float64) []float64 { + vals := make([]float64, 0, 64) + vals = append(vals, edges...) + for _, e := range edges { + vals = append(vals, + math.Nextafter(e, math.Inf(1)), + math.Nextafter(e, math.Inf(-1)), + e+0, + e-0) + } + vals = append(vals, + 0, math.Copysign(0, -1), + math.Inf(1), math.Inf(-1), + math.NaN(), math.Copysign(math.NaN(), -1), + 5e-324, -5e-324, 1e-320, -1e-320, + 1e300, -1e300) + g := NewGenerator(3) + rnd, err := Floats(g, 40) + if err != nil { + panic(err) + } + vals = append(vals, rnd.RawFloats()[:40]...) + return vals +} + +// TestBinSearchMatchesReference pins AssignBins' float bin selection +// against the bisection walk on every edge class: uniform edges, a +// single bin, repeated edges, negative ranges, infinite outer edges, +// subnormal edges and both large and small magnitudes. NaN edges are out +// of scope: an ascending edge set holds none. +func TestBinSearchMatchesReference(t *testing.T) { + edgeSets := [][]float64{ + {0, 1}, + {0, 0.25, 0.5, 0.75, 1}, + {1, 2, 3}, + {0, 0, 0, 1, 2, 2}, + {-5, -3.5, -1, 0, 0, 2}, + {math.Inf(-1), -1, 0, 1, math.Inf(1)}, + {0, 5e-324, 1e-320, 1}, + {1e300, 2e300, 3e300}, + linspaceVals(0, 1, 256), + linspaceVals(-3, 7, 257), + } + for _, ev := range edgeSets { + vals := probeSearchValues(ev) + edges := mustFloats(t, ev) + a := mustFloats(t, vals) + out, err := AssignBins(a, edges) + if err != nil { + t.Fatalf("AssignBins m=%d: %v", len(ev), err) + } + got := out.RawInts()[:out.Len()] + for i, v := range vals { + want := refBinFloat(ev, v) + if got[i] != int64(want) { + t.Fatalf("edges %v…%v (m=%d), v=%v: bin %d, want %d", + ev[0], ev[len(ev)-1], len(ev), v, got[i], want) + } + } + } +} + +// TestBinSearchMatchesReferenceInt pins the int bin selection, native +// int64 comparisons included: the extremes, the neighbours above 2^53 +// that a float64 detour would fold together, and repeated edges. +func TestBinSearchMatchesReferenceInt(t *testing.T) { + edgeSets := [][]int64{ + {0, 1}, + {1, 2, 3}, + {math.MinInt64, -1, 0, 1, math.MaxInt64}, + {1 << 53, 1<<53 + 1, 1<<53 + 2}, + {5, 5, 7, 9, 9}, + {math.MinInt64, math.MaxInt64}, + } + values := []int64{ + 0, 1, 2, 3, 5, 7, 9, + -1, 4, 6, 8, 10, + math.MinInt64, math.MaxInt64, math.MinInt64 + 1, math.MaxInt64 - 1, + 1 << 53, 1<<53 + 1, 1<<53 + 2, 1<<53 + 3, + -(1 << 53), -(1<<53 + 1), + } + for _, ev := range edgeSets { + edges, err := FromInts(ev, len(ev)) + if err != nil { + t.Fatal(err) + } + a, err := FromInts(values, len(values)) + if err != nil { + t.Fatal(err) + } + out, err := AssignBins(a, edges) + if err != nil { + t.Fatalf("AssignBins int m=%d: %v", len(ev), err) + } + got := out.RawInts()[:out.Len()] + for i, v := range values { + want := refBinInt(ev, v) + if got[i] != int64(want) { + t.Fatalf("int edges %v, v=%d: bin %d, want %d", ev, v, got[i], want) + } + } + } +} + +// TestAssignBinsPinsDocumentedSemantics fixes the observable edge rules +// with literals: below-range clamps to bin 0, at-or-above the last edge +// clamps to the outermost bin, an exact edge lands in the bin it opens, +// and a NaN keeps the outermost bin whatever its sign. +func TestAssignBinsPinsDocumentedSemantics(t *testing.T) { + edges := mustFloats(t, []float64{0, 1.0 / 3, 2.0 / 3, 1}) + vals := []float64{ + -1, 0, 1.0 / 3, 0.5, 2.0 / 3, 0.999, 1, 2, + math.NaN(), math.Copysign(math.NaN(), -1), + math.Inf(1), math.Inf(-1), math.Copysign(0, -1), + } + a := mustFloats(t, vals) + out, err := AssignBins(a, edges) + if err != nil { + t.Fatalf("AssignBins: %v", err) + } + want := []int64{0, 0, 1, 1, 2, 2, 2, 2, 2, 2, 2, 0, 0} + for i, w := range want { + if got := out.RawInts()[i]; got != w { + t.Fatalf("value %v (bits %#x): bin %d, want %d", vals[i], math.Float64bits(vals[i]), got, w) + } + } +} + +// TestSearchSortedMatchesReference pins the rightmost insertion rule +// against the bisection for both operand kinds through the public entry +// point, including the int64 range above 2^53, NaN needles, infinite +// needles and infinite haystack ends, ties and the two zeros. +func TestSearchSortedMatchesReference(t *testing.T) { + haySetsF := [][]float64{ + {1}, + {0, 0}, + {1, 2}, + {-1, 0, 0, 1, 4, 4, 4, 9}, + {math.Inf(-1), -3, 0, 3, math.Inf(1)}, + {math.Copysign(0, -1), 0, 0, 1}, + linspaceVals(0, 1, 33), + } + needleSetsF := [][]float64{ + {0, 0.5, 1, 4, 9, 10}, + {-1, math.Copysign(0, -1), 0, 3}, + {math.NaN(), math.Copysign(math.NaN(), -1), math.Inf(1), math.Inf(-1)}, + probeSearchValues([]float64{-1, 0, 1, 4}), + } + for _, h := range haySetsF { + for _, qs := range needleSetsF { + hay := mustFloats(t, h) + nd := mustFloats(t, qs) + out, err := SearchSorted(hay, nd) + if err != nil { + t.Fatalf("SearchSorted: %v", err) + } + got := out.RawInts()[:out.Len()] + for i, q := range qs { + want := refUpperFloat(h, q) + if got[i] != int64(want) { + t.Fatalf("haystack %v, needle %v: position %d, want %d", h, q, got[i], want) + } + } + } + } + haySetsI := [][]int64{ + {1}, + {0, 0, 1}, + {1 << 53, 1<<53 + 1, 1<<53 + 1, 1 << 54}, + {math.MinInt64, -1, 0, 1, math.MaxInt64}, + } + needleValsI := []int64{ + 0, 1, 2, 1 << 53, 1<<53 + 1, 1<<53 + 2, 1 << 54, + math.MinInt64, math.MaxInt64, -1, + } + for _, h := range haySetsI { + hay, err := FromInts(h, len(h)) + if err != nil { + t.Fatal(err) + } + nd, err := FromInts(needleValsI, len(needleValsI)) + if err != nil { + t.Fatal(err) + } + out, err := SearchSorted(hay, nd) + if err != nil { + t.Fatalf("SearchSorted int: %v", err) + } + got := out.RawInts()[:out.Len()] + for i, q := range needleValsI { + want := refUpperInt(h, q) + if got[i] != int64(want) { + t.Fatalf("int haystack %v, needle %d: position %d, want %d", h, q, got[i], want) + } + } + } + // An int haystack above 2^53 must search natively, and an empty + // haystack reports position zero everywhere. + hi, err := FromInts([]int64{1 << 53, 1<<53 + 1, 1 << 54}, 3) + if err != nil { + t.Fatal(err) + } + ni, err := FromInts([]int64{1 << 53, 1<<53 + 1, 1<<53 + 2}, 3) + if err != nil { + t.Fatal(err) + } + out, err := SearchSorted(hi, ni) + if err != nil { + t.Fatalf("SearchSorted: %v", err) + } + for i, w := range []int64{1, 2, 2} { + if got := out.RawInts()[i]; got != w { + t.Fatalf("int needle %d: position %d, want %d", ni.RawInts()[i], got, w) + } + } + empty, err := FromInts(nil, 0) + if err != nil { + t.Fatal(err) + } + out, err = SearchSorted(empty, ni) + if err != nil { + t.Fatalf("SearchSorted empty: %v", err) + } + for i := range 3 { + if got := out.RawInts()[i]; got != 0 { + t.Fatalf("empty haystack position %d = %d, want 0", i, got) + } + } +} + +// TestInterpolateSegmentMatchesReference pins the whole per-point +// pipeline against the bisection-derived reference across knot classes: +// strict ascent, repeated knots at the edges and in the interior, the +// minimal two-knot set, and queries on knots, between them, outside the +// range and at both infinities. The comparison is the result value's +// exact bits, which is the observable contract. +func TestInterpolateSegmentMatchesReference(t *testing.T) { + knotSets := [][]float64{ + {0, 1}, + {0, 1, 2, 3}, + {0, 0.5, 0.5, 0.5, 1}, + {0, 0, 1, 2}, + {0, 1, 2, 2}, + {-2, -1, -1, 0}, + {0, 1e-300, 2e-300, 3e-300}, + linspaceVals(0, 1, 33), + } + for _, xv := range knotSets { + ys := make([]float64, len(xv)) + for i := range ys { + ys[i] = float64(i*i-3*i+1) / 7 + } + queries := probeSearchValues([]float64{xv[0], xv[len(xv)/2], xv[len(xv)-1]}) + qv := make([]float64, 0, len(queries)) + for _, q := range queries { + if q != q { + continue // NaN queries are refused before the search + } + qv = append(qv, q) + } + xArr := mustFloats(t, xv) + yArr := mustFloats(t, ys) + qArr := mustFloats(t, qv) + out, err := Interpolate(xArr, yArr, qArr) + if err != nil { + t.Fatalf("Interpolate knots %v…%v: %v", xv[0], xv[len(xv)-1], err) + } + for i, q := range qv { + lo := refInterpSegment(xv, q) + x0, x1 := xv[lo], xv[lo+1] + y0, y1 := ys[lo], ys[lo+1] + tt := 0.0 + if x1 > x0 { + tt = (q - x0) / (x1 - x0) + } else if q > x0 { + tt = 1 + } + if tt < 0 { + tt = 0 + } + if tt > 1 { + tt = 1 + } + want := y0 + tt*(y1-y0) + got := out.RawFloats()[i] + if math.Float64bits(got) != math.Float64bits(want) { + t.Fatalf("knots %v…%v (n=%d), q=%v: value %#x, want %#x", + xv[0], xv[len(xv)-1], len(xv), q, math.Float64bits(got), math.Float64bits(want)) + } + } + } +} + +// TestInterpolateMatchesReferencePointwise pins the whole per-point +// pipeline: same segment, same t expression, same bits, for random +// queries against repeated-knot and strictly ascending knot sets. +func TestInterpolateMatchesReferencePointwise(t *testing.T) { + for _, knots := range [][]float64{ + {0, 0.5, 0.5, 1, 2, 2, 3}, + linspaceVals(-1, 2, 17), + } { + ys := make([]float64, len(knots)) + for i := range ys { + ys[i] = float64(i*i-3*i+1) / 7 + } + g := NewGenerator(9) + rnd, err := Floats(g, 200) + if err != nil { + t.Fatal(err) + } + queries := make([]float64, 0, 220) + queries = append(queries, rnd.RawFloats()[:200]...) + for _, k := range knots { + queries = append(queries, k) + } + queries = append(queries, -2, 5, math.Inf(1), math.Inf(-1), math.Copysign(0, -1)) + xv := mustFloats(t, knots) + yv := mustFloats(t, ys) + qv := mustFloats(t, queries) + out, err := Interpolate(xv, yv, qv) + if err != nil { + t.Fatalf("Interpolate: %v", err) + } + n := len(knots) + for i, q := range queries { + lo := refInterpSegment(knots, q) + x0, x1 := knots[lo], knots[lo+1] + y0, y1 := ys[lo], ys[lo+1] + tt := 0.0 + if x1 > x0 { + tt = (q - x0) / (x1 - x0) + } else if q > x0 { + tt = 1 + } + if tt < 0 { + tt = 0 + } + if tt > 1 { + tt = 1 + } + want := y0 + tt*(y1-y0) + got := out.RawFloats()[i] + if math.Float64bits(got) != math.Float64bits(want) { + t.Fatalf("knots n=%d, query %v: value %#x, want %#x", + n, q, math.Float64bits(got), math.Float64bits(want)) + } + } + } +} + +// TestInterpolateNaNQueryReportsSmallestIndex pins the error contract +// under the parallel walk: the first NaN in order names the error. +func TestInterpolateNaNQueryReportsSmallestIndex(t *testing.T) { + xs := mustFloats(t, []float64{0, 1, 2}) + ys := mustFloats(t, []float64{0, 1, 2}) + q := mustFloats(t, []float64{0.5, 0.5, math.NaN(), 0.5, math.NaN(), 0.5}) + if _, err := Interpolate(xs, ys, q); err == nil || !strings.Contains(err.Error(), "query 2 is NaN") { + t.Fatalf("Interpolate NaN error = %v, want the query 2 report", err) + } +} + +// TestArgwhereMatchesReference pins the coordinate packing of the +// count-then-fill against the single odometer walk, across dtypes, +// ranks, densities and a chunk-boundary size; -0.0 counts as zero and +// NaN counts as non-zero, as the value tests have always had it. +func TestArgwhereMatchesReference(t *testing.T) { + shapes := [][]int{ + {1}, {7}, {5, 4}, {3, 5, 4}, {2, 2, 2, 2}, {80, 80}, {1, 1, 5}, + } + for _, shape := range shapes { + n := 1 + for _, d := range shape { + n *= d + } + for _, density := range []int{0, 1, 2, 3, 64} { + // density: 0 all-zero, 1 all-non-zero, k every k-th non-zero. + valsF := make([]float64, n) + valsI := make([]int64, n) + for i := range n { + if density == 1 || (density > 1 && i%density == 0) { + valsF[i] = float64(i+1) / 3 + valsI[i] = int64(i) + 1 + } + } + if density > 0 && n > 3 { + valsF[2] = math.NaN() + valsF[3] = math.Copysign(0, -1) // zero + valsF[4] = math.Inf(1) // non-zero + } + af, _ := FromFloats(valsF, shape...) + ai, _ := FromInts(valsI, shape...) + for _, a := range []*Array{af, ai} { + want := refArgwhere(a) + out, err := Argwhere(a) + if err != nil { + t.Fatalf("Argwhere %v density %d: %v", shape, density, err) + } + got := out.RawInts()[:out.Len()] + if len(got) != len(want) { + t.Fatalf("Argwhere %v density %d: %d coordinates, want %d", shape, density, len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("Argwhere %v density %d: coordinate %d = %d, want %d", shape, density, i, got[i], want[i]) + } + } + if out.Shape()[0]*out.Shape()[1] != len(want) || out.Shape()[1] != a.NDim() { + t.Fatalf("Argwhere %v: shape %v does not pack %d coordinates of rank %d", + shape, out.Shape(), len(want), a.NDim()) + } + } + } + } + // A rebased view: the walk is bounded by the view's own extent, and + // the coordinates are the view's. + base := mustFloats(t, []float64{0, 5, 0, 0, 7, 9, 0, 1, 0, 0, 0, 2}) + view, err := Slice(base, 0, 3, 11) + if err != nil { + t.Fatal(err) + } + v2, err := Reshape(view, 2, 4) + if err != nil { + t.Fatal(err) + } + want := refArgwhere(v2) + out, err := Argwhere(v2) + if err != nil { + t.Fatalf("Argwhere view: %v", err) + } + got := out.RawInts()[:out.Len()] + if len(got) != len(want) { + t.Fatalf("Argwhere view: %d coordinates, want %d", len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("Argwhere view: coordinate %d = %d, want %d", i, got[i], want[i]) + } + } +} + +// TestGatherTakeMatchReference pins the parallel payload walks against +// the per-element accessor walks, values and error reports alike, over +// every dtype, several ranks, repeated and out-of-range indices, view +// indices and view sources. +func TestGatherTakeMatchReference(t *testing.T) { + mk := func(vals []float64, shape ...int) *Array { + a, err := FromFloats(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a + } + srcF := mk([]float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + }, 3, 4) + srcI, _ := FromInts([]int64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + }, 3, 4) + srcC, _ := FromComplexes([]complex128{1, 2i, 3, 4i, 5, 6i}, 3, 2) + src3d := mk([]float64{ + 1, 2, 3, 4, 5, 6, 7, 8, + 9, 10, 11, 12, 13, 14, 15, 16, + }, 2, 2, 4) + idxSets := []*Array{ + mustInts(t, []int64{3, 0, 2, 2, 1, 3}, 3, 2), // for dim 1: (3, 2) + mustInts(t, []int64{2, 0, 2, 1}, 2, 2), // for dim 0: (2, 4) reshaped below + mustInts(t, []int64{0, 3, 2, 1, 1, 3, 0, 2}, 2, 4), + } + for _, src := range []*Array{srcF, srcI} { + for dim := range 2 { + for _, idx := range idxSets { + if !gatherCompatible(src.shape, idx.shape, dim) { + continue + } + want, wantErr := refGather(src, dim, idx) + got, gotErr := Gather(src, dim, idx) + if (wantErr == nil) != (gotErr == nil) || (wantErr != nil && wantErr.Error() != gotErr.Error()) { + t.Fatalf("Gather dim %d: error %v, want %v", dim, gotErr, wantErr) + } + if wantErr == nil && !Equal(got, want) { + t.Fatalf("Gather dim %d idx %v: %v, want %v", dim, idx.RawInts(), got, want) + } + } + } + } + // Complex source gathers unchanged; the index keeps the source's + // non-dim extent, as gatherCompatible requires. + idxC := mustInts(t, []int64{2, 0, 1, 1}, 2, 2) + wantC, _ := refGather(srcC, 0, idxC) + gotC, err := Gather(srcC, 0, idxC) + if err != nil || !Equal(gotC, wantC) { + t.Fatalf("Gather complex: %v vs %v (%v)", gotC, wantC, err) + } + // 3-D source, gather along the interior dimension: the index keeps + // the outer and inner extents and varies along dim 1. + idx3 := mustInts(t, []int64{1, 0, 1, 1, 0, 1, 0, 0}, 2, 1, 4) + want3, _ := refGather(src3d, 1, idx3) + got3, err := Gather(src3d, 1, idx3) + if err != nil || !Equal(got3, want3) { + t.Fatalf("Gather 3-D dim 1: %v vs %v (%v)", got3, want3, err) + } + // Out-of-range and negative indices: same error, same position. + bad := mustInts(t, []int64{0, 4, 2, 1, 0, 0, 0, 0}, 2, 4) + _, wantErr := refGather(srcF, 0, bad) + _, gotErr := Gather(srcF, 0, bad) + if wantErr == nil || gotErr == nil || wantErr.Error() != gotErr.Error() { + t.Fatalf("Gather range error: %v, want %v", gotErr, wantErr) + } + neg := mustInts(t, []int64{1, -1, 0, 0, 0, 0, 0, 0}, 2, 4) + _, wantErr = refGather(srcF, 0, neg) + _, gotErr = Gather(srcF, 0, neg) + if wantErr == nil || gotErr == nil || wantErr.Error() != gotErr.Error() { + t.Fatalf("Gather negative error: %v, want %v", gotErr, wantErr) + } + // A rebased view on the index side: the walk is bounded by the + // view's own extent, and the invisible payload tail never reads. + idxBase := mustInts(t, []int64{0, 1, 2, 0, 2, 1, 0, 2, 99, 99, 99, 99}, 3, 4) + idxView, err := Slice(idxBase, 0, 1, 2) + if err != nil { + t.Fatal(err) + } + wantV, _ := refGather(srcF, 0, idxView) + gotV, err := Gather(srcF, 0, idxView) + if err != nil || !Equal(gotV, wantV) { + t.Fatalf("Gather view index: %v vs %v (%v)", gotV, wantV, err) + } + + // Take: every dtype, view source, view indices, range errors. + for _, src := range []*Array{srcF, srcI, srcC} { + tk := mustInts(t, []int64{int64(src.Len()) - 1, 0, 3, 3, 1}, 5) + wantT, wantErr := refTake(src, tk) + gotT, gotErr := Take(src, tk) + if (wantErr == nil) != (gotErr == nil) || (wantErr != nil && wantErr.Error() != gotErr.Error()) { + t.Fatalf("Take: error %v, want %v", gotErr, wantErr) + } + if wantErr == nil && !Equal(gotT, wantT) { + t.Fatalf("Take: %v, want %v", gotT, wantT) + } + } + // Take on a rebased view source: flat indices count from the view's + // origin, and the walk is bounded by the view's own extent. + longBase := mustFloats(t, []float64{99, 99, 99, 99, 10, 11, 12, 13, 14, 15, 16, 17}) + viewSrc, err := Slice(longBase, 0, 4, 12) + if err != nil { + t.Fatal(err) + } + tk := mustInts(t, []int64{7, 0, 3}, 3) + wantT, _ := refTake(viewSrc, tk) + gotT, err := Take(viewSrc, tk) + if err != nil || !Equal(gotT, wantT) { + t.Fatalf("Take view source: %v vs %v (%v)", gotT, wantT, err) + } + badT := mustInts(t, []int64{1, 2, 99}, 3) + _, wantErr = refTake(srcF, badT) + _, gotErr = Take(srcF, badT) + if wantErr == nil || gotErr == nil || wantErr.Error() != gotErr.Error() { + t.Fatalf("Take range error: %v, want %v", gotErr, wantErr) + } +} + +// TestIntegrateHalvingIsBitIdentical proves the multiplication by 0.5 +// reproduces the division by 2 bit for bit on adversarial payloads: +// subnormals, the signed zeros, infinities, NaN payloads and random +// spreads, for a range of spacings. The accumulation order is untouched +// either way. +func TestIntegrateHalvingIsBitIdentical(t *testing.T) { + payloads := [][]float64{ + {1, 2, 3, 4}, + {5e-324, 1e-320, 2.5e-320, 1e-308}, + {math.Copysign(0, -1), 0, 1, math.Copysign(0, -1)}, + {math.Inf(1), 1, math.Inf(-1), 2}, + {math.NaN(), 1, 2, math.Copysign(math.NaN(), -1)}, + {1e308, 1.5e308, 2, 3}, + {-1e308, 1e308, 1, -1}, + } + g := NewGenerator(21) + rnd, err := Floats(g, 128) + if err != nil { + t.Fatal(err) + } + payloads = append(payloads, rnd.RawFloats()[:128]) + for _, dx := range []float64{1, 0.5, 2, 1e-300, 1e300, -3, math.NaN(), math.Inf(1)} { + for _, y := range payloads { + yArr := mustFloats(t, y) + gotTotal, err := Integrate(yArr, dx) + if err != nil { + t.Fatalf("Integrate: %v", err) + } + var want float64 + for i := 1; i < len(y); i++ { + want += (y[i-1] + y[i]) / 2 + } + want *= dx + if math.Float64bits(gotTotal) != math.Float64bits(want) { + t.Fatalf("Integrate dx=%v payload %v…: %#x, want %#x", + dx, y[0], math.Float64bits(gotTotal), math.Float64bits(want)) + } + gotCum, err := CumulativeIntegrate(yArr, dx) + if err != nil { + t.Fatalf("CumulativeIntegrate: %v", err) + } + ov := make([]float64, len(y)) + for i := 1; i < len(y); i++ { + ov[i] = ov[i-1] + (y[i-1]+y[i])/2*dx + } + for i := range y { + if math.Float64bits(gotCum.RawFloats()[i]) != math.Float64bits(ov[i]) { + t.Fatalf("CumulativeIntegrate dx=%v at %d: %#x, want %#x", + dx, i, math.Float64bits(gotCum.RawFloats()[i]), math.Float64bits(ov[i])) + } + } + } + } +} + +// TestArgwhereVariantsMatchReference pins every packing variant the +// probe walks, production included, against the single odometer walk: +// the coordinates, their order and the packed shape must match whatever +// the buffers do on the way. +func TestArgwhereVariantsMatchReference(t *testing.T) { + variants := append(argwhereStyles(), + struct { + name string + run func(*Array) *Array + }{"merge-capped-1024", func(a *Array) *Array { return probeArgwhereMerge(a, 1024) }}, + ) + shapes := [][]int{{1}, {7}, {5, 4}, {3, 5, 4}, {80, 80}, {2, 2, 2, 2}} + for _, shape := range shapes { + n := 1 + for _, d := range shape { + n *= d + } + for _, density := range []int{0, 1, 2, 64} { + valsF := make([]float64, n) + for i := range n { + if density == 1 || (density > 1 && i%density == 0) { + valsF[i] = float64(i+1) / 3 + } + } + if density > 0 && n > 4 { + valsF[2] = math.NaN() + valsF[3] = math.Copysign(0, -1) + valsF[4] = math.Inf(1) + } + a, _ := FromFloats(valsF, shape...) + want := refArgwhere(a) + for _, v := range variants { + out := v.run(a) + if out == nil { + t.Fatalf("%s %v density %d: nil result", v.name, shape, density) + } + got := out.RawInts()[:out.Len()] + if len(got) != len(want) { + t.Fatalf("%s %v density %d: %d coordinates, want %d", v.name, shape, density, len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("%s %v density %d: coordinate %d = %d, want %d", v.name, shape, density, i, got[i], want[i]) + } + } + } + } + } + // A rebased view: bounded by the view's own extent for every variant. + base := mustFloats(t, []float64{0, 5, 0, 0, 7, 9, 0, 1, 0, 0, 0, 2}) + view, err := Slice(base, 0, 3, 11) + if err != nil { + t.Fatal(err) + } + v2, err := Reshape(view, 2, 4) + if err != nil { + t.Fatal(err) + } + want := refArgwhere(v2) + for _, v := range variants { + out := v.run(v2) + got := out.RawInts()[:out.Len()] + if len(got) != len(want) { + t.Fatalf("%s view: %d coordinates, want %d", v.name, len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("%s view: coordinate %d = %d, want %d", v.name, i, got[i], want[i]) + } + } + } +} + +// TestSortRadixConfigsBitIdentical pins every probe radix configuration +// against the production Sort and ArgSort on adversarial fixtures: the +// sorted values bit for bit (NaN payloads included) and the +// permutations exactly, whatever crew cap and digit width carry them. +func TestSortRadixConfigsBitIdentical(t *testing.T) { + mkInts := func(vals []int64) *Array { + a, err := FromInts(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + return a + } + intRand := benchIntsShape(1000, 5000) + intWide := benchIntsShape(1<<30, 5000) + intParallel := benchIntsShape(1000, 20000) + floatRand := benchFloatsShape(5000) + floatParallel := benchFloatsShape(20000) + fixtures := []*Array{ + intRand, + intWide, + intParallel, + mkInts([]int64{math.MinInt64, math.MaxInt64, 0, -1, 1, 1 << 53, 1<<53 + 1, math.MinInt64 + 1, math.MaxInt64 - 1, 5, 5, 5}), + mkInts([]int64{42, 42, 42, 42, 42}), + floatRand, + floatParallel, + mustFloats(t, []float64{math.NaN(), math.Copysign(math.NaN(), -1), math.Copysign(0, -1), 0, math.Inf(1), math.Inf(-1), 1, 1, 2.5, 2.5, -3, 0.1}), + mustFloats(t, []float64{7, 7, 7, 7, 7}), + } + cfgs := probeSortConfigs() + for _, a := range fixtures { + wantSort, err := Sort(a) + if err != nil { + t.Fatal(err) + } + wantIdx, err := ArgSort(a) + if err != nil { + t.Fatal(err) + } + for _, c := range cfgs { + var gotSort *Array + if a.dt == Int { + gotSort = probeSortInt(a, c.wcap, c.bits) + } else { + gotSort = probeSortFloat(a, c.wcap, c.bits) + } + if gotSort.Len() != wantSort.Len() { + t.Fatalf("%v %s: sorted length %d, want %d", a.Shape(), c.name, gotSort.Len(), wantSort.Len()) + } + if a.dt == Int { + g, w := gotSort.RawInts()[:gotSort.Len()], wantSort.RawInts()[:wantSort.Len()] + for i := range w { + if g[i] != w[i] { + t.Fatalf("int fixture %v %s: sorted[%d] = %d, want %d", a.Shape(), c.name, i, g[i], w[i]) + } + } + } else { + g, w := gotSort.RawFloats()[:gotSort.Len()], wantSort.RawFloats()[:wantSort.Len()] + for i := range w { + if math.Float64bits(g[i]) != math.Float64bits(w[i]) { + t.Fatalf("float fixture %v %s: sorted[%d] = %#x, want %#x", a.Shape(), c.name, i, math.Float64bits(g[i]), math.Float64bits(w[i])) + } + } + } + gotIdx := probeArgSort(a, c.wcap, c.bits) + g, w := gotIdx.RawInts()[:gotIdx.Len()], wantIdx.RawInts()[:wantIdx.Len()] + for i := range w { + if g[i] != w[i] { + t.Fatalf("fixture %v %s: argsort[%d] = %d, want %d", a.Shape(), c.name, i, g[i], w[i]) + } + } + } + } +} + +// linspaceVals builds n evenly spaced values, the deterministic edge +// ladder the bin equivalence tests walk. +func linspaceVals(start, stop float64, n int) []float64 { + a, err := Linspace(start, stop, n) + if err != nil { + panic(err) + } + return append([]float64(nil), a.RawFloats()[:n]...) +} + +func mustInts(t *testing.T, vals []int64, shape ...int) *Array { + t.Helper() + a, err := FromInts(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} diff --git a/internal/core/reshape.go b/internal/core/reshape.go new file mode 100644 index 0000000..a7f415e --- /dev/null +++ b/internal/core/reshape.go @@ -0,0 +1,375 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +// Reshape and padding utilities for the tensor numeric surface. Flatten, +// Squeeze, Unsqueeze and TransposeAxes change the shape without changing +// the element order; Copy returns a fresh copy with the same layout +// (every array is row-major, so Copy is always a real allocation); +// Pad extends arrays along one or more dimensions with constant, +// reflect, replicate or circular modes. + +// Flatten returns a copy with the dimensions in [startDim, endDim] +// collapsed into one. Negative indices count from the end (endDim = -1 +// means the last dimension). startDim > endDim is an error. +func Flatten(a *Array, startDim, endDim int) (*Array, error) { + ndim := a.NDim() + if startDim < 0 { + startDim += ndim + } + if endDim < 0 { + endDim += ndim + } + if startDim < 0 || startDim >= ndim || endDim < 0 || endDim >= ndim || startDim > endDim { + return nil, errf("Flatten: range [%d, %d] out of bounds for shape %s", startDim, endDim, shapeText(a.shape)) + } + flat := 1 + for d := startDim; d <= endDim; d++ { + flat *= a.shape[d] + } + newShape := make([]int, 0, ndim-(endDim-startDim)) + newShape = append(newShape, a.shape[:startDim]...) + newShape = append(newShape, flat) + newShape = append(newShape, a.shape[endDim+1:]...) + return Reshape(a, newShape...) +} + +// Squeeze returns a copy with size-1 dimensions removed. When dim is -1 +// every size-1 dimension is dropped; otherwise only that one is (and it +// must have size 1). +func Squeeze(a *Array, dim int) (*Array, error) { + if dim == -1 { + newShape := make([]int, 0, a.NDim()) + for _, d := range a.shape { + if d != 1 { + newShape = append(newShape, d) + } + } + if len(newShape) == 0 { + newShape = []int{1} + } + return Reshape(a, newShape...) + } + if dim < 0 || dim >= a.NDim() { + return nil, errf("Squeeze: dimension %d is out of range for shape %s", dim, shapeText(a.shape)) + } + if a.shape[dim] != 1 { + return nil, errf("Squeeze: dimension %d has size %d, must be 1", dim, a.shape[dim]) + } + newShape := make([]int, 0, a.NDim()-1) + newShape = append(newShape, a.shape[:dim]...) + newShape = append(newShape, a.shape[dim+1:]...) + if len(newShape) == 0 { + newShape = []int{1} + } + return Reshape(a, newShape...) +} + +// Unsqueeze returns a copy with a new size-1 dimension inserted at dim. +// dim is in [0, NDim]; negative values count from the end of the result +// rank, so dim = -1 appends the new axis at the end. +func Unsqueeze(a *Array, dim int) (*Array, error) { + ndim := a.NDim() + 1 + if dim < 0 { + dim += a.NDim() + 1 + } + if dim < 0 || dim >= ndim { + return nil, errf("Unsqueeze: dimension %d is out of range for inserting into shape %s", dim, shapeText(a.shape)) + } + newShape := make([]int, 0, ndim) + newShape = append(newShape, a.shape[:dim]...) + newShape = append(newShape, 1) + newShape = append(newShape, a.shape[dim:]...) + return Reshape(a, newShape...) +} + +// TransposeAxes returns a copy with the dimensions reordered according to +// dims. dims must be a permutation of [0, NDim); an error names the +// bad permutation. Renamed from Permute to avoid the clash with the +// random-number Fisher-Yates permutation in the generator. +func TransposeAxes(a *Array, dims ...int) (*Array, error) { + if len(dims) != a.NDim() { + return nil, errf("TransposeAxes: needs %d dimensions, got %d", a.NDim(), len(dims)) + } + seen := make([]bool, a.NDim()) + for _, d := range dims { + if d < 0 || d >= a.NDim() || seen[d] { + return nil, errf("TransposeAxes: invalid permutation %v for shape %s", dims, shapeText(a.shape)) + } + seen[d] = true + } + newShape := make([]int, a.NDim()) + for i, d := range dims { + newShape[i] = a.shape[d] + } + total := a.Len() + out := &Array{shape: newShape, dt: a.dt} + out.alloc(total) + srcCoord := make([]int, a.NDim()) + // The dtype dispatch keeps setFrom's per-element switch out of the + // walk; the destination index still needs the coordinate fold. Bool + // and the narrow integer widths take the setFrom walk directly: the + // same values in the same order, with no per-dtype loop of their own. + switch a.dt { + case Int: + for i := range total { + dst := 0 + for k := range a.NDim() { + dst = dst*newShape[k] + srcCoord[dims[k]] + } + out.ints[dst] = a.ints[i] + advanceOdometer(srcCoord, a.shape) + } + case Float16: + for i := range total { + dst := 0 + for k := range a.NDim() { + dst = dst*newShape[k] + srcCoord[dims[k]] + } + out.halves[dst] = a.halves[i] + advanceOdometer(srcCoord, a.shape) + } + case Float32: + for i := range total { + dst := 0 + for k := range a.NDim() { + dst = dst*newShape[k] + srcCoord[dims[k]] + } + out.floats32[dst] = a.floats32[i] + advanceOdometer(srcCoord, a.shape) + } + case Float: + for i := range total { + dst := 0 + for k := range a.NDim() { + dst = dst*newShape[k] + srcCoord[dims[k]] + } + out.floats[dst] = a.floats[i] + advanceOdometer(srcCoord, a.shape) + } + case Complex: + for i := range total { + dst := 0 + for k := range a.NDim() { + dst = dst*newShape[k] + srcCoord[dims[k]] + } + out.complexes[dst] = a.complexes[i] + advanceOdometer(srcCoord, a.shape) + } + default: + for i := range total { + dst := 0 + for k := range a.NDim() { + dst = dst*newShape[k] + srcCoord[dims[k]] + } + out.setFrom(dst, a, i) + advanceOdometer(srcCoord, a.shape) + } + } + return out, nil +} + +// Copy returns a fresh array with the same data, shape and dtype as a. +// Every tensor is already row-major in storage, so Copy is +// always a real allocation: there are no strides to collapse, no +// views to flatten. Useful when a caller wants a guaranteed-new +// buffer they can hand off without worrying about aliasing. +func Copy(a *Array) *Array { + // cloneArray carries every payload the dtype owns; the five-slice + // cloneData leaves the narrow element types empty. + return a.cloneArray() +} + +// Pad extends an array along its trailing dimensions. pad is a flat +// sequence of pairs in reverse spatial order: +// - 2-D: pad = (left, right, top, bottom) +// - 3-D: pad = (left, right, top, bottom, front, back) +// +// mode is one of "constant" (fill with value), "reflect" (mirror without +// repeating the edge), "replicate" (repeat the edge), "circular" (wrap +// around). Pad returns an error for an unsupported mode, a malformed +// pad argument, or a complex input under non-constant modes. +func Pad(a *Array, pad []int, mode string, value float64) (*Array, error) { + if len(pad)%2 != 0 { + return nil, errf("Pad: pad must hold a pre/post pair per dimension, got the odd count %d values", len(pad)) + } + switch mode { + case "constant", "reflect", "replicate", "circular": + default: + return nil, errf("Pad: unknown mode %q", mode) + } + if mode != "constant" { + // The folding modes index back into the source: an empty + // dimension has no element to mirror, repeat or wrap, and the + // fold loops below would never terminate (or would panic) on + // one. + for d := range a.shape { + if a.shape[d] == 0 { + return nil, errf("Pad: mode %q is not defined for an empty dimension (%d has size 0)", mode, d) + } + } + } + if a.dt == Complex && mode != "constant" { + return nil, errf("Pad: mode %q is not supported for complex arrays", mode) + } + // Pair count must equal rank (each dim gets a left/right). + if len(pad)/2 != a.NDim() { + return nil, errf("Pad: shape %s needs %d pad values, got %d", shapeText(a.shape), a.NDim()*2, len(pad)) + } + // Reverse pad so it pairs with shape left-to-right: the pad list is + // laid out from the last dimension to the first (2-D: left, right, + // top, bottom = dim 1 pre/post, dim 0 pre/post). rpad[i*2], + // rpad[i*2+1] is the pre/post pair for dimension i. + rpad := make([]int, a.NDim()*2) + for d := range a.shape { + srcIdx := (a.NDim() - 1 - d) * 2 + rpad[2*d] = pad[srcIdx] + rpad[2*d+1] = pad[srcIdx+1] + } + newShape := make([]int, a.NDim()) + for d := range a.shape { + // A negative pad has no meaning for the fold below and would + // drive the new shape negative, which alloc would then reject + // with a makeslice panic instead of an error. + if rpad[2*d] < 0 || rpad[2*d+1] < 0 { + return nil, errf("Pad: pad values must be non-negative, got %d and %d for dimension %d", + rpad[2*d], rpad[2*d+1], d) + } + // Bound each extent before adding: hostile pad values must be + // an error, not a wrapped extent that alloc turns into a + // makeslice panic. + extent := a.shape[d] + rpad[2*d] + if extent < rpad[2*d] { + return nil, errf("Pad: pad values overflow dimension %d", d) + } + extent += rpad[2*d+1] + if extent < rpad[2*d+1] { + return nil, errf("Pad: pad values overflow dimension %d", d) + } + newShape[d] = extent + } + if mode == "reflect" { + // A reflection can only fold once: a pad of n or more on an + // axis of length n has no source left to mirror, and folding + // again would hand setFrom a negative offset. + for d := range a.shape { + if a.shape[d] > 1 && (rpad[2*d] >= a.shape[d] || rpad[2*d+1] >= a.shape[d]) { + return nil, errf("Pad: reflect pad %d exceeds dimension %d of length %d", + max(rpad[2*d], rpad[2*d+1]), d, a.shape[d]) + } + } + } + // The padded extents are checked like every constructor's shape: a + // product that wraps must be an error, not a tiny allocation paired + // with a huge shape. + total, _, terr := checkedDims(newShape) + if terr != nil { + return nil, base.WrapErr("Pad", terr) + } + out := &Array{shape: newShape, dt: a.dt} + out.alloc(total) + // Fill by source coordinates. + dstCoord := make([]int, a.NDim()) + srcCoord := make([]int, a.NDim()) + for i := range total { + for d := range a.NDim() { + s := dstCoord[d] - rpad[2*d] + switch mode { + case "constant": + if s < 0 || s >= a.shape[d] { + s = -1 // marker: fill with value + } + case "reflect": + if a.shape[d] == 1 { + s = 0 + break + } + for s < 0 { + s = -s + } + for s >= a.shape[d] { + s = 2*a.shape[d] - 2 - s + } + case "replicate": + if s < 0 { + s = 0 + } + if s >= a.shape[d] { + s = a.shape[d] - 1 + } + case "circular": + // One modulo pair folds any offset into range in + // constant time: Go's % keeps the sign of s, so the + // addition of n normalises the negative side and the + // second % lands in [0, n). The fold loops this + // replaces walked one step at a time, which turned a + // large pad on a short axis into a quadratic fold. + s = ((s % a.shape[d]) + a.shape[d]) % a.shape[d] + default: + return nil, errf("Pad: unknown mode %q", mode) + } + srcCoord[d] = s + } + if mode == "constant" && containsNeg(srcCoord) { + out.setFromValue(i, value) + } else { + off := 0 + for d := range srcCoord { + off = off*a.shape[d] + srcCoord[d] + } + out.setFrom(i, a, off) + } + advanceOdometer(dstCoord, newShape) + } + return out, nil +} + +// setFromValue sets element i of the result to the given float, widening +// to the array's dtype: a half payload narrows under the +// HalfFromFloat64 contract, and the narrow integer widths and bool take +// the same implicit-store cast filled() carries (Go's conversion through +// int64, v != 0 for bool). Used by Pad's constant mode. +func (a *Array) setFromValue(i int, v float64) { + switch a.dt { + case Int: + a.ints[i] = int64(v) + case Bool: + a.bools[i] = v != 0 + case Int8: + a.i8s[i] = int8(int64(v)) + case Uint8: + a.u8s[i] = uint8(int64(v)) + case Int16: + a.i16s[i] = int16(int64(v)) + case Uint16: + a.u16s[i] = uint16(int64(v)) + case Int32: + a.i32s[i] = int32(int64(v)) + case Uint32: + a.u32s[i] = uint32(int64(v)) + case Float16: + a.halves[i] = HalfFromFloat64(v) + case Float32: + a.floats32[i] = float32(v) + case Float: + a.floats[i] = v + case Complex: + a.complexes[i] = complex(v, 0) + default: + // The alloc convention: an unlisted ordinal carries the uint32 + // payload, and no validated constructor can produce one. + a.u32s[i] = uint32(int64(v)) + } +} + +func containsNeg(c []int) bool { + for _, v := range c { + if v < 0 { + return true + } + } + return false +} diff --git a/internal/core/reshape_pad_test.go b/internal/core/reshape_pad_test.go new file mode 100644 index 0000000..35b12c2 --- /dev/null +++ b/internal/core/reshape_pad_test.go @@ -0,0 +1,544 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +func TestFlatten(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + got, err := Flatten(a, 0, -1) + if err != nil { + t.Fatal(err) + } + want, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 6) + if got.String() != want.String() { + t.Errorf("flatten all: got %s, want %s", got, want) + } + // Flatten only middle dimension. + got, err = Flatten(a, 1, 1) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 2 || got.Shape()[1] != 3 { + t.Errorf("flatten 1: shape = %v", got.Shape()) + } + // Out-of-range. + if _, err := Flatten(a, 0, 5); err == nil { + t.Error("flatten: expected error for out-of-range endDim") + } +} + +func TestSqueeze(t *testing.T) { + a, _ := FromFloats([]float64{1, 2, 3, 4}, 1, 2, 1, 2) + got, err := Squeeze(a, 0) + if err != nil { + t.Fatal(err) + } + if len(got.Shape()) != 3 || got.Shape()[0] != 2 { + t.Errorf("squeeze dim 0: shape = %v", got.Shape()) + } + // Squeeze all (-1). + got, err = Squeeze(a, -1) + if err != nil { + t.Fatal(err) + } + if len(got.Shape()) != 2 || got.Shape()[0] != 2 { + t.Errorf("squeeze all: shape = %v", got.Shape()) + } + // Cannot squeeze non-1 dimension. + if _, err := Squeeze(a, 1); err == nil { + t.Error("squeeze: expected error for non-1 dim") + } +} + +func TestUnsqueeze(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 3) + got, err := Unsqueeze(a, 0) + if err != nil { + t.Fatal(err) + } + if len(got.Shape()) != 2 || got.Shape()[0] != 1 || got.Shape()[1] != 3 { + t.Errorf("unsqueeze 0: shape = %v", got.Shape()) + } + // Negative dim. + got, err = Unsqueeze(a, -1) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 3 || got.Shape()[1] != 1 { + t.Errorf("unsqueeze -1: shape = %v", got.Shape()) + } +} + +func TestTransposeAxes(t *testing.T) { + a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + got, err := TransposeAxes(a, 1, 0) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 3 || got.Shape()[1] != 2 { + t.Errorf("transpose: shape = %v", got.Shape()) + } + // Check element: a[1,0] = 4 should land at out[0,1]. + v, _ := FloatAt(got, 0, 1) + if v != 4 { + t.Errorf("transpose element: got %v, want 4", v) + } + // Invalid permutation. + if _, err := TransposeAxes(a, 0, 0); err == nil { + t.Error("transpose: expected error for duplicate dims") + } +} + +func TestCopy(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + got := Copy(a) + if got.Shape()[0] != 2 || got.Shape()[1] != 2 { + t.Errorf("copy: shape = %v", got.Shape()) + } +} + +func TestPadConstant(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + // 2-D: pad = (left, right, top, bottom). + got, err := Pad(a, []int{1, 1, 1, 1}, "constant", 0) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 4 || got.Shape()[1] != 4 { + t.Errorf("pad constant: shape = %v", got.Shape()) + } + // Corners should be zero. + c, _ := FloatAt(got, 0, 0) + if c != 0 { + t.Errorf("pad corner: got %v, want 0", c) + } + // Centre should be original a[0,0] = 1. + c, _ = FloatAt(got, 1, 1) + if c != 1 { + t.Errorf("pad centre: got %v, want 1", c) + } + // Wrong pad length. + if _, err := Pad(a, []int{1, 1, 1}, "constant", 0); err == nil { + t.Error("pad: expected error for odd-length pad") + } + // Unknown mode. + if _, err := Pad(a, []int{1, 1, 1, 1}, "weird", 0); err == nil { + t.Error("pad: expected error for unknown mode") + } +} + +func TestPadReflect(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 1, 3) + got, err := Pad(a, []int{2, 0, 0, 0}, "reflect", 0) + if err != nil { + t.Fatal(err) + } + // Reflect without repeating edge: at index 0 we mirror index 2 -> 3, + // at index 1 we mirror index 1 -> 2. + v, _ := FloatAt(got, 0, 0) + if v != 3 { + t.Errorf("pad reflect [0]: got %v, want 3", v) + } + v, _ = FloatAt(got, 0, 1) + if v != 2 { + t.Errorf("pad reflect [1]: got %v, want 2", v) + } + // Original index 0 = 1 should land at the position equal to the pre-pad. + v, _ = FloatAt(got, 0, 2) + if v != 1 { + t.Errorf("pad reflect [2]: got %v, want 1", v) + } +} + +func TestPadReplicate(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + got, err := Pad(a, []int{2, 1, 0, 0}, "replicate", 0) + if err != nil { + t.Fatal(err) + } + // Shape (2, 5). At (0, 0) replicate original (0, 0) = 1. + v, _ := FloatAt(got, 0, 0) + if v != 1 { + t.Errorf("pad replicate [0,0]: got %v, want 1", v) + } + v, _ = FloatAt(got, 0, 1) + if v != 1 { + t.Errorf("pad replicate [0,1]: got %v, want 1", v) + } + // Post-pad replicates last column (index 1) for the last 1 column. + v, _ = FloatAt(got, 1, 4) + if v != 4 { + t.Errorf("pad replicate [1,4]: got %v, want 4", v) + } +} + +func TestPadCircular(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + got, err := Pad(a, []int{2, 1, 0, 0}, "circular", 0) + if err != nil { + t.Fatal(err) + } + // Shape (2, 5). Pad left=2, right=1 on dim 1. + // Row 0 = [1, 2]. Wrapped with pre=2 -> [..., 1, 2, 1, 2, 1] then trim to 5: + // index 0 = wrap(0-2=-2) = 0 -> 1 + // index 1 = wrap(0-1=-1) = 1 -> 2 + // index 2 = 1 + // index 3 = 2 + // index 4 = wrap(0+2=2) mod 2 = 0 -> 1 + v, _ := FloatAt(got, 0, 0) + if v != 1 { + t.Errorf("pad circular [0,0]: got %v, want 1", v) + } + v, _ = FloatAt(got, 0, 1) + if v != 2 { + t.Errorf("pad circular [0,1]: got %v, want 2", v) + } + v, _ = FloatAt(got, 0, 4) + if v != 1 { + t.Errorf("pad circular [0,4]: got %v, want 1", v) + } +} + +func TestGather(t *testing.T) { + a, _ := FromFloats([]float64{10, 20, 30, 40, 50, 60}, 2, 3) + idx := mustFromInts(t, []int64{2, 0}, 2, 1) + got, err := Gather(a, 1, idx) + if err != nil { + t.Fatal(err) + } + // Expect [a[0,2], a[1,0]] = [30, 40]. + v0, _ := FloatAt(got, 0, 0) + v1, _ := FloatAt(got, 1, 0) + if v0 != 30 || v1 != 40 { + t.Errorf("gather: got [%v, %v], want [30, 40]", v0, v1) + } + // Out-of-range index. + badIdx := mustFromInts(t, []int64{99, 0}, 2, 1) + if _, err := Gather(a, 1, badIdx); err == nil { + t.Error("gather: expected error for out-of-range index") + } + // Wrong dtype. + badDtype := mustFromFloats(t, []float64{1, 1}, 2, 1) + if _, err := Gather(a, 1, badDtype); err == nil { + t.Error("gather: expected error for non-int index") + } +} + +func TestScatter(t *testing.T) { + a, _ := FromFloats([]float64{10, 20, 30, 40, 50, 60}, 2, 3) + idx := mustFromInts(t, []int64{2, 0}, 2, 1) + src := mustFromFloats(t, []float64{300, 400}, 2, 1) + got, err := Scatter(a, 1, idx, src) + if err != nil { + t.Fatal(err) + } + // Expect [10, 20, 300, 400, 50, 60]. + for i, want := range []float64{10, 20, 300, 400, 50, 60} { + v, _ := FloatAt(got, i/3, i%3) + if v != want { + t.Errorf("scatter [%d]: got %v, want %v", i, v, want) + } + } + // Index/src shape mismatch. + badIdx := mustFromInts(t, []int64{0, 0, 0}, 3, 1) + if _, err := Scatter(a, 0, badIdx, idx); err == nil { + t.Error("scatter: expected error for shape mismatch") + } +} + +func TestNonzero(t *testing.T) { + a, _ := FromFloats([]float64{0, 1, 0, 2, 3, 0}, 2, 3) + got, err := Nonzero(a) + if err != nil { + t.Fatal(err) + } + if len(got) != 2 { + t.Fatalf("nonzero dims: got %d, want 2", len(got)) + } + // Flat layout [0, 1, 0, 2, 3, 0] -> row 0 [0,1,0], row 1 [2,3,0]. + // Non-zero positions: flat 1 = (0,1), flat 3 = (1,0), flat 4 = (1,1). + want := [][2]int{{0, 1}, {1, 0}, {1, 1}} + for i, w := range want { + if got[0][i] != w[0] || got[1][i] != w[1] { + t.Errorf("nonzero [%d]: got (%d, %d), want %v", i, got[0][i], got[1][i], w) + } + } + // Complex rejected. + c, _ := FromComplexes([]complex128{1, 0}, 2) + if _, err := Nonzero(c); err == nil { + t.Error("nonzero: expected error for complex input") + } +} + +func TestTake(t *testing.T) { + a := mustFromFloats(t, []float64{10, 20, 30, 40, 50}, 5) + idx := mustFromInts(t, []int64{4, 0, 2}, 3) + got, err := Take(a, idx) + if err != nil { + t.Fatal(err) + } + for i, want := range []float64{50, 10, 30} { + v, _ := FloatAt(got, i) + if v != want { + t.Errorf("take [%d]: got %v, want %v", i, v, want) + } + } + // Out-of-range. + badIdx := mustFromInts(t, []int64{99}, 1) + if _, err := Take(a, badIdx); err == nil { + t.Error("take: expected error for out-of-range index") + } + // Non-1-D indices. + badShape := mustFromInts(t, []int64{0, 0}, 2, 1) + if _, err := Take(a, badShape); err == nil { + t.Error("take: expected error for non-1-D indices") + } +} + +func TestCumSum(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + got, err := CumSum(a, 1) + if err != nil { + t.Fatal(err) + } + want, _ := FromFloats([]float64{1, 3, 6, 4, 9, 15}, 2, 3) + for i := range 6 { + v, _ := FloatAt(got, i/3, i%3) + w, _ := FloatAt(want, i/3, i%3) + if v != w { + t.Errorf("cumsum [%d]: got %v, want %v", i, v, w) + } + } + // Out-of-range dim. + if _, err := CumSum(a, 5); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Errorf("cumsum: expected out-of-range error, got %v", err) + } +} + +func TestCumProd(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + got, err := CumProd(a, 0) + if err != nil { + t.Fatal(err) + } + want, _ := FromFloats([]float64{1, 2, 3, 4, 10, 18}, 2, 3) + for i := range 6 { + v, _ := FloatAt(got, i/3, i%3) + w, _ := FloatAt(want, i/3, i%3) + if v != w { + t.Errorf("cumprod [%d]: got %v, want %v", i, v, w) + } + } +} + +func TestProd(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + got, err := Prod(a, 1, false) + if err != nil { + t.Fatal(err) + } + // Prod across dim 1: row 0 = [1,2,3] -> 6, row 1 = [4,5,6] -> 120. + for i, w := range []float64{6, 120} { + v, _ := FloatAt(got, i) + if v != w { + t.Errorf("prod [%d]: got %v, want %v", i, v, w) + } + } + // With keepDim. + got, err = Prod(a, 1, true) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 2 || got.Shape()[1] != 1 { + t.Errorf("prod keepDim: shape = %v", got.Shape()) + } +} + +func TestNorm(t *testing.T) { + a := mustFromFloats(t, []float64{3, 4}, 2) + got, err := Norm(a, 2, 0, false) + if err != nil { + t.Fatal(err) + } + if math.Abs(got.RawFloats()[0]-5) > 1e-9 { + t.Errorf("L2 norm [3,4]: got %v, want 5", got.RawFloats()[0]) + } + // L1. + got, _ = Norm(a, 1, 0, false) + if got.RawFloats()[0] != 7 { + t.Errorf("L1 norm: got %v, want 7", got.RawFloats()[0]) + } + // L-inf. + got, _ = Norm(a, math.Inf(1), 0, false) + if got.RawFloats()[0] != 4 { + t.Errorf("L-inf norm: got %v, want 4", got.RawFloats()[0]) + } + // Negative p. + if _, err := Norm(a, -1, 0, false); err == nil { + t.Error("norm: expected error for negative p") + } +} + +func TestTrace(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3) + got, err := Trace(a) + if err != nil { + t.Fatal(err) + } + if got != 15 { + t.Errorf("trace: got %v, want 15", got) + } + // Non-square. + bad, _ := FromFloats([]float64{1, 2, 3, 4}, 2, 2) + // 2x2 is square; try non-square. + bad2, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if _, err := Trace(bad2); err == nil { + t.Error("trace: expected error for non-square") + } + // Just to silence "unused" for bad. + _ = bad +} + +func TestDiagonal(t *testing.T) { + a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3) + main, err := Diagonal(a, 0) + if err != nil { + t.Fatal(err) + } + for i, w := range []float64{1, 5, 9} { + v, _ := FloatAt(main, i) + if v != w { + t.Errorf("diag main [%d]: got %v, want %v", i, v, w) + } + } + // Super-diagonal offset 1. + sup, err := Diagonal(a, 1) + if err != nil { + t.Fatal(err) + } + for i, w := range []float64{2, 6} { + v, _ := FloatAt(sup, i) + if v != w { + t.Errorf("diag +1 [%d]: got %v, want %v", i, v, w) + } + } + // Sub-diagonal offset -1. + sub, err := Diagonal(a, -1) + if err != nil { + t.Fatal(err) + } + for i, w := range []float64{4, 8} { + v, _ := FloatAt(sub, i) + if v != w { + t.Errorf("diag -1 [%d]: got %v, want %v", i, v, w) + } + } +} + +func TestKron(t *testing.T) { + a, _ := FromFloats([]float64{1, 2}, 1, 2) + b, _ := FromFloats([]float64{0, 5, 6, 7}, 2, 2) + got, err := Kron(a, b) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 2 || got.Shape()[1] != 4 { + t.Errorf("kron: shape = %v", got.Shape()) + } + // Expect: a = [[1, 2]], b = [[0,5],[6,7]] + // Result = [[0,5,0,10],[6,7,12,14]]. + for i, w := range []float64{0, 5, 0, 10, 6, 7, 12, 14} { + v, _ := FloatAt(got, i/4, i%4) + if v != w { + t.Errorf("kron [%d]: got %v, want %v", i, v, w) + } + } + // Int input: exercises setFromValue with int dtype. + aInt, _ := FromInts([]int64{1, 2}, 1, 2) + gotInt, err := Kron(aInt, b) + if err != nil { + t.Fatal(err) + } + v0, _ := IntAt(gotInt, 0, 0) + if v0 != 0 { + t.Errorf("kron int [0,0]: got %v, want 0", v0) + } +} + +// TestPadReflectRejectsOversizedPads pins the reflect guard: a pad of +// the full axis length has no source to mirror, and the single fold +// used to emit a negative offset into the payload (panic). +func TestPadReflectRejectsOversizedPads(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 3) + if _, err := Pad(a, []int{6, 0}, "reflect", 0); err == nil { + t.Error("Pad reflect pad=6 on length 3: expected error") + } + // A pad of n-1 is still mirrorable and must keep working. + got, err := Pad(a, []int{2, 0}, "reflect", 0) + if err != nil { + t.Fatalf("Pad reflect pad=2: %v", err) + } + want := mustFromFloats(t, []float64{3, 2, 1, 2, 3}, 5) + if !Equal(want, got) { + t.Errorf("Pad reflect pad=2: %s", got) + } +} + +// TestScanFloat32CarryNarrows pins the float32 scan's carry rule: the +// running value narrows to float32 at every step, exactly as the value +// the walk stored fed the next one, so a term below the current +// precision is absorbed before the next term lands on it. +func TestScanFloat32CarryNarrows(t *testing.T) { + // 1 + 1e-9 rounds back to 1, so the -1 lands on a clean zero. + sums := mustFromFloat32s(t, []float32{1, 1e-9, -1}, 3) + cs, err := CumSum(sums, 0) + if err != nil { + t.Fatalf("CumSum float32: %v", err) + } + for i, want := range []float32{1, 1, 0} { + if got := cs.RawFloat32s()[i]; got != want { + t.Errorf("CumSum float32 [%d] = %v, want %v", i, got, want) + } + } + // A longer chain keeps the rule: the absorbed terms never + // accumulate into the running sum. + chain := mustFromFloat32s(t, []float32{1, 1e-9, 1e-9, 1e-9, -1}, 5) + cs2, err := CumSum(chain, 0) + if err != nil { + t.Fatalf("CumSum float32 chain: %v", err) + } + for i, want := range []float32{1, 1, 1, 1, 0} { + if got := cs2.RawFloat32s()[i]; got != want { + t.Errorf("CumSum float32 chain [%d] = %v, want %v", i, got, want) + } + } + // The product twin narrows the same way, so a product wider than + // float32 is rounded before the next factor multiplies it. + vals := []float32{1 + 1.0/4096, 1 + 1.0/4096, 1 + 1.0/4096, 1 + 1.0/4096} + prod := mustFromFloat32s(t, vals, len(vals)) + cp, err := CumProd(prod, 0) + if err != nil { + t.Fatalf("CumProd float32: %v", err) + } + var want []float32 + acc := 1.0 + for i, v := range vals { + if i == 0 { + acc = float64(v) + } else { + acc = float64(float32(acc)) * float64(v) + } + want = append(want, float32(acc)) + } + for i := range want { + if got := cp.RawFloat32s()[i]; got != want[i] { + t.Errorf("CumProd float32 [%d] = %v, want %v", i, got, want[i]) + } + } +} diff --git a/internal/core/runtime.go b/internal/core/runtime.go new file mode 100644 index 0000000..1928915 --- /dev/null +++ b/internal/core/runtime.go @@ -0,0 +1,59 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/engine" + +// CPU policy. The heavy loops of the library, matrix products, +// convolutions, element-wise maps, axis reductions, scans and the FFT, +// run in parallel across the worker count chosen here. The default is +// NumCPU, which is right for almost every machine; call SetNumCPU once +// at startup to override (e.g. for a container with a CPU quota, or a +// small embedded box). + +// SetNumCPU sets the number of goroutines the parallel kernels may use +// and returns the previous value. Values below 1 reset to NumCPU. It +// is safe to call at any time; running kernels finish with their old +// worker count. +func SetNumCPU(n int) int { return engine.SetNumWorkers(n) } + +// NumWorkers returns the current worker count. +func NumWorkers() int { return engine.NumWorkers() } + +// workersFor bounds the worker count for a workload of n items. +func workersFor(n int) int { return engine.WorkersFor(n) } + +// parallel splits the [0, n) range into chunks across workers. +func parallel(n int, fn func(start, end int)) { engine.Parallel(n, fn) } + +// parallelMin splits the [0, n) range across workers only while each +// worker keeps at least minPerWorker elements; a workload below that +// floor runs on the calling goroutine. +func parallelMin(n, minPerWorker int, fn func(start, end int)) { + engine.ParallelMin(n, minPerWorker, fn) +} + +// The spawn floors below are constraints on when a parallel kernel may +// spawn at all: a worker must carry enough elements for its own +// creation and scheduling to stay a small fraction of the work it runs, +// and below that floor the whole map is faster on the calling +// goroutine. The floor scales with the per-element cost of the kernel, +// so each kernel family carries its own measured value; the sweep in +// math_bench_extra_test.go guards the sizes around both crossovers. + +// elementwiseMinPerWorker is the per-worker chunk floor for the +// element-independent arithmetic maps (Add, Sub, Mul, Div and the +// scalar maps). Their per-element work is a couple of machine +// instructions, so a spawned worker needs about a thousand elements +// before the chunk outweighs its own start-up; below roughly 32k +// elements the maps measure faster serial, above that parallel. +const elementwiseMinPerWorker = 1024 + +// mathFuncMinPerWorker is the per-worker chunk floor for the +// per-element math maps (realFunc, roundFunc and the specialised +// Sqrt). Their per-element work costs tens of cycles, so a worker +// amortises its start-up at far smaller chunks than the arithmetic +// maps; a floor in the thousands would keep workloads of a few thousand +// elements needlessly serial. +const mathFuncMinPerWorker = 64 diff --git a/internal/core/runtime_test.go b/internal/core/runtime_test.go new file mode 100644 index 0000000..03a3189 --- /dev/null +++ b/internal/core/runtime_test.go @@ -0,0 +1,89 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "runtime" + "testing" +) + +func TestSetNumCPU(t *testing.T) { + // SetNumCPU(0) resets to runtime.NumCPU(): the invariant that must + // hold regardless of what previous tests set. Test against that, + // not against a captured "previous" value (tests run in any order). + ncpu := runtime.NumCPU() + if NumWorkers() < 1 { + t.Fatalf("NumWorkers: %d", NumWorkers()) + } + got := SetNumCPU(4) + if got < 1 { + t.Errorf("SetNumCPU return: %d, want ≥ 1", got) + } + if NumWorkers() != 4 { + t.Errorf("NumWorkers after SetNumCPU(4): %d", NumWorkers()) + } + // workersFor bounds against the explicit worker count, independent + // of the host's CPU count (the CI runner has one core). + if w := workersFor(2); w != 2 { + t.Errorf("workersFor(2) with 4 workers: %d, want 2", w) + } + if w := workersFor(1); w != 1 { + t.Errorf("workersFor(1): %d, want 1", w) + } + if w := workersFor(10); w != 4 { + t.Errorf("workersFor(10) with 4 workers: %d, want 4", w) + } + SetNumCPU(0) // resets to NumCPU + if NumWorkers() != ncpu { + t.Errorf("NumWorkers after reset: %d, want %d", NumWorkers(), ncpu) + } + // Universal invariants after the reset: at least one worker, never + // more than the item count. + if w := workersFor(0); w != 1 { + t.Errorf("workersFor(0): %d, want 1 (floor of one worker)", w) + } + if w := workersFor(1); w != 1 { + t.Errorf("workersFor(1) after reset: %d, want 1", w) + } +} + +// TestParallelCoverage forces the parallel branch of `parallel` to run +// even on a single-core CI runner, where workersFor(n) would otherwise +// collapse to 1 and skip the goroutine-spawning path entirely, which +// silently drops the measured coverage below the gate. It also covers +// the parallel-only merge branch of the axis reductions (reduceAxis), +// whose private-scratch/merge code has no serial equivalent. +func TestParallelCoverage(t *testing.T) { + prev := SetNumCPU(2) + defer SetNumCPU(prev) + a, err := FromFloats(make([]float64, 1<<16), 1<<16) + if err != nil { + t.Fatal(err) + } + b, err := FromFloats(make([]float64, 1<<16), 1<<16) + if err != nil { + t.Fatal(err) + } + // With numWorkers=2 and 65536 items the parallel branch spawns + // goroutines even on a one-core host. + sum, err := Add(a, b) + if err != nil { + t.Fatal(err) + } + if sum.Len() != a.Len() { + t.Fatalf("Add len: %d", sum.Len()) + } + // Axis reductions take the parallel merge path with 2 workers, + // covering the private-scratch and merge code in reduceAxis. + mat, err := Reshape(a, 256, 256) + if err != nil { + t.Fatal(err) + } + if _, err := SumAxis(mat, 1); err != nil { + t.Fatal(err) + } + if _, err := MaxAxis(mat, 0); err != nil { + t.Fatal(err) + } +} diff --git a/internal/core/scalar_clamp_pow_bench_test.go b/internal/core/scalar_clamp_pow_bench_test.go new file mode 100644 index 0000000..d087445 --- /dev/null +++ b/internal/core/scalar_clamp_pow_bench_test.go @@ -0,0 +1,71 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// The scalar clamps, the int-exponent power and the integer quotient at +// the million-element size the element-wise family reports at, so the +// walks behind them are judged at the size where a serial pass and a +// fanned-out pass part ways. + +// clampBenchFloats fills an n-element float64 slice with a +// deterministic mix straddling the clamp window. +func clampBenchFloats(n int) []float64 { + v := make([]float64, n) + for i := range v { + v[i] = float64(i%2001) - 1000.5 + } + return v +} + +// clampBenchInts fills an n-element int64 slice with a deterministic mix +// straddling the clamp window. +func clampBenchInts(n int) []int64 { + v := make([]int64, n) + for i := range v { + v[i] = int64(i%2001) - 1000 + } + return v +} + +func BenchmarkClipI1M(b *testing.B) { + a, _ := FromFloats(clampBenchFloats(1<<20), 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := ClipI(a, -100, 100); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkClipF1M(b *testing.B) { + a, _ := FromFloats(clampBenchFloats(1<<20), 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := ClipF(a, -100, 100); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkPowIInt1M(b *testing.B) { + a, _ := FromInts(clampBenchInts(1<<20), 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := PowI(a, 3); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkQuoI1M(b *testing.B) { + a, _ := FromInts(clampBenchInts(1<<20), 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := QuoI(a, 7); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/core/scan_compensation_test.go b/internal/core/scan_compensation_test.go new file mode 100644 index 0000000..fd6c151 --- /dev/null +++ b/internal/core/scan_compensation_test.go @@ -0,0 +1,249 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "fmt" + "math" + "math/big" + mrand "math/rand/v2" + "testing" +) + +// The compensated-scan evidence file: the sequential float64 CumSum +// chain against the Neumaier-compensated candidate, both measured +// against an exact big.Float referent (256 bits of precision, every +// widening exact). The twins below carry the two walks side by side so +// one binary answers both the accuracy question and the cost question. + +// scanLegacyFloats is the sequential chain the Float64 scan arm keeps: +// out[i] = out[i-1] + src[i], one rounding per step. +func scanLegacyFloats(src, dst []float64) { + acc := src[0] + dst[0] = acc + for i := 1; i < len(src); i++ { + acc += src[i] + dst[i] = acc + } +} + +// scanNeumaierFloats is the compensated candidate: a running high part +// and a correction, the output the corrected running value. A non-finite +// partial freezes the correction, so an overflow sticks to infinity and +// a NaN poisons the tail exactly as the plain chain does. +func scanNeumaierFloats(src, dst []float64) { + sum, comp := src[0], 0.0 + dst[0] = sum + for i := 1; i < len(src); i++ { + x := src[i] + t := sum + x + if !math.IsInf(t, 0) && t == t { + if math.Abs(sum) >= math.Abs(x) { + comp += (sum - t) + x + } else { + comp += (x - t) + sum + } + } + sum = t + dst[i] = sum + comp + } +} + +// scanPrefixExact folds the prefixes in big.Float at the given +// precision; every input is a float64, so the widening is exact. +func scanPrefixExact(src []float64, prec uint) []*big.Float { + out := make([]*big.Float, len(src)) + acc := new(big.Float).SetPrec(prec) + for i, v := range src { + acc = new(big.Float).SetPrec(prec).Add(acc, new(big.Float).SetPrec(prec).SetFloat64(v)) + out[i] = acc + } + return out +} + +// scanError reports the largest deviation of the computed prefixes from +// the exact referent, both absolute and relative to Σ|src|. +func scanError(got []float64, exact []*big.Float) (abs, rel float64) { + scale := 0.0 + for _, v := range got { + scale += math.Abs(v) + } + for i := range got { + ref, _ := exact[i].Float64() + if e := math.Abs(got[i] - ref); e > abs { + abs = e + } + } + if scale > 0 { + rel = abs / scale + } + return abs, rel +} + +// scanAlternatingData builds the cancellation pattern [1, 1e100, 1, +// -1e100] repeated: the plain chain loses every small addend it meets, +// the compensated walk carries them. +func scanAlternatingData(n int) []float64 { + src := make([]float64, n) + for i := 0; i < n; i += 4 { + src[i], src[i+1], src[i+2], src[i+3] = 1, 1e100, 1, -1e100 + } + return src +} + +// scanSpreadData builds a deterministic log-uniform random sample whose +// magnitudes span fifty orders, the shape of data that ages a plain +// chain fastest. +func scanSpreadData(n int, seed uint64) []float64 { + rng := mrand.New(mrand.NewPCG(seed, seed)) + src := make([]float64, n) + for i := range src { + mag := math.Pow(10, -25+50*rng.Float64()) + if rng.Float64() < 0.5 { + mag = -mag + } + src[i] = mag + } + return src +} + +// TestScanCompensationAccuracy measures both walks against the exact +// referent and pins the compensated result's correctness on the +// cancellation pattern, where the plain chain reports a bare zero. +func TestScanCompensationAccuracy(t *testing.T) { + // The referent runs at 512 bits: the alternating pattern's values + // reach 1e100, whose mantissa needs 333 bits, and every prefix must + // stay exactly representable for the referent to judge ulp errors. + const prec = 512 + cases := []struct { + name string + src []float64 + }{ + {"alternating 2^20", scanAlternatingData(1 << 20)}, + {"spread 2^20", scanSpreadData(1<<20, 0xC0FFEE)}, + {"spread 2^16", scanSpreadData(1<<16, 7)}, + {"unit 2^20", func() []float64 { + rng := mrand.New(mrand.NewPCG(42, 42)) + src := make([]float64, 1<<20) + for i := range src { + src[i] = rng.Float64() + } + return src + }()}, + } + for _, tc := range cases { + exact := scanPrefixExact(tc.src, prec) + oldGot := make([]float64, len(tc.src)) + newGot := make([]float64, len(tc.src)) + scanLegacyFloats(tc.src, oldGot) + scanNeumaierFloats(tc.src, newGot) + oldAbs, oldRel := scanError(oldGot, exact) + newAbs, newRel := scanError(newGot, exact) + t.Logf("%s: legacy abs=%.3e rel=%.3e | neumaier abs=%.3e rel=%.3e | improvement %.0fx (abs)", + tc.name, oldAbs, oldRel, newAbs, newRel, oldAbs/newAbs) + // A regression guard: the compensated walk is never the worse of + // the two on these data. + if newAbs > oldAbs { + t.Errorf("%s: compensated error %.3e exceeds the chain's %.3e", tc.name, newAbs, oldAbs) + } + } + // The cancellation pin: every complete pattern's prefix returns to + // 2·k exactly, which the plain chain reports as zero. This pin fails + // against the uncompenated walk by construction. + src := []float64{1, 1e100, 1, -1e100, 1, 1e100, 1, -1e100} + exact := scanPrefixExact(src, prec) + got := make([]float64, len(src)) + scanNeumaierFloats(src, got) + for i := range src { + want, _ := exact[i].Float64() + if math.Abs(got[i]-want) > 1e-9*math.Max(1, math.Abs(want)) { + t.Errorf("compensated prefix %d: got %v, want %v", i, got[i], want) + } + } + if last, _ := exact[len(src)-1].Float64(); last != 4 { + t.Fatalf("referent check: final exact prefix %v, want 4", last) + } + // The plain chain's final prefix is 0 here: the document of what the + // compensation buys. + legacy := make([]float64, len(src)) + scanLegacyFloats(src, legacy) + if legacy[len(src)-1] == 4 { + t.Log("note: the plain chain happened to land exactly on this sample") + } +} + +// TestScanNaNInfSemantics pins the non-finite contract of the +// compensated walk: an overflow sticks to infinity and a NaN poisons +// the tail, exactly as the plain chain reports them. +func TestScanNaNInfSemantics(t *testing.T) { + cases := [][]float64{ + {math.Inf(1), 1, 2, 3}, + {1, 2, math.Inf(-1), 4}, + {1, math.NaN(), 3, 4}, + {math.MaxFloat64, math.MaxFloat64, 1, 2}, + {1, -1, math.MaxFloat64, math.MaxFloat64}, + } + for _, src := range cases { + oldGot := make([]float64, len(src)) + newGot := make([]float64, len(src)) + scanLegacyFloats(src, oldGot) + scanNeumaierFloats(src, newGot) + for i := range src { + wantClass, gotClass := math.Signbit(oldGot[i]), math.Signbit(newGot[i]) + finiteW, finiteG := oldGot[i]-oldGot[i] == 0, newGot[i]-newGot[i] == 0 + if math.IsInf(oldGot[i], 0) != math.IsInf(newGot[i], 0) || + math.IsNaN(oldGot[i]) != math.IsNaN(newGot[i]) || + (finiteW != finiteG) || (finiteW && wantClass != gotClass) { + t.Errorf("src %v prefix %d: legacy %v, compensated %v", src, i, oldGot[i], newGot[i]) + } + } + } +} + +// TestCumSumCompensationProduction pins the shipped entry point on the +// cancellation pattern: the final prefix of two full patterns is 4, +// which the uncompensated chain reports as 0 (the whole sum cancels and +// every small addend is lost on the way). This pin fails against the +// plain-chain walk by construction. +func TestCumSumCompensationProduction(t *testing.T) { + a, err := FromFloats([]float64{1, 1e100, 1, -1e100, 1, 1e100, 1, -1e100}, 8) + if err != nil { + t.Fatal(err) + } + got, err := CumSum(a, 0) + if err != nil { + t.Fatal(err) + } + if last := got.FloatAt(7); math.Abs(last-4) > 1e-9 { + t.Errorf("compensated CumSum final prefix: %v, want 4", last) + } + // The intermediate prefixes carry their compensation too: after the + // first pattern the running value is 2, not the chain's 0. + if mid := got.FloatAt(3); math.Abs(mid-2) > 1e-9 { + t.Errorf("compensated CumSum prefix 3: %v, want 2", mid) + } +} + +// BenchmarkScanCompensated runs the two walks side by side over one +// line of the given length; the sub-benchmarks alternate so one binary +// answers the cost question. +func BenchmarkScanCompensated(b *testing.B) { + for _, n := range []int{1 << 10, 1 << 16, 1 << 20} { + src := scanSpreadData(n, 3) + b.Run(fmt.Sprintf("legacy/n=%d", n), func(b *testing.B) { + dst := make([]float64, n) + for b.Loop() { + scanLegacyFloats(src, dst) + } + b.SetBytes(int64(n) * 8) + }) + b.Run(fmt.Sprintf("neumaier/n=%d", n), func(b *testing.B) { + dst := make([]float64, n) + for b.Loop() { + scanNeumaierFloats(src, dst) + } + b.SetBytes(int64(n) * 8) + }) + } +} diff --git a/internal/core/shape_overflow_pin_test.go b/internal/core/shape_overflow_pin_test.go new file mode 100644 index 0000000..ae9e147 --- /dev/null +++ b/internal/core/shape_overflow_pin_test.go @@ -0,0 +1,66 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +// The element-count guard: a shape whose product wraps must be a loud +// refusal, never a tiny allocation under a huge declared shape. + +// TestShapeOverflowRefused pins the element-count guard: a shape +// whose product wraps past the int range must be refused, because a +// wrapped total pairs a huge declared shape with a tiny allocation and +// every later index write then lands outside the payload. A hostile +// HDF5 dataspace of 2^32 x 2^32, whose product wraps to zero, reaches +// the same guard. +func TestShapeOverflowRefused(t *testing.T) { + const big = 1 << 32 + for _, tc := range []struct { + name string + shape []int + fits bool + }{ + {"wrapping product", []int{big, big}, false}, + {"wrapping via three dims", []int{big, big, 2}, false}, + {"boundary below overflow", []int{1 << 31, 1}, true}, + {"zero factor absorbs", []int{big, 0}, true}, + } { + t.Run(tc.name, func(t *testing.T) { + total, _, err := checkedDims(tc.shape) + if tc.fits { + if err != nil { + t.Fatalf("checkedDims(%v) = %v, want success", tc.shape, err) + } + if total < 0 { + t.Fatalf("checkedDims(%v) = %d, want a non-negative count", tc.shape, total) + } + return + } + if err == nil { + t.Fatalf("checkedDims(%v) accepted a wrapping shape (total %d)", tc.shape, total) + } + if !strings.Contains(err.Error(), "more elements than fit") { + t.Fatalf("checkedDims(%v) = %v, want the overflow message", tc.shape, err) + } + }) + } + + // The validating constructor must refuse the same shape rather than + // allocate an empty payload for it. + if _, err := Zeros(Float, big, big); err == nil { + t.Fatal("Zeros(Float, 2^32, 2^32) accepted a wrapping shape") + } + if _, err := FromFloats(nil, big, big); err == nil { + t.Fatal("FromFloats(nil, 2^32, 2^32) accepted a wrapping shape") + } + // The largest representable single-dimension shape stays legal: the + // guard must not shrink the usable range for shapes that do not wrap. + if _, _, err := checkedDims([]int{math.MaxInt}); err != nil { + t.Fatalf("checkedDims(MaxInt) = %v, want success", err) + } +} diff --git a/internal/core/smallops_test.go b/internal/core/smallops_test.go new file mode 100644 index 0000000..8f1fd65 --- /dev/null +++ b/internal/core/smallops_test.go @@ -0,0 +1,198 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// TestInterpolateMonotoneShape checks the two promises of the +// Fritsch-Carlson tangents: the curve passes through every knot, and +// interpolating monotone data never overshoots the data's range. +func TestInterpolateMonotoneShape(t *testing.T) { + xv := []float64{0, 1, 2, 3, 4, 5, 6, 7, 8, 9} + yv := []float64{0, 0.1, 1.9, 2.0, 2.1, 2.2, 3.9, 4.0, 4.1, 8} + xs := mustFloats(t, xv, len(xv)) + ys := mustFloats(t, yv, len(yv)) + dense := make([]float64, 0, 801) + for i := range 801 { + dense = append(dense, -1+float64(i)*10/800) + } + q := mustFloats(t, dense, len(dense)) + out, err := InterpolateMonotone(xs, ys, q) + if err != nil { + t.Fatalf("InterpolateMonotone: %v", err) + } + for k := range xv { + if got := out.FloatAt((k + 1) * 80); math.Abs(got-yv[k]) > 1e-12 { + t.Fatalf("knot %d: %.12g, want %.12g", k, got, yv[k]) + } + } + for i := range out.Len() { + if v := out.FloatAt(i); v < -1e-12 || v > 8+1e-12 { + t.Fatalf("overshoot at %g: %.12g outside the data range", dense[i], v) + } + } +} + +// TestInterpolateMonotoneBasics checks the two-knot limit (a +// straight line), the flat clamp outside, and the refusal of +// repeated knots. +func TestInterpolateMonotoneBasics(t *testing.T) { + xs := mustFloats(t, []float64{0, 1}, 2) + ys := mustFloats(t, []float64{10, 20}, 2) + q := mustFloats(t, []float64{0.25, -3, 4}, 3) + out, err := InterpolateMonotone(xs, ys, q) + if err != nil { + t.Fatalf("InterpolateMonotone: %v", err) + } + if got := out.FloatAt(0); math.Abs(got-12.5) > 1e-12 { + t.Fatalf("linear read = %.12g, want 12.5", got) + } + if got := out.FloatAt(1); got != 10 { + t.Fatalf("clamp below = %.12g, want 10", got) + } + if got := out.FloatAt(2); got != 20 { + t.Fatalf("clamp above = %.12g, want 20", got) + } + bad := mustFloats(t, []float64{1, 1, 2}, 3) + if _, err := InterpolateMonotone(bad, ys, q); err == nil { + t.Fatal("repeated knots accepted") + } +} + +// TestInterpolateGrid checks multilinear reads: exact at the nodes, +// the bilinear average at a cell centre, clamped outside, and the +// argument guards. +func TestInterpolateGrid(t *testing.T) { + grid := mustFloats(t, []float64{ + 0, 10, + 20, 30, + }, 2, 2) + origins := []float64{0, 0} + steps := []float64{1, 0.5} + // Node reads land exactly; the cell centre averages the corners. + queries := mustFloats(t, []float64{ + 0, 0, + 1, 0.5, + 0.5, 0.25, + -9, 9, + }, 4, 2) + out, err := InterpolateGrid(grid, origins, steps, queries) + if err != nil { + t.Fatalf("InterpolateGrid: %v", err) + } + if got := out.FloatAt(0); got != 0 { + t.Fatalf("node (0,0) = %.12g, want 0", got) + } + if got := out.FloatAt(1); got != 30 { + t.Fatalf("node (1,0.5) = %.12g, want 30", got) + } + if got := out.FloatAt(2); math.Abs(got-15) > 1e-12 { + t.Fatalf("cell centre = %.12g, want 15", got) + } + if got := out.FloatAt(3); got != 10 { + t.Fatalf("clamp = %.12g, want 10", got) + } + if _, err := InterpolateGrid(grid, origins, []float64{1, -1}, queries); err == nil { + t.Fatal("negative step accepted") + } + badQueries := mustFloats(t, []float64{0, 0, 0}, 3, 1) + if _, err := InterpolateGrid(grid, origins, steps, badQueries); err == nil { + t.Fatal("wrong query columns accepted") + } +} + +// TestDiff checks the finite differences on a known series and along +// a chosen axis of a matrix, plus the argument guards. +func TestDiff(t *testing.T) { + a := mustFloats(t, []float64{1, 4, 9, 16}, 4) + d1, err := Diff(a, 1, 0) + if err != nil { + t.Fatalf("Diff: %v", err) + } + want := []float64{3, 5, 7} + for i := range want { + if d1.FloatAt(i) != want[i] { + t.Fatalf("first difference %d = %g, want %g", i, d1.FloatAt(i), want[i]) + } + } + d2, err := Diff(a, 2, 0) + if err != nil { + t.Fatalf("Diff order 2: %v", err) + } + if d2.FloatAt(0) != 2 || d2.FloatAt(1) != 2 { + t.Fatalf("second difference = [%g, %g], want [2, 2]", d2.FloatAt(0), d2.FloatAt(1)) + } + m := mustFloats(t, []float64{1, 4, 2, 8}, 2, 2) + dm, err := Diff(m, 1, 0) + if err != nil { + t.Fatalf("Diff axis 0: %v", err) + } + if dm.Shape()[0] != 1 || dm.Shape()[1] != 2 { + t.Fatalf("shape %v, want [1 2]", dm.Shape()) + } + if dm.FloatAt(1) != 4 { + t.Fatalf("axis-0 difference column 1 = %g, want 4", dm.FloatAt(1)) + } + if _, err := Diff(a, 4, 0); err == nil { + t.Fatal("order at the axis length accepted") + } + if _, err := Diff(m, 1, 5); err == nil { + t.Fatal("axis past the rank accepted") + } +} + +// TestDiffNonLeadingAxis pins the run walk of Diff along an axis that is +// neither leading nor trailing: several outer positions (head) each hold +// a run of trailing elements (tail), so a run's destination is not its +// own start and the source offsets step by the line length. +func TestDiffNonLeadingAxis(t *testing.T) { + vals := make([]int64, 24) + for i := range vals { + vals[i] = int64(i) + } + src := mustFromInts(t, vals, 2, 3, 4) + + // Axis 1: head 2, line 3, tail 4. + got, err := Diff(src, 1, 1) + if err != nil { + t.Fatalf("Diff axis 1: %v", err) + } + if sh := got.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 2 || sh[2] != 4 { + t.Fatalf("Diff axis 1 shape: %v, want [2 2 4]", sh) + } + for i := range 2 { + for j := range 2 { + for k := range 4 { + hi, _ := IntAt(src, i, j+1, k) + lo, _ := IntAt(src, i, j, k) + if v, _ := IntAt(got, i, j, k); v != hi-lo { + t.Fatalf("Diff axis 1 [%d %d %d] = %d, want %d", i, j, k, v, hi-lo) + } + } + } + } + + // Axis 2: head 6, line 4, tail 1. + got2, err := Diff(src, 1, 2) + if err != nil { + t.Fatalf("Diff axis 2: %v", err) + } + if sh := got2.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 3 || sh[2] != 3 { + t.Fatalf("Diff axis 2 shape: %v, want [2 3 3]", sh) + } + for i := range 2 { + for j := range 3 { + for k := range 3 { + hi, _ := IntAt(src, i, j, k+1) + lo, _ := IntAt(src, i, j, k) + if v, _ := IntAt(got2, i, j, k); v != hi-lo { + t.Fatalf("Diff axis 2 [%d %d %d] = %d, want %d", i, j, k, v, hi-lo) + } + } + } + } +} diff --git a/internal/core/sort.go b/internal/core/sort.go new file mode 100644 index 0000000..6e9e2e1 --- /dev/null +++ b/internal/core/sort.go @@ -0,0 +1,642 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "cmp" + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Ordering. Sort returns an ascending copy, Reverse flips the +// element order of any dtype, ArgSort returns the stable index +// permutation that would sort. Floats sort with NaN at the end; complex +// arrays have no ordering and error. Descending order composes as +// Sort().Reverse(). + +// Sort returns an ascending copy; NaN elements land at the end. The +// two zeros compare equal, as they always have, and the copy carries +// +0.0 wherever the input held −0.0. A NaN's payload bits are not +// preserved either: the sorted NaNs are fresh quiet NaNs, since the +// radix keys carry only the finite values' order. +func Sort(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("Sort: complex arrays have no ordering") + } + if narrowRefused(a.dt) { + // The narrow radix digits carry no kernel; the + // refusal is loud and named, and Astype is the documented + // route. + return nil, errf("Sort: dtype %s is not supported; convert with Astype", a.dt) + } + out := &Array{shape: a.Shape(), dt: a.dt} + switch a.dt { + case Int: + ints, _, _, _, _ := a.cloneData() + // The sign-bit flip lets the digit passes compare in uint64; + // the unflip restores the original values afterwards. The key + // range is folded into the same walk, so the passes only cover + // the bytes the keys actually differ in. + keys := make([]uint64, len(ints)) + hi, lo := uint64(0), ^uint64(0) + for i, v := range ints { + k := uint64(v) ^ (1 << 63) + keys[i] = k + hi |= k + lo &= k + } + radixSortUint64(keys, hi^lo) + for i, k := range keys { + ints[i] = int64(k ^ (1 << 63)) + } + out.ints = ints + case Float16: + // Sort in float64 (exact conversions) and narrow back: NaN + // handling identical to the float64 path. The payload arrives + // through cloneData like its Int and Float32 siblings, so a + // strided view sorts its own elements, never the raw window. + _, halves, _, _, _ := a.cloneData() + vals := make([]float64, a.Len()) + for i := range vals { + vals[i] = HalfToFloat64(halves[i]) + } + sorted := sortFloatsRadix(vals) + out.halves = make([]uint16, len(sorted)) + for i, v := range sorted { + out.halves[i] = HalfFromFloat64(v) + } + case Float32: + // Sort in float64 (exact conversions) and round back: NaN + // handling identical to the float64 path. The payload + // arrives through cloneData like its Int and Float siblings, so + // a strided view sorts its own elements, never the raw window. + _, _, f32, _, _ := a.cloneData() + vals := make([]float64, a.Len()) + for i := range vals { + vals[i] = float64(f32[i]) + } + sorted := sortFloatsRadix(vals) + out.floats32 = make([]float32, len(sorted)) + for i, v := range sorted { + out.floats32[i] = float32(v) + } + default: + _, _, _, floats, _ := a.cloneData() + out.floats = sortFloatsRadix(floats) + } + return out, nil +} + +// sortFloatsRadix partitions NaNs out, sorts the finite values +// ascending through the order-preserving uint64 keys and appends the +// NaNs at the end. The input slice is reused as the result, exactly +// like the compaction the comparator sort it replaces used. Keys and +// their range are folded in the one compaction walk: the NaNs take no +// pass and the digit passes cover only the bytes the finite keys differ +// in. +func sortFloatsRadix(floats []float64) []float64 { + keys := make([]uint64, len(floats)) + finite := floats[:0] + m := 0 + hi, lo := uint64(0), ^uint64(0) + for _, v := range floats { + if v != v { + continue + } + k := floatSortKey(v) + finite = append(finite, v) + keys[m] = k + m++ + hi |= k + lo &= k + } + keys = keys[:m] + radixSortUint64(keys, hi^lo) + for i, k := range keys { + finite[i] = floatKeyToFloat(k) + } + for range len(floats) - m { + finite = append(finite, math.NaN()) + } + return finite +} + +// Reverse returns a copy with the elements in the opposite order; it +// works for every dtype, complex included. +func Reverse(a *Array) *Array { + out := &Array{shape: a.Shape(), dt: a.dt} + n := a.Len() + out.alloc(n) + if !a.isContiguous() { + // A strided source has no payload run to mirror against, so each + // element reads through its own accessor. + for i := range n { + out.setFrom(i, a, n-1-i) + } + return out + } + switch a.dt { + case Int: + reverseInto(out.ints, a.ints[:n]) + case Bool: + reverseInto(out.bools, a.bools[:n]) + case Int8: + reverseInto(out.i8s, a.i8s[:n]) + case Uint8: + reverseInto(out.u8s, a.u8s[:n]) + case Int16: + reverseInto(out.i16s, a.i16s[:n]) + case Uint16: + reverseInto(out.u16s, a.u16s[:n]) + case Int32: + reverseInto(out.i32s, a.i32s[:n]) + case Uint32: + reverseInto(out.u32s, a.u32s[:n]) + case Float16: + reverseInto(out.halves, a.halves[:n]) + case Float32: + reverseInto(out.floats32, a.floats32[:n]) + case Float: + reverseInto(out.floats, a.floats[:n]) + default: + reverseInto(out.complexes, a.complexes[:n]) + } + return out +} + +// reverseInto fills dst[i] with src[n-1-i] over a parallel split of the +// destination: the chunks are disjoint and each mirrors only the range +// it owns, so no two workers touch one slot. +func reverseInto[T any](dst, src []T) { + n := len(dst) + parallelMin(n, copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + dst[i] = src[n-1-i] + } + }) +} + +// reverseIntoSerial is reverseInto for one row of a larger walk, where +// spawning workers per row would cost more than the row. +func reverseIntoSerial[T any](dst, src []T) { + n := len(dst) + for i := range dst { + dst[i] = src[n-1-i] + } +} + +// ArgSort returns the stable int permutation of indices that would sort +// the array ascending, NaN at the end. +func ArgSort(a *Array) (*Array, error) { + if a.dt == Complex { + return nil, errf("ArgSort: complex arrays have no ordering") + } + if narrowRefused(a.dt) { + // The same narrow-dtype boundary Sort carries: the narrow radix + // digits carry no kernel. + return nil, errf("ArgSort: dtype %s is not supported; convert with Astype", a.dt) + } + n := a.Len() + idx := make([]int, n) + m := a.materialise() + if a.dt == Int { + // The sign-bit flip maps int64 order onto uint64 order, so the + // radix digits compare exactly as the signed values do. The + // loop is bound to n, not to the payload length: a rebased view + // carries a payload longer than its own Len. The key range is + // folded into the same walk, so the passes only cover the bytes + // the keys differ in. + ints := m.ints + keys := make([]uint64, n) + hi, lo := uint64(0), ^uint64(0) + for i := range n { + k := uint64(ints[i]) ^ (1 << 63) + keys[i] = k + idx[i] = i + hi |= k + lo &= k + } + argSortRadixUint64(keys, idx, hi^lo) + } else { + // Each value is materialised once through the same widening + // floatAt applies (float32 and float16 convert exactly). The NaN + // rules are unchanged, and the stable radix keeps the tie order + // (equal values and NaN-to-NaN alike), so the index sequence is + // identical to the comparator sort this replaced: -0.0 and +0.0 + // fold to one key because the comparator has always tied them. + // NaN indices are partitioned out and re-appended in their + // original order, so they need no keys at all. + var vals []float64 + if m.dt == Float32 { + vals = make([]float64, n) + f32 := m.floats32 + for i := range vals { + vals[i] = float64(f32[i]) + } + } else if m.dt == Float16 { + vals = make([]float64, n) + halves := m.halves + for i := range vals { + vals[i] = HalfToFloat64(halves[i]) + } + } else { + vals = m.floats + } + // The keys are held at the original positions, so the radix + // permutes the original indices directly; a NaN slot keeps the + // zero key and is never read, because only the finite indices + // enter the permutation. + keys := make([]uint64, n) + kept := make([]int, 0, n) // original indices of the finite values + var nans []int + hi, lo := uint64(0), ^uint64(0) + // Bound to n as above: materialise returns a contiguous view + // unchanged, so vals can be longer than the logical array. + for i := range n { + v := vals[i] + if v != v { + nans = append(nans, i) + continue + } + k := floatSortKey(v) + keys[i] = k + kept = append(kept, i) + hi |= k + lo &= k + } + argSortRadixUint64(keys, kept, hi^lo) + // kept now holds the sorted finite block, in place; the NaN + // block lands behind it untouched. + copy(idx, kept) + copy(idx[len(kept):], nans) + } + out := make([]int64, len(idx)) + parallelMin(len(out), copyMinPerWorker, func(s, e int) { + for i := s; i < e; i++ { + out[i] = int64(idx[i]) + } + }) + return &Array{shape: []int{len(out)}, dt: Int, ints: out}, nil +} + +// floatSortKey folds a float64 onto a uint64 that carries the same +// order: positives get the sign bit set, negatives are complemented. +// The zero check folds -0.0 onto the +0.0 key because the comparator +// this radix replaced always tied the two zeros. NaN never reaches the +// function; its key would sort outside the finite range either way. +func floatSortKey(v float64) uint64 { + bits := math.Float64bits(v) + if v == 0 { + bits = 0 + } + if bits&(1<<63) != 0 { + return ^bits + } + return bits | (1 << 63) +} + +// floatKeyToFloat inverts floatSortKey. The fold that tied -0.0 to +// +0.0 is not undone: the key carries +0.0, so a sorted -0.0 comes +// back as +0.0, the same representative the two zeros have always +// collapsed to under equality. +func floatKeyToFloat(k uint64) float64 { + if k&(1<<63) != 0 { + k &^= 1 << 63 + } else { + k = ^k + } + return math.Float64frombits(k) +} + +// LSD radix thresholds. Below radixSeqMin the comparator sort wins: +// no scratch buffers and no digit passes. From radixParMin the +// multi-pass histogram and scatter form wins: the histograms, the prefix +// sums and the scatter all scale with the worker count, while the +// serial counting sort is one core walking the element stream twice. +const ( + radixSeqMin = 24 + radixParMin = 1 << 13 +) + +// radixMaxWorkers caps the crew of one parallel value digit pass. A +// pass is memory-bound: it streams the element buffer once for the +// histogram and once for the scatter, and the histogram slab grows with +// the worker count, so a full-machine crew pays more in goroutine spawns +// and slab traffic than its extra workers stream. +const radixMaxWorkers = 8 + +// radixIndexedMaxWorkers caps the crew of one permutation digit pass. +// The permutation scatter reads the keys through the index stream, a +// random access per element; measured against the value scatter across +// the five sort fixtures, a crew beyond this loses to spawn and +// histogram traffic on every one of them, int fixtures included, so the +// permutation pass takes the same cap as the value pass. +const radixIndexedMaxWorkers = 8 + +// radixBuckets is the bucket count of one parallel digit. Twelve-bit +// digits cut the 64-bit key into at most six passes whose histogram +// stays cache-resident, against the eight 8-bit passes the serial path +// walks; the digit grouping never moves a result (see +// argSortRadixParallel). +const radixBuckets = 1 << 12 + +// radix12Shifts are the bit positions of the parallel radix digits. +var radix12Shifts = [6]uint{0, 12, 24, 36, 48, 60} + +// digitMaskOf folds the accumulated key range diff, the bitwise OR of +// every key pair's differences, into the byte positions at which the +// keys are not all identical. A digit pass over a byte every key shares +// scatters every element into one bucket in input order, so it is a +// stable no-op and can be skipped; zero means every key is equal and +// the order in place is already the stable one. A byte the caller never +// folded (an empty key set) reads as constant zero and is skipped too. +func digitMaskOf(diff uint64) uint32 { + var mask uint32 + for b := range 8 { + if diff>>(8*b)&0xFF != 0 { + mask |= 1 << b + } + } + return mask +} + +// digitMask12Of marks the parallel radix digits within whose bits the +// key range diff carries a set bit. The fold is exact in the key bits: +// a digit whose bits are constant across every key scatters stably in +// place and is skipped, so no pass runs without cause. +func digitMask12Of(diff uint64) uint32 { + var mask uint32 + for j, shift := range radix12Shifts { + width := uint(12) + if j == len(radix12Shifts)-1 { + width = 64 - shift + } + bits := diff >> shift + if width < 64 { + bits &= ^(^uint64(0) << width) + } + if bits != 0 { + mask |= 1 << j + } + } + return mask +} + +// argSortRadixUint64 stably permutes idx so that keys ascend. Below +// radixParMin it runs an 8-bit-digit counting sort; above it the +// parallel form takes over (see argSortRadixParallel). Both skip the +// digits their mask reports as constant (see digitMaskOf). diff is the +// bitwise OR of every key pair's differences, the range the digits are +// masked against. The keys must be order-preserving maps of the sort +// values, so unsigned key order equals value order exactly and the +// stable digit scatter reproduces the unique stable permutation the +// comparator sort returned. idx carries the element indices to permute, +// so it may be a strict subset of the keys' positions: every entry must +// index a key the permutation is meant to order. +func argSortRadixUint64(keys []uint64, idx []int, diff uint64) { + n := len(idx) + if n < radixSeqMin { + slices.SortStableFunc(idx, func(x, y int) int { + return cmp.Compare(keys[x], keys[y]) + }) + return + } + if diff == 0 { + return // every key is equal: the order in place is the stable one + } + if n < radixParMin { + argSortRadixSerial(keys, idx, digitMaskOf(diff)) + return + } + argSortRadixParallel(keys, idx, diff) +} + +// argSortRadixSerial runs the classic counting sort per digit: one +// histogram pass, an exclusive prefix sum over the 256 buckets and one +// scatter pass, alternating between the permutation and a scratch +// buffer so no per-element swaps are needed. Only the digits mask marks +// run; the parity of that count decides which buffer holds the result. +func argSortRadixSerial(keys []uint64, idx []int, mask uint32) { + n := len(idx) + buf := make([]int, n) + src, dst := idx, buf + var count [256]int + passes := 0 + for b := range 8 { + if mask&(1<>shift&0xFF]++ + } + sum := 0 + for bb, c := range count { + count[bb], sum = sum, sum+c + } + for _, i := range src { + d := keys[i] >> shift & 0xFF + dst[count[d]] = i + count[d]++ + } + src, dst = dst, src + passes++ + } + if passes%2 == 1 { + copy(idx, buf) + } +} + +// forEachRadixChunk runs fn once per chunk of the caller's own split of +// [0, n). chunk is the width the caller snapshotted from the same +// WorkersFor read that sized its histogram, and fn receives the row its +// range owns, so the row index is a property of the split rather than a +// second reading of the global worker count: a concurrent SetNumCPU can +// move rows between the spawned goroutines, never two live chunks onto +// one histogram row. Rows arrive in increasing order, which is +// what keeps the scatter stable, and a chunk width below radixParMin +// runs the whole range on the calling goroutine exactly as the spawn +// floor of engine.ParallelMin did. The width comes from the caller's +// snapshot, so the chunk count never exceeds the worker count that sized +// the caller's histogram rows. +func forEachRadixChunk(n, chunk int, fn func(row, start, end int)) { + if chunk < radixParMin || chunk >= n { + // One chunk: either the whole range sits below the parallel + // floor, or a single worker owns all of it. + fn(0, 0, n) + return + } + nchunks := (n + chunk - 1) / chunk + engine.Parallel(nchunks, func(lo, hi int) { + for row := lo; row < hi; row++ { + start := row * chunk + fn(row, start, min(start+chunk, n)) + } + }) +} + +// argSortRadixParallel runs each 12-bit digit as two worker waves. The +// first wave computes one histogram row per chunk; the rows are +// exclusive prefix-summed per bucket, which gives every chunk its own +// disjoint segment of every bucket; the second wave scatters each chunk +// into its segments. The segments never overlap and the scatter inside a +// chunk is sequential, so the result is stable across chunks as well as +// within one and no lock is needed. Neither the digit grouping nor the +// chunk count can move a result: the scatter is stable under both, so +// the permutation is the unique stable order by key whatever digit size +// and crew carry it. Only the digits the key range marks as varying run. +func argSortRadixParallel(keys []uint64, idx []int, diff uint64) { + n := len(idx) + buf := make([]int, n) + src, dst := idx, buf + // The permutation scatter reads the keys through the index stream, a + // random access per element, so it takes a wider crew than the value + // radix's streaming scatter; the histogram slab stays bounded by the + // 12-bit digit count. + w := min(engine.WorkersFor(n), radixIndexedMaxWorkers) + chunk := (n + w - 1) / w + if chunk < radixParMin || chunk >= n { + // One chunk: the same fallback forEachRadixChunk applies, so the + // histogram slab is sized for the rows the scatter will touch. + w = 1 + chunk = n + } + digits := digitMask12Of(diff) + // One slab for both the per-chunk counts and, rewritten in place, the + // per-chunk scatter offsets: every slot is read once and overwritten + // once by the exclusive prefix sum below. + hist := make([]int, w*radixBuckets) + passes := 0 + for j, shift := range radix12Shifts { + if digits&(1<>shift)&(radixBuckets-1))]++ + } + }) + sum := 0 + for bb := range radixBuckets { + for c := range w { + cnt := hist[c*radixBuckets+bb] + hist[c*radixBuckets+bb] = sum + sum += cnt + } + } + forEachRadixChunk(n, chunk, func(row, start, end int) { + base := row * radixBuckets + for _, i := range src[start:end] { + d := base + (int(keys[i]>>shift) & (radixBuckets - 1)) + dst[hist[d]] = i + hist[d]++ + } + }) + src, dst = dst, src + passes++ + } + if passes%2 == 1 { + copy(idx, buf) + } +} + +// radixSortUint64 sorts vals ascending in place with the same LSD digit +// passes and thresholds the index radix uses, skipping the digits the +// key range marks as constant (see digitMaskOf). Equal values are +// indistinguishable, so stability is never observable here, but the +// scatter still walks each chunk in order so the two radices stay in +// step; any digit grouping and any crew produce the same ascending +// sequence. +func radixSortUint64(vals []uint64, diff uint64) { + n := len(vals) + if n < radixSeqMin { + slices.Sort(vals) + return + } + if diff == 0 { + return + } + if n < radixParMin { + buf := make([]uint64, n) + src, dst := vals, buf + mask := digitMaskOf(diff) + passes := 0 + var count [256]int + for b := range 8 { + if mask&(1<>shift&0xFF]++ + } + sum := 0 + for bb, c := range count { + count[bb], sum = sum, sum+c + } + for _, v := range src { + d := v >> shift & 0xFF + dst[count[d]] = v + count[d]++ + } + src, dst = dst, src + passes++ + } + if passes%2 == 1 { + copy(vals, buf) + } + return + } + buf := make([]uint64, n) + src, dst := vals, buf + w := min(engine.WorkersFor(n), radixMaxWorkers) + chunk := (n + w - 1) / w + if chunk < radixParMin || chunk >= n { + w = 1 + chunk = n + } + digits := digitMask12Of(diff) + // As above: counts first, then the same slab carries the offsets. + hist := make([]int, w*radixBuckets) + passes := 0 + for j, shift := range radix12Shifts { + if digits&(1<>shift)&(radixBuckets-1))]++ + } + }) + sum := 0 + for bb := range radixBuckets { + for c := range w { + cnt := hist[c*radixBuckets+bb] + hist[c*radixBuckets+bb] = sum + sum += cnt + } + } + forEachRadixChunk(n, chunk, func(row, start, end int) { + base := row * radixBuckets + for _, v := range src[start:end] { + d := base + (int(v>>shift) & (radixBuckets - 1)) + dst[hist[d]] = v + hist[d]++ + } + }) + src, dst = dst, src + passes++ + } + if passes%2 == 1 { + copy(vals, buf) + } +} diff --git a/internal/core/sort_test.go b/internal/core/sort_test.go new file mode 100644 index 0000000..3de2b91 --- /dev/null +++ b/internal/core/sort_test.go @@ -0,0 +1,387 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "cmp" + "fmt" + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +func TestSort(t *testing.T) { + a := mustFromInts(t, []int64{3, 1, 2}, 3) + s, err := Sort(a) + if err != nil { + t.Fatalf("Sort: %v", err) + } + if !Equal(mustFromInts(t, []int64{1, 2, 3}, 3), s) { + t.Fatalf("Sort: %s", s) + } + // The receiver is untouched. + if v, _ := IntAt(a, 0); v != 3 { + t.Fatalf("Sort mutated the receiver: %d", v) + } + + f := mustFromFloats(t, []float64{2.5, math.NaN(), 0.5, 1.5}, 4) + fs, _ := Sort(f) + want := []float64{0.5, 1.5, 2.5, math.NaN()} + for i := range 4 { + v, _ := FloatAt(fs, i) + if math.IsNaN(want[i]) != math.IsNaN(v) || (!math.IsNaN(v) && v != want[i]) { + t.Fatalf("Sort NaN-at-end: %s", fs) + } + } + + if _, err := Sort(mustFromComplexes(t, []complex128{1}, 1)); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("Sort complex: %v", err) + } +} + +func TestReverse(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3}, 3) + r := Reverse(a) + if !Equal(mustFromInts(t, []int64{3, 2, 1}, 3), r) { + t.Fatalf("Reverse: %s", r) + } + + // Descending composes as Sort().Reverse(). + desc, _ := Sort(a) + desc = Reverse(desc) + if !Equal(mustFromInts(t, []int64{3, 2, 1}, 3), desc) { + t.Fatalf("descending: %s", desc) + } + + // Reverse works for complex too. + c := mustFromComplexes(t, []complex128{1, complex(2, 2)}, 2) + cr := Reverse(c) + if v, _ := ComplexAt(cr, 0); v != complex(2, 2) { + t.Fatalf("Reverse complex: %v", v) + } +} + +func TestArgSort(t *testing.T) { + a := mustFromInts(t, []int64{30, 10, 20, 10}, 4) + idx, err := ArgSort(a) + if err != nil { + t.Fatalf("ArgSort: %v", err) + } + // Stable: the two 10s keep their relative order (1 before 3). + want := mustFromInts(t, []int64{1, 3, 2, 0}, 4) + if !Equal(want, idx) { + t.Fatalf("ArgSort: %s", idx) + } + + f := mustFromFloats(t, []float64{2.0, math.NaN(), 1.0}, 3) + fidx, _ := ArgSort(f) + if !Equal(mustFromInts(t, []int64{2, 0, 1}, 3), fidx) { + t.Fatalf("ArgSort NaN last: %s", fidx) + } + + // float32 reads through the float64 accessor: no panic, correct order. + f32 := mustFromFloat32s(t, []float32{3, 1, 2}, 3) + f32idx, err := ArgSort(f32) + if err != nil { + t.Fatalf("ArgSort float32: %v", err) + } + if !Equal(mustFromInts(t, []int64{1, 2, 0}, 3), f32idx) { + t.Fatalf("ArgSort float32: %s", f32idx) + } + + if _, err := ArgSort(mustFromComplexes(t, []complex128{1}, 1)); err == nil || !strings.Contains(err.Error(), "no ordering") { + t.Fatalf("ArgSort complex: %v", err) + } +} + +// argSortReference is the comparator sort ArgSort ran before the LSD +// radix replaced it. The stable permutation is unique, so the tests +// below keep the comparator as the oracle every radix path must +// reproduce index for index. +func argSortReference(vals []float64) []int64 { + idx := make([]int, len(vals)) + for i := range idx { + idx[i] = i + } + slices.SortStableFunc(idx, func(x, y int) int { + vx, vy := vals[x], vals[y] + switch { + case vx != vx && vy != vy: + return 0 + case vx != vx: + return 1 + case vy != vy: + return -1 + case vx < vy: + return -1 + case vx > vy: + return 1 + default: + return 0 + } + }) + out := make([]int64, len(idx)) + for i, v := range idx { + out[i] = int64(v) + } + return out +} + +func argSortReferenceInt(vals []int64) []int64 { + idx := make([]int, len(vals)) + for i := range idx { + idx[i] = i + } + slices.SortStableFunc(idx, func(x, y int) int { + return cmp.Compare(vals[x], vals[y]) + }) + out := make([]int64, len(idx)) + for i, v := range idx { + out[i] = int64(v) + } + return out +} + +func checkArgSortFloat(t *testing.T, vals []float64) { + t.Helper() + want := argSortReference(vals) + got, err := ArgSort(mustFromFloats(t, vals, len(vals))) + if err != nil { + t.Fatalf("ArgSort: %v", err) + } + if !slices.Equal(want, got.RawInts()) { + t.Fatalf("ArgSort diverged from the comparator oracle:\n got %v\nwant %v", + got.RawInts(), want) + } +} + +func checkArgSortInt(t *testing.T, vals []int64) { + t.Helper() + want := argSortReferenceInt(vals) + got, err := ArgSort(mustFromInts(t, vals, len(vals))) + if err != nil { + t.Fatalf("ArgSort: %v", err) + } + if !slices.Equal(want, got.RawInts()) { + t.Fatalf("ArgSort diverged from the comparator oracle:\n got %v\nwant %v", + got.RawInts(), want) + } +} + +// TestArgSortMatchesComparator drives every radix path against the +// oracle: the comparator fallback below radixSeqMin, the serial +// counting sort, and the parallel histogram/scatter either side of +// radixParMin under several worker counts. The special values force +// ties, both zeros, infinities and NaN blocks into every size class. +func TestArgSortMatchesComparator(t *testing.T) { + sizes := []int{0, 1, 2, 23, 24, 25, 100, 4095, 8191, 8192, 8193, 40_000} + specials := []float64{ + math.NaN(), math.Copysign(0, -1), 0, math.Inf(1), math.Inf(-1), 42, 42, math.NaN(), + } + workers := []int{1, 3, 0} // 0 restores the NumCPU default + for _, w := range workers { + prev := engine.SetNumWorkers(w) + for _, n := range sizes { + t.Run(fmt.Sprintf("float64/w%d/n%d", w, n), func(t *testing.T) { + g := NewGenerator(int64(n)*31 + 7) + src, err := Floats(g, n) + if err != nil { + t.Fatalf("Floats: %v", err) + } + vals := src.RawFloats() + for i := 0; i < n; i += 5 { + vals[i] = specials[(i/5)%len(specials)] + } + checkArgSortFloat(t, vals) + }) + t.Run(fmt.Sprintf("int64/w%d/n%d", w, n), func(t *testing.T) { + g := NewGenerator(int64(n)*13 + 3) + src, err := Floats(g, n) + if err != nil { + t.Fatalf("Floats: %v", err) + } + f := src.RawFloats() + vals := make([]int64, n) + for i := range vals { + vals[i] = int64(f[i]*100) - 50 // few distinct: heavy ties + } + if n > 0 { + vals[0] = math.MinInt64 + } + if n > 1 { + vals[1] = math.MaxInt64 + } + if n > 3 { + vals[3] = -1 + } + checkArgSortInt(t, vals) + }) + } + engine.SetNumWorkers(prev) + } + + // A wide tie field: every value recurs across all parallel chunks, + // which is where a chunk-stability mistake in the scatter would show. + t.Run("ties-across-chunks", func(t *testing.T) { + const n = 100_000 + g := NewGenerator(99) + src, _ := Floats(g, n) + f := src.RawFloats() + for i := range f { + f[i] = float64(i % 4) // four values, 25k ties each + } + checkArgSortFloat(t, f) + }) + + // A rebased view reads through materialise, not the raw payload: + // Slice(0, 1, 5) yields Len 4 over a payload of 5, so any loop + // walking the payload instead of Len would read past the view. + t.Run("rebased-view", func(t *testing.T) { + full := mustFromInts(t, []int64{5, 4, 3, 2, 1, 0}, 6) + view, err := Slice(full, 0, 1, 5) + if err != nil { + t.Fatalf("Slice: %v", err) + } + got, err := ArgSort(view) + if err != nil { + t.Fatalf("ArgSort: %v", err) + } + want := mustFromInts(t, []int64{3, 2, 1, 0}, 4) + if !Equal(want, got) { + t.Fatalf("ArgSort view: %s", got) + } + }) + + // The float32 path sorts through the exact float64 widening, so the + // f64 oracle on the widened values is the right reference. + t.Run("float32", func(t *testing.T) { + const n = 20_000 + g := NewGenerator(5) + src, _ := Floats(g, n) + f := src.RawFloats() + f32vals := make([]float32, n) + widened := make([]float64, n) + for i := range f32vals { + f32vals[i] = float32(f[i]) + widened[i] = float64(f32vals[i]) + } + f32vals[n/2] = float32(math.NaN()) + widened[n/2] = math.NaN() + want := argSortReference(widened) + got, err := ArgSort(mustFromFloat32s(t, f32vals, n)) + if err != nil { + t.Fatalf("ArgSort: %v", err) + } + if !slices.Equal(want, got.RawInts()) { + t.Fatalf("ArgSort float32 diverged from the oracle") + } + }) +} + +// sortReference is the comparator Sort ran before the digit passes: +// finite values ascending, NaN appended at the end. +func sortReference(vals []float64) []float64 { + out := append([]float64(nil), vals...) + nans := 0 + finite := out[:0] + for _, v := range out { + if v != v { + nans++ + continue + } + finite = append(finite, v) + } + slices.Sort(finite) + for range nans { + finite = append(finite, math.NaN()) + } + return finite +} + +// equalSortedFloats compares sort output value for value. The zeros +// compare equal under ==, which is the right bar: the comparator sort +// this oracle stands for was not stable, so either zero sign is a +// valid result for a tie the radix folds to +0.0. +func equalSortedFloats(a, b []float64) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if math.IsNaN(a[i]) != math.IsNaN(b[i]) { + return false + } + if !math.IsNaN(a[i]) && a[i] != b[i] { + return false + } + } + return true +} + +// TestSortMatchesComparator drives the value radix against the +// comparator oracle across the same threshold and worker grid the +// permutation radix uses. +func TestSortMatchesComparator(t *testing.T) { + sizes := []int{0, 1, 2, 23, 24, 25, 100, 4095, 8191, 8192, 8193, 40_000} + specials := []float64{ + math.NaN(), math.Copysign(0, -1), 0, math.Inf(1), math.Inf(-1), 42, 42, math.NaN(), + } + workers := []int{1, 3, 0} // 0 restores the NumCPU default + for _, w := range workers { + prev := engine.SetNumWorkers(w) + for _, n := range sizes { + t.Run(fmt.Sprintf("float64/w%d/n%d", w, n), func(t *testing.T) { + g := NewGenerator(int64(n)*17 + 5) + src, err := Floats(g, n) + if err != nil { + t.Fatalf("Floats: %v", err) + } + vals := src.RawFloats() + for i := 0; i < n; i += 5 { + vals[i] = specials[(i/5)%len(specials)] + } + want := sortReference(vals) + got, err := Sort(mustFromFloats(t, vals, n)) + if err != nil { + t.Fatalf("Sort: %v", err) + } + if !equalSortedFloats(want, got.RawFloats()) { + t.Fatalf("Sort diverged from the comparator oracle:\n got %v\nwant %v", + got.RawFloats(), want) + } + }) + t.Run(fmt.Sprintf("int64/w%d/n%d", w, n), func(t *testing.T) { + g := NewGenerator(int64(n)*29 + 11) + src, err := Floats(g, n) + if err != nil { + t.Fatalf("Floats: %v", err) + } + f := src.RawFloats() + vals := make([]int64, n) + for i := range vals { + vals[i] = int64(f[i]*100) - 50 + } + if n > 0 { + vals[0] = math.MinInt64 + } + if n > 1 { + vals[1] = math.MaxInt64 + } + want := append([]int64(nil), vals...) + slices.Sort(want) + got, err := Sort(mustFromInts(t, vals, n)) + if err != nil { + t.Fatalf("Sort: %v", err) + } + if !slices.Equal(want, got.RawInts()) { + t.Fatalf("Sort diverged from the comparator oracle:\n got %v\nwant %v", + got.RawInts(), want) + } + }) + } + engine.SetNumWorkers(prev) + } +} diff --git a/internal/core/sparse.go b/internal/core/sparse.go new file mode 100644 index 0000000..d65d0ff --- /dev/null +++ b/internal/core/sparse.go @@ -0,0 +1,373 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +// Sparse COO arrays. Used for embedding tables, attention +// masks, recommender features: anywhere the data is sparse enough that +// a dense representation would waste memory and bandwidth. +// +// A SparseCOO stores non-zero values as (indices, values) pairs; the +// full shape is part of the metadata. Operations preserve sparsity: +// multiplying a SparseCOO by a dense array stays sparse when the +// dense side does not introduce new non-zeros. + +// SparseCOO is a coordinate-format sparse array. Indices has shape +// (nnz, ndim) and dtype int; values has shape (nnz,) and the +// element dtype of the array. Shape is the dense shape. +type SparseCOO struct { + Indices *Array + Values *Array + Shape []int +} + +// NewSparseCOO creates a SparseCOO from explicit indices, values, and +// shape. Returns an error if either array is nil, if indices and values +// disagree on nnz, or if indices doesn't have the right rank. +func NewSparseCOO(indices, values *Array, shape []int) (*SparseCOO, error) { + // A nil array is what a caller holds after a constructor refused its + // arguments, so it is a refusal here rather than a nil-pointer panic + // on the fields read below. + if indices == nil { + return nil, errf("NewSparseCOO: indices must not be nil") + } + if values == nil { + return nil, errf("NewSparseCOO: values must not be nil") + } + if indices.dt != Int { + return nil, errf("NewSparseCOO: indices must be int, got %s", indices.dt) + } + if indices.NDim() != 2 { + return nil, errf("NewSparseCOO: indices must be 2-D (nnz, ndim), got shape %s", shapeText(indices.shape)) + } + if indices.shape[1] != len(shape) { + return nil, errf("NewSparseCOO: indices dim %d does not match shape %v", indices.shape[1], shape) + } + if values.Len() != indices.shape[0] { + return nil, errf("NewSparseCOO: values len %d does not match nnz %d", values.Len(), indices.shape[0]) + } + if narrowRefused(values.dt) { + return nil, errf("NewSparseCOO: dtype %s is not supported; convert with Astype", values.dt) + } + // The dense shape is metadata the arithmetic trusts (SpMatMul + // allocates from it directly), so it gets the same validation and + // private copy every constructor applies. + if _, _, serr := checkedDims(shape); serr != nil { + return nil, base.WrapErr("NewSparseCOO", serr) + } + sh := make([]int, len(shape)) + copy(sh, shape) + return &SparseCOO{Indices: indices, Values: values, Shape: sh}, nil +} + +// SparseFrom extracts a SparseCOO from a dense array by keeping only +// the non-zero elements. Renamed from `SparseFromDense` to mirror +// the symmetry with `SparseCOO.Dense`. The values array keeps the +// dense array's dtype, so the round trip preserves the element type. +func SparseFrom(dense *Array) (*SparseCOO, error) { + if dense.dt == Complex { + return nil, errf("SparseFrom: complex arrays are not supported") + } + if narrowRefused(dense.dt) { + return nil, errf("SparseFrom: dtype %s is not supported; convert with Astype", dense.dt) + } + shape := dense.Shape() + var coords []int64 + for flat := range dense.Len() { + if !isZero(dense, flat) { + coord := make([]int, dense.NDim()) + rem := flat + for d := range dense.NDim() { + stride := 1 + for k := d + 1; k < dense.NDim(); k++ { + stride *= shape[k] + } + coord[d] = rem / stride + rem %= stride + } + for d := range dense.NDim() { + coords = append(coords, int64(coord[d])) + } + } + } + nnz := len(coords) / dense.NDim() + indices, err := FromInts(coords, nnz, dense.NDim()) + if err != nil { + return nil, err + } + var values *Array + switch dense.dt { + case Int: + ints := make([]int64, nnz) + k := 0 + for flat := range dense.Len() { + if !isZero(dense, flat) { + ints[k] = dense.ints[flat] + k++ + } + } + values, err = FromInts(ints, nnz) + case Float16: + halves := make([]uint16, nnz) + k := 0 + for flat := range dense.Len() { + if !isZero(dense, flat) { + halves[k] = dense.halves[flat] + k++ + } + } + values, err = HalvesFromArray(halves, nnz) + case Float32: + f32 := make([]float32, nnz) + k := 0 + for flat := range dense.Len() { + if !isZero(dense, flat) { + f32[k] = dense.floats32[flat] + k++ + } + } + values, err = FromFloat32s(f32, nnz) + default: + floats := make([]float64, nnz) + k := 0 + for flat := range dense.Len() { + if !isZero(dense, flat) { + floats[k] = dense.floats[flat] + k++ + } + } + values, err = FromFloats(floats, nnz) + } + if err != nil { + return nil, err + } + return &SparseCOO{Indices: indices, Values: values, Shape: shape}, nil +} + +// consistent validates the index/value pairing of a hand-built +// SparseCOO literal, whose exported fields the constructor's checks +// never saw. Every entry point that walks the stored coordinates +// calls this first, so a mismatched pair is an error, never an +// out-of-bounds panic. +func (s *SparseCOO) consistent(name string) error { + if s.Values.Len() != s.Indices.Shape()[0] { + return errf("%s: %d values for %d index rows", name, s.Values.Len(), s.Indices.Shape()[0]) + } + return nil +} + +// Dense materialises the sparse array as a dense *Array. Renamed +// from `ToDense` for symmetry with `SparseFrom`. +func (s *SparseCOO) Dense() (*Array, error) { + if err := s.consistent("Dense"); err != nil { + return nil, err + } + if narrowRefused(s.Values.dt) { + return nil, errf("Dense: dtype %s is not supported; convert with Astype", s.Values.dt) + } + out, err := Zeros(s.Values.dt, s.Shape...) + if err != nil { + return nil, err + } + strides := denseStrides(s.Shape) + for i := range s.Indices.shape[0] { + flat, err := s.flatEntry(i, strides, "Dense") + if err != nil { + return nil, err + } + switch s.Values.dt { + case Int: + out.ints[flat] = s.Values.ints[i] + case Float16: + out.halves[flat] = s.Values.halves[i] + case Float32: + out.floats32[flat] = s.Values.floats32[i] + case Float: + out.floats[flat] = s.Values.floats[i] + default: + out.complexes[flat] = s.Values.complexes[i] + } + } + return out, nil +} + +// NNZ returns the number of non-zero entries. +func (s *SparseCOO) NNZ() int { + return s.Values.Len() +} + +// SpMul multiplies a sparse array element-wise by a dense array of the +// same shape. Returns a dense array whose dtype follows the promotion +// ladder. +func SpMul(s *SparseCOO, dense *Array) (*Array, error) { + if err := s.consistent("SpMul"); err != nil { + return nil, err + } + if narrowRefused(s.Values.dt) { + return nil, errf("SpMul: dtype %s is not supported; convert with Astype", s.Values.dt) + } + if narrowRefused(dense.dt) { + return nil, errf("SpMul: dtype %s is not supported; convert with Astype", dense.dt) + } + if !sameShape(s.Shape, dense.shape) { + return nil, errf("SpMul: shape mismatch %v vs %s", s.Shape, shapeText(dense.shape)) + } + dt := promote(s.Values.dt, dense.dt) + out, err := Zeros(dt, s.Shape...) + if err != nil { + return nil, err + } + strides := denseStrides(s.Shape) + for i := range s.Indices.shape[0] { + flat, err := s.flatEntry(i, strides, "SpMul") + if err != nil { + return nil, err + } + switch dt { + case Int: + out.ints[flat] = s.Values.ints[i] * dense.ints[flat] + case Float16: + // Computed in float64, narrowed once, like the float32 + // branch. + out.halves[flat] = HalfFromFloat64(s.Values.FloatAt(i) * dense.FloatAt(flat)) + case Float32: + out.floats32[flat] = float32(s.Values.FloatAt(i) * dense.FloatAt(flat)) + case Float: + out.floats[flat] = s.Values.FloatAt(i) * dense.FloatAt(flat) + default: + out.complexes[flat] = s.Values.complexAt(i) * dense.complexAt(flat) + } + } + return out, nil +} + +// SpAdd returns a dense array equal to the element-wise sum of two +// sparse arrays. The result is dense because addition can collapse +// zeros into non-zeros. Both arrays must have the same shape. +func SpAdd(a, b *SparseCOO) (*Array, error) { + for _, s := range []*SparseCOO{a, b} { + if narrowRefused(s.Values.dt) { + return nil, errf("SpAdd: dtype %s is not supported; convert with Astype", s.Values.dt) + } + } + if !sameShape(a.Shape, b.Shape) { + return nil, errf("SpAdd: shape mismatch %v vs %v", a.Shape, b.Shape) + } + da, err := a.Dense() + if err != nil { + return nil, err + } + db, err := b.Dense() + if err != nil { + return nil, err + } + return Add(da, db) +} + +// SpMatMul multiplies a sparse matrix (n×k) by a dense matrix (k×m). +// Returns a dense n×m result. Only valid for 2-D sparse × 2-D dense. +// The result dtype follows the promotion ladder, like SpMul: int with +// int stays int, any float or complex operand promotes. Every stored +// coordinate is validated, so an out-of-range index is an error naming +// the entry, never an out-of-bounds panic. +func SpMatMul(s *SparseCOO, dense *Array) (*Array, error) { + if err := s.consistent("SpMatMul"); err != nil { + return nil, err + } + if s.Values.dt == Complex { + return nil, errf("SpMatMul: complex sparse is not supported") + } + if s.Values.dt == Float16 || dense.dt == Float16 { + // Like MatMul, the matrix product kernels are not offered for + // the half dtype yet; the refusal is loud, not a silent + // misread of the payload. + return nil, errf("SpMatMul: float16 is not supported; convert with Astype") + } + for _, op := range []Dtype{s.Values.dt, dense.dt} { + if narrowRefused(op) { + return nil, errf("SpMatMul: dtype %s is not supported; convert with Astype", op) + } + } + if len(s.Shape) != 2 || dense.NDim() != 2 { + return nil, errf("SpMatMul: needs 2-D sparse and 2-D dense, got %d-D and %d-D", + len(s.Shape), dense.NDim()) + } + if s.Shape[1] != dense.shape[0] { + return nil, errf("SpMatMul: inner dim mismatch %d vs %d", s.Shape[1], dense.shape[0]) + } + dt := promote(s.Values.dt, dense.dt) + n := s.Shape[0] + k := s.Shape[1] + m := dense.shape[1] + out := &Array{shape: []int{n, m}, dt: dt} + out.alloc(n * m) + strides := denseStrides(s.Shape) + // The Float32 branch accumulates its products in float64 and narrows + // once per stored coordinate, so one scratch row outlives the + // whole walk. + var acc []float64 + if dt == Float32 { + acc = make([]float64, m) + } + for i := range s.Indices.shape[0] { + flat, err := s.flatEntry(i, strides, "SpMatMul") + if err != nil { + return nil, err + } + row := flat / k + col := flat % k + // Add s.Values[i] * dense[col, j] to out[row, j] for each j. + switch dt { + case Int: + for j := range m { + out.ints[row*m+j] += s.Values.ints[i] * dense.ints[col*m+j] + } + case Float32: + vf := s.Values.FloatAt(i) + for j := range m { + acc[j] = float64(out.floats32[row*m+j]) + vf*dense.FloatAt(col*m+j) + } + for j := range m { + out.floats32[row*m+j] = float32(acc[j]) + } + case Float: + for j := range m { + out.floats[row*m+j] += s.Values.FloatAt(i) * dense.FloatAt(col*m+j) + } + default: + for j := range m { + out.complexes[row*m+j] += s.Values.complexAt(i) * dense.complexAt(col*m+j) + } + } + } + return out, nil +} + +// denseStrides returns the row-major strides for a shape slice. +func denseStrides(shape []int) []int { + strides := make([]int, len(shape)) + stride := 1 + for i := len(shape) - 1; i >= 0; i-- { + strides[i] = stride + stride *= shape[i] + } + return strides +} + +// flatEntry returns the flat dense offset of the i-th stored +// coordinate, checking every index component against the shape. name +// names the calling entry point in the error. +func (s *SparseCOO) flatEntry(i int, strides []int, name string) (int, error) { + flat := 0 + for d := range s.Shape { + idx := int(s.Indices.ints[i*len(s.Shape)+d]) + if idx < 0 || idx >= s.Shape[d] { + return 0, errf("%s: index [%d]=%d out of range for dim %d of size %d", + name, i, idx, d, s.Shape[d]) + } + flat += idx * strides[d] + } + return flat, nil +} diff --git a/internal/core/special.go b/internal/core/special.go new file mode 100644 index 0000000..e00abbd --- /dev/null +++ b/internal/core/special.go @@ -0,0 +1,410 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/cmplx" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Special functions for physics, cosmology and chemistry: the gamma +// and beta families, the error function family, Legendre polynomials, +// spherical harmonics and spherical Bessel functions. Element-wise +// functions run on real arrays and promote ints and float32 per the +// ladder; complex inputs are refused. Values outside a function's +// domain follow the IEEE behaviour of the underlying implementation +// (math.Gamma, for example, returns +Inf at non-positive integers and +// NaN where undefined). + +// Gamma returns the gamma function Γ(x) of each element. +func Gamma(a *Array) (*Array, error) { return a.realFunc("Gamma", math.Gamma) } + +// LnGamma returns the natural logarithm of |Γ(x)| of each element. The +// sign of Γ for negative arguments is dropped; callers that need it +// evaluate Gamma directly. +func LnGamma(a *Array) (*Array, error) { + return a.realFunc("LnGamma", func(x float64) float64 { + lg, _ := math.Lgamma(x) + return lg + }) +} + +// Beta returns the Euler beta function B(x, y) = Γ(x)Γ(y)/Γ(x+y) +// element-wise. Both arrays must have the same shape. +func Beta(x, y *Array) (*Array, error) { + if !sameShape(x.shape, y.shape) { + return nil, errf("Beta: shape mismatch %s vs %s", shapeText(x.shape), shapeText(y.shape)) + } + if x.dt == Complex || y.dt == Complex { + return nil, errf("Beta: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv, yv := x.floatAt(i), y.floatAt(i) + // In log space: the Γ product overflows to Inf well before + // B itself leaves the float64 range (B(100, 80) ≈ 3e-67), + // and Inf/Inf would turn the answer into NaN. + lx, sx := math.Lgamma(xv) + ly, sy := math.Lgamma(yv) + lxy, sxy := math.Lgamma(xv + yv) + sign := sx * sy * sxy + out.floats[i] = float64(sign) * math.Exp(lx+ly-lxy) + } + }) + return out, nil +} + +// Erf and Erfc return the error function and the complementary error +// function of each element. +func Erf(a *Array) (*Array, error) { return a.realFunc("Erf", math.Erf) } +func Erfc(a *Array) (*Array, error) { return a.realFunc("Erfc", math.Erfc) } + +// Sinc returns the normalised sinc sin(πx)/(πx) of each element, with +// Sinc(0) = 1. The π-normalised convention is the one signal +// processing and interpolation use. +func Sinc(a *Array) (*Array, error) { + return a.realFunc("Sinc", func(x float64) float64 { + if x == 0 { + return 1 + } + return math.Sin(math.Pi*x) / (math.Pi * x) + }) +} + +// Cosm1 returns cos(x) − 1 of each element, kept accurate where the +// direct subtraction loses: as a float64 subtraction cos(x) − 1 has +// no correct significant bit under about |x| ≈ 1e-8, while the factored +// even series in x² holds full precision down to the smallest +// representable argument, where the answer is exactly −x²/2 as far as +// the format can see it. The standard library carries log1p and expm1 +// for the logarithmic and exponential side of the same problem and no +// cosinusoidal counterpart, which is the gap this fills. +func Cosm1(a *Array) (*Array, error) { + return a.realFunc("Cosm1", cosm1) +} + +// cosm1Crossover splits the series and the direct subtraction: below +// it the factored series ends well past the rounding floor, above it +// the cancellation the subtraction suffers still leaves every bit the +// answer has. +const cosm1Crossover = math.Pi / 4 + +// cosm1 evaluates cos(x) − 1 at one point. With t = x² the series +// cos x − 1 = −(t/2)·(1 − t/12·(1 − t/30·(1 − t/56·(1 − t/90· +// (1 − t/132·(1 − t/182·(1 − t/240))))))) truncates around 2e-18 +// relative at the crossover and falls quadratically below it; the +// leading −t/2 factor carries the sign, so negative arguments need no +// branch. +func cosm1(x float64) float64 { + t := x * x + if math.Abs(x) >= cosm1Crossover { + return math.Cos(x) - 1 + } + return -0.5 * t * (1 - t/12*(1-t/30*(1-t/56*(1-t/90*(1-t/132*(1-t/182*(1-t/240))))))) +} + +// Legendre returns the Legendre polynomial P_l(x) of each element, +// computed by the Bonnet recurrence. The degree l must be ≥ 0; the +// convention is P_0 = 1, P_1 = x, no extra scaling. +func Legendre(l int, x *Array) (*Array, error) { + if l < 0 { + return nil, errf("Legendre: degree must be ≥ 0, got %d", l) + } + if x.dt == Complex { + return nil, errf("Legendre: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv := x.floatAt(i) + p0, p1 := 1.0, xv + for k := 1; k < l; k++ { + p0, p1 = p1, ((2*float64(k)+1)*xv*p1-float64(k)*p0)/float64(k+1) + } + if l == 0 { + p1 = 1 + } + out.floats[i] = p1 + } + }) + return out, nil +} + +// LegendreAssociated returns the associated Legendre function +// P_l^m(x) of each element, with the Condon-Shortley phase (−1)^m +// folded in and no extra normalisation. Requires |m| ≤ l; a negative m +// uses the standard relation P_l^(−m) = (−1)^m (l−m)!/(l+m)! P_l^m. +func LegendreAssociated(l, m int, x *Array) (*Array, error) { + if l < 0 { + return nil, errf("LegendreAssociated: degree must be ≥ 0, got %d", l) + } + if m < -l || m > l { + return nil, errf("LegendreAssociated: order %d out of range for degree %d", m, l) + } + if x.dt == Complex { + return nil, errf("LegendreAssociated: complex arrays are not supported") + } + am, neg := m, false + if am < 0 { + am, neg = -m, true + } + mm := float64(am) + // Seed sign from the Condon-Shortley phase (−1)^m. + seedSign := 1.0 + if am%2 == 1 { + seedSign = -1 + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv := x.floatAt(i) + // P_m^m = (−1)^m (2m−1)!! (1−x²)^(m/2). + fact := 1.0 + for k := 1; k <= am; k++ { + fact *= float64(2*k - 1) + } + pmm := seedSign * fact * math.Pow(1-xv*xv, mm/2) + if l == am { + out.floats[i] = pmm + continue + } + // P_{m+1}^m = x (2m+1) P_m^m, then the fixed-order + // recurrence up the degree. + pmp1 := xv * (2*mm + 1) * pmm + pl := pmp1 + for k := am + 2; k <= l; k++ { + pl = ((2*float64(k)-1)*xv*pmp1 - float64(k+am-1)*pmm) / float64(k-am) + pmm, pmp1 = pmp1, pl + } + out.floats[i] = pl + } + }) + if neg { + // P_l^(−m) = (−1)^m (l−m)!/(l+m)! P_l^m. + ratio := math.Exp(LnFactorial(l-am) - LnFactorial(l+am)) + engine.Parallel(out.Len(), func(s, e int) { + for i := s; i < e; i++ { + out.floats[i] *= seedSign * ratio + } + }) + } + return out, nil +} + +// LnFactorial returns the natural logarithm of n!. +func LnFactorial(n int) float64 { + if n < 0 { + return math.NaN() + } + lg, _ := math.Lgamma(float64(n + 1)) + return lg +} + +// SphericalHarmonic returns the complex spherical harmonic +// Y_l^m(θ, φ) element-wise over the polar angle θ and azimuth φ, +// using the Condon-Shortley phase and the orthonormal convention +// N·P_l^m(cos θ)·e^{imφ} with N = √((2l+1)/(4π)·(l−m)!/(l+m)!). +// theta and phi must have the same shape. +func SphericalHarmonic(l, m int, theta, phi *Array) (*Array, error) { + if l < 0 { + return nil, errf("SphericalHarmonic: degree must be ≥ 0, got %d", l) + } + if m < -l || m > l { + return nil, errf("SphericalHarmonic: order %d out of range for degree %d", m, l) + } + if !sameShape(theta.shape, phi.shape) { + return nil, errf("SphericalHarmonic: shape mismatch %s vs %s", + shapeText(theta.shape), shapeText(phi.shape)) + } + if theta.dt == Complex || phi.dt == Complex { + return nil, errf("SphericalHarmonic: complex arrays are not supported") + } + am, conj := m, false + if am < 0 { + am, conj = -m, true + } + // P_l^m takes the cosine of the polar angle, not the angle. + cosTheta, err := Cos(theta) + if err != nil { + return nil, err + } + p, err := LegendreAssociated(l, am, cosTheta) + if err != nil { + return nil, err + } + norm := math.Sqrt((2*float64(l) + 1) / (4 * math.Pi) * + math.Exp(LnFactorial(l-am)-LnFactorial(l+am))) + // The Condon-Shortley phase (−1)^m is already folded into the + // associated Legendre function; the normalisation carries no + // extra sign. + out := &Array{shape: append([]int{}, theta.shape...), dt: Complex} + out.alloc(theta.Len()) + engine.Parallel(theta.Len(), func(s, e int) { + for i := s; i < e; i++ { + y := complex(norm*p.floatAt(i), 0) * + cmplx.Exp(complex(0, float64(am)*phi.floatAt(i))) + if conj { + // Y_{l,−m} = (−1)^m conj(Y_{l,m}). + phase := 1.0 + if am%2 == 1 { + phase = -1 + } + y = complex(phase, 0) * cmplx.Conj(y) + } + out.complexes[i] = y + } + }) + return out, nil +} + +// SphericalBesselJ returns the spherical Bessel function of the first +// kind j_l(x) of each element. Orders at or below |x| climb the upward +// recurrence from the exact closed forms j₀ = sin x/x and +// j₁ = sin x/x² − cos x/x, the stable direction there at O(l) cost; +// higher orders run the downward Miller recurrence, which is what the +// small roots need and whose start sits above the turning point at +// order l + |x|. +func SphericalBesselJ(l int, x *Array) (*Array, error) { + if l < 0 { + return nil, errf("SphericalBesselJ: degree must be ≥ 0, got %d", l) + } + if x.dt == Complex { + return nil, errf("SphericalBesselJ: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + out.floats[i] = sphericalBesselJ(l, x.floatAt(i)) + } + }) + return out, nil +} + +// SphericalBesselY returns the spherical Bessel function of the second +// kind y_l(x) of each element, evaluated by the upward recurrence from +// y_0 = −cos(x)/x, the stable direction. Every y_l diverges to −Inf at +// x = 0. +func SphericalBesselY(l int, x *Array) (*Array, error) { + if l < 0 { + return nil, errf("SphericalBesselY: degree must be ≥ 0, got %d", l) + } + if x.dt == Complex { + return nil, errf("SphericalBesselY: complex arrays are not supported") + } + out := &Array{shape: append([]int{}, x.shape...), dt: Float} + out.alloc(x.Len()) + engine.Parallel(x.Len(), func(s, e int) { + for i := s; i < e; i++ { + xv := x.floatAt(i) + if xv == 0 { + // Divergence at the origin: fill this element and carry + // on, the rest of the chunk still needs its values. + out.floats[i] = math.Inf(-1) + continue + } + y0 := -math.Cos(xv) / xv + if l == 0 { + // Assign and carry on: a return here would abandon the + // rest of the worker's chunk at zero. + out.floats[i] = y0 + continue + } + y1 := -math.Cos(xv)/(xv*xv) - math.Sin(xv)/xv + for k := 1; k < l; k++ { + y0, y1 = y1, (2*float64(k)+1)/xv*y1-y0 + } + out.floats[i] = y1 + } + }) + return out, nil +} + +// sphericalBesselJ evaluates j_l at one point. Below the argument the +// upward recurrence climbs from the exact j₀ and j₁ at O(l) cost; at +// and above it the downward Miller recurrence, started well above the +// turning point, is the stable direction. +func sphericalBesselJ(l int, x float64) float64 { + if x == 0 { + if l == 0 { + return 1 + } + return 0 + } + if math.Abs(x) < 1e-4 { + // Small-x power series. Below ~1e-6 the unscaled Miller seed + // overflows on the way down and the renormalisation turns into + // Inf/Inf = NaN; the two-term series is relative ~1e-8 exact + // over the whole branch and underflows to the true limit for + // large l. + num := math.Pow(x, float64(l)) + dbl := 1.0 // (2l+1)!! + for k := 3; k <= 2*l+1; k += 2 { + dbl *= float64(k) + } + return num / dbl * (1 - x*x/(2*float64(2*l+3))) + } + if float64(l) <= math.Abs(x) { + return sphericalJUpward(l, x) + } + // |x| keeps the start above l for negative arguments too: a start + // at or below l would never pass the order on the way down and + // renormalise against a garbage seed. + ceiling := l + int(math.Abs(x)) + 40 + jp1, j := 0.0, 1.0 // j_{k+1}, j_k, seeded at k = ceiling + jl := 0.0 + for k := ceiling; k >= 1; k-- { + if k == l { + jl = j + } + jp1, j = j, (2*float64(k)+1)/x*j-jp1 + if aj := math.Abs(j); aj > 1e200 { + // The unscaled walk overflows for degree-argument + // combinations whose true value is representable; the + // renormalisation cancels any common factor. + jp1 /= aj + j /= aj + jl /= aj + } + } + // j now holds the unscaled j_0, which for l = 0 is the answer. + if l == 0 { + jl = j + } + return sphericalJ0(x) * jl / j +} + +// sphericalJUpward climbs the spherical Bessel functions of the first +// kind from the exact closed forms j₀ = sin x/x and +// j₁ = sin x/x² − cos x/x by the upward recurrence +// jₖ₊₁ = (2k+1)/x·jₖ − jₖ₋₁, the stable direction while the order +// stays at or below the argument. Both seeds carry the absolute scale, +// so no renormalisation is needed, and they carry the parity for free: +// the recurrence reproduces j_l(−x) = (−1)^l j_l(x). The caller +// guarantees x ≠ 0 and 0 ≤ l ≤ |x|. +func sphericalJUpward(l int, x float64) float64 { + j0 := sphericalJ0(x) + j1 := math.Sin(x)/(x*x) - math.Cos(x)/x + if l == 0 { + return j0 + } + for k := 1; k < l; k++ { + j0, j1 = j1, (2*float64(k)+1)/x*j1-j0 + } + return j1 +} + +// sphericalJ0 is the exact j_0 = sin(x)/x with the x = 0 limit. +func sphericalJ0(x float64) float64 { + if x == 0 { + return 1 + } + return math.Sin(x) / x +} diff --git a/internal/core/special2.go b/internal/core/special2.go new file mode 100644 index 0000000..c2fdb56 --- /dev/null +++ b/internal/core/special2.go @@ -0,0 +1,449 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "math" + +// Ordinary Bessel functions of the first and second kind at one real +// point, the scalar companions of the modified pair in besselmod.go. +// J grows out of a convergent power series below the crossover and a +// downward Miller recurrence above it, the direction that amplifies +// no rounding once the order passes the argument. Y climbs the +// upward recurrence from the seeds Y₀ and Y₁, the stable direction +// for the second kind, with the seeds themselves from the Frobenius +// series below the crossover and from the asymptotic expansion above +// it, where the ascending series would start paying for cancellation. + +// besselCrossover splits the power-series and recurrence regimes. At +// the crossover both sides still carry twelve significant digits, so +// the exact split point is a matter of taste rather than accuracy. +const besselCrossover = 15.0 + +// BesselJ returns the Bessel function of the first kind of integer +// order n at the real point x, Jₙ(x). Arguments with |x| below the +// crossover are served by the convergent power series +// Σ (−1)^k (x/2)^{2k+n}/(k!·Γ(k+n+1)); above it the recurrence runs in +// its stable direction, orders at or below the argument climbing +// upward from the large-argument asymptotic J₀ and J₁ at O(n) cost and +// higher orders running the downward Miller walk anchored on the same +// asymptotic J₀, whose start sits above the turning point at order +// n + |x|. A negative argument follows the parity law +// Jₙ(−x) = (−1)ⁿ Jₙ(x) and a negative order the law +// J₋ₙ(x) = (−1)ⁿ Jₙ(x); Jₙ is finite for every real x, so nothing here +// can fail. +func BesselJ(n int, x float64) float64 { + if x == 0 { + if n == 0 { + return 1 + } + return 0 + } + // Fold both parity laws into one sign: (−1)ⁿ only cares about the + // order's parity, which |n| preserves. + order, ax, flip := n, math.Abs(x), 1.0 + if order < 0 { + if order%2 != 0 { + flip = -1 + } + order = -order + } + if x < 0 && order%2 != 0 { + flip = -flip + } + if ax < besselCrossover { + return flip * besselJSeries(order, ax) + } + if float64(order) <= ax { + return flip * besselJUpward(order, ax) + } + return flip * besselJMiller(order, ax) +} + +// BesselY returns the Bessel function of the second kind of integer +// order n at the real point x, Yₙ(x), defined for x > 0; Y diverges +// at the origin and a non-positive argument is an error, not a NaN. +// Y₀ and Y₁ are seeded below the crossover by the Frobenius series +// (Abramowitz & Stegun 9.1.11 with ψ(k+1) = H_k − γ written out +// through the harmonic numbers H_k) and above it by the +// large-argument asymptotic expansion; higher orders climb the upward +// recurrence Yₙ₊₁ = 2n/x·Yₙ − Yₙ₋₁, stable for the second kind +// because the parasitic J component the seeds carry decays relative +// to Y at every step. A negative order follows the parity law +// Y₋ₙ(x) = (−1)ⁿ Yₙ(x). +// +// Errors: x ≤ 0 or NaN. +func BesselY(n int, x float64) (float64, error) { + if math.IsNaN(x) || x <= 0 { + return 0, errf("BesselY: the argument must be positive, got %g", x) + } + flip, order := 1.0, n + if order < 0 { + if order%2 != 0 { + flip = -1 + } + order = -order + } + var v float64 + switch { + case order == 0: + v = besselY0(x) + case order == 1: + v = besselY1(x) + default: + ym, y := besselY0(x), besselY1(x) + for k := 1; k < order; k++ { + ym, y = y, 2*float64(k)/x*y-ym + } + v = y + } + return flip * v, nil +} + +// besselJSeries evaluates Jₙ(x) by the convergent power series, +// stepping the summand along tₖ = tₖ₋₁·(−(x/2)²)/(k(k+n)). The first +// term (x/2)ⁿ/n! goes through logarithms so a large order never +// overflows the factorial on the way to a small answer. +func besselJSeries(n int, x float64) float64 { + half := 0.5 * x + // The k = 0 term (x/2)ⁿ/n! is positive; the (−1)^k alternation + // enters through the recursion step below. + t := math.Exp(float64(n)*math.Log(half) - LnFactorial(n)) + sum := t + q := half * half + for k := 1; k <= 400; k++ { + t *= -q / (float64(k) * float64(k+n)) + sum += t + if math.Abs(t) <= 1e-17*math.Abs(sum) { + break + } + } + return sum +} + +// besselJUpward climbs J₀ and J₁ from the large-argument asymptotic to +// order n by the upward recurrence Jₖ₊₁ = (2k/x)·Jₖ − Jₖ₋₁. The +// asymptotic seeds carry the absolute scale, so nothing needs +// renormalising, and the climb is the stable direction while the order +// stays at or below the argument: there the parasitic Y component the +// seeds carry stays bounded relative to J, while the downward walk +// would start above the turning point at order x and pay O(x) steps +// for an answer of order n. The caller guarantees x ≥ besselCrossover +// and 0 ≤ n ≤ x. +func besselJUpward(n int, x float64) float64 { + j0, _ := besselAsymptotic(0, x) + if n == 0 { + return j0 + } + j1, _ := besselAsymptotic(1, x) + for k := 1; k < n; k++ { + j0, j1 = j1, 2*float64(k)/x*j1-j0 + } + return j1 +} + +// besselJMiller evaluates Jₙ(x) for |x| above the crossover by the +// downward Miller recurrence, the same stable scheme +// besselInMiller uses for the modified kind: start well above n with +// an arbitrary scale, recurse down through J_{k−1} = (2k/x)·J_k − +// J_{k+1}, then renormalise the arbitrary seed scale against an +// independent J₀, here the asymptotic value the way besselInMiller +// leans on its quadrature I₀. The caller guarantees x ≠ 0. +func besselJMiller(n int, x float64) float64 { + start := n + int(x) + 40 + jp, j := 0.0, 1.0 // J_{k+1}, J_k, seeded at k = start + jn := 0.0 + for k := start; k >= 1; k-- { + if k == n { + jn = j + } + jp, j = j, 2*float64(k)/x*j-jp + if aj := math.Abs(j); aj > 1e200 { + // The unscaled seed grows like k!/x^k on the way down and + // would overflow for high orders whose true value is + // representable; the renormalisation cancels any common + // factor, so rescaling the running pair (and the captured + // order-n value) is exact up to rounding. + jp /= aj + j /= aj + jn /= aj + } + } + if n == 0 { + jn = j + } + anchor, _ := besselAsymptotic(0, x) + return anchor * jn / j +} + +// besselY0 evaluates Y₀(x) for x > 0. Below the crossover the +// Frobenius series in the form the digamma reduction gives, +// +// Y₀ = (2/π)·[(ln(x/2) + γ)·J₀(x) + Σ (−1)^{k+1} H_k (x/2)^{2k}/(k!)²], +// +// above it the asymptotic expansion, whose terms still fall fast +// enough at the crossover to keep twelve significant digits. +func besselY0(x float64) float64 { + if x >= besselCrossover { + _, y := besselAsymptotic(0, x) + return y + } + half := 0.5 * x + q := half * half + u := q // u_k = (x/2)^{2k}/(k!)², starting at k = 1 + h := 1.0 // H_k + sum := u // the k = 1 term carries the + sign + for k := 2; k <= 200; k++ { + u *= q / float64(k*k) + h += 1 / float64(k) + term := h * u + if k%2 == 1 { + sum += term + } else { + sum -= term + } + if math.Abs(term) <= 1e-17*math.Abs(sum) { + break + } + } + return (2 / math.Pi) * ((math.Log(half)+eulerGamma)*besselJSeries(0, x) + sum) +} + +// besselY1 evaluates Y₁(x) for x > 0, the order-one twin of +// besselY0's series branch: +// +// Y₁ = (2/π)(ln(x/2) + γ)·J₁(x) − 2/(πx) +// − (1/π)·Σ (−1)^k (H_k + H_{k+1})·(x/2)^{2k+1}/(k!(k+1)!), +// +// where the 2/(πx) term is the one-entry finite sum of A&S 9.1.11 +// and the ascending series converges for every x, paying only the +// cancellation that caps its usable range at the crossover. +func besselY1(x float64) float64 { + if x >= besselCrossover { + _, y := besselAsymptotic(1, x) + return y + } + half := 0.5 * x + q := half * half + u := half // u_k = (x/2)^{2k+1}/(k!(k+1)!), starting at k = 0 + h := 1.0 // H_{k+1}, starting at H_1 = 1 (H_0 = 0) + sum := u // the k = 0 term: (H_0 + H_1)·u₀ with the + sign + for k := 1; k <= 200; k++ { + hk := h // H_k, before the update below + u *= q / float64(k*(k+1)) + h += 1 / float64(k+1) + term := (hk + h) * u + if k%2 == 1 { + sum -= term + } else { + sum += term + } + if math.Abs(term) <= 1e-17*math.Abs(sum) { + break + } + } + return (2/math.Pi)*((math.Log(half)+eulerGamma)*besselJSeries(1, x)) - + 2/(math.Pi*x) - sum/math.Pi +} + +// besselAsymptotic evaluates the pair J_ν, Y_ν above the crossover +// from the large-argument expansion (Abramowitz & Stegun 9.2.5 through +// 9.2.6): +// +// J_ν ~ sqrt(2/πx)·[cos ω·Σ(−1)^k a_{2k} − sin ω·Σ(−1)^k a_{2k+1}] +// Y_ν ~ sqrt(2/πx)·[sin ω·Σ(−1)^k a_{2k} + cos ω·Σ(−1)^k a_{2k+1}], +// +// with ω = x − νπ/2 − π/4 and aₘ = ∏(μ − (2s−1)²)/(m!·(8x)^m), +// μ = 4ν². The order is a float64: the integer callers pass exact +// integers, whose float64 arithmetic reproduces the int path bit for +// bit, and the real-order BesselJRealOrder passes the fractional seed +// orders. Both kinds share the one coefficient sweep, and the sum +// stops at the optimal truncation where the terms turn around and +// start growing again, which at the crossover still leaves about +// twelve significant digits. +func besselAsymptotic(nu float64, x float64) (j, y float64) { + mu := 4 * nu * nu + omega := besselPhase(nu, x) + eighth := 8 * x + a, prev := 1.0, math.Inf(1) + even, odd := 1.0, 0.0 // Σ(−1)^k a_{2k}, Σ(−1)^k a_{2k+1} + for m := 1; m <= 200; m++ { + a *= (mu - float64((2*m-1)*(2*m-1))) / (float64(m) * eighth) + // m = 2k and m = 2k+1 share the integer k = m/2 and the sign + // (−1)^k of their sum's k-th term. + if (m/2)%2 == 0 { + if m%2 == 0 { + even += a + } else { + odd += a + } + } else { + if m%2 == 0 { + even -= a + } else { + odd -= a + } + } + if math.Abs(a) > prev { + break + } + prev = math.Abs(a) + } + factor := math.Sqrt(2 / (math.Pi * x)) + sin, cos := math.Sin(omega), math.Cos(omega) + j = factor * (cos*even - sin*odd) + y = factor * (sin*even + cos*odd) + return j, y +} + +// 2π split into two float64 parts, the sum of which is 2π to about +// 1e-32: the phase reduction below needs the low part to hold the +// fraction of a large argument. +const ( + twoPiHi = 6.283185307179586 // fl(2π) + twoPiLo = 2.4492935982947064e-16 // 2π − fl(2π) +) + +// besselPhase returns ω = x − νπ/2 − π/4 reduced modulo 2π for the +// large-argument expansion. The reduction is what keeps the phase of a +// large argument meaningful: the plain float64 difference rounds the +// fraction of x away, at x = 1e9 to about 1e-7 absolute, which the +// amplitude √(2/πx) turns into a relative error of 1e-8, and at +// x = 1e12 into 1e-4. The remainder is taken with the two-part 2π: +// x − hi is exact by Sterbenz's lemma, math.FMA gives the exact +// residual of n·2π, and the rest is a handful of flops on quantities +// below π, so the phase keeps its own last ulp while |ω| < 2^53·2π +// (beyond that the float64 grid of x is coarser than a radian and the +// plain difference is as good as anything). The caller passes |x|, +// which is at least the crossover. +func besselPhase(nu float64, x float64) float64 { + nuPi2 := nu * math.Pi / 2 + quarterPi := math.Pi / 4 + omega := x - nuPi2 - quarterPi + n := math.Round(omega / twoPiHi) + if math.Abs(n) >= 1<<53 { + return omega + } + hi := n * twoPiHi + lo := math.FMA(n, twoPiHi, -hi) + return x - hi - lo - n*twoPiLo - nuPi2 - quarterPi +} + +// BesselJRealOrder returns the Bessel function of the first kind of +// real order ν at the real point x, J_ν(x). The order must be ≥ 0 and +// the argument positive; J diverges at the origin for ν > 0 and a +// non-positive argument or a NaN order is an error, not a NaN. +// +// The regimes follow the integer BesselJ's, with the fractional part +// of the order taking the role of the seeds: below the crossover the +// ascending Frobenius series carries the answer, above it the order +// resolves into its integer and fractional parts, the fractional pair +// (ν₀, ν₀+1) is seeded from the large-argument expansion and the +// three-term recurrence runs in its stable direction, upward while the +// order stays at or below the argument and by the downward Miller walk +// anchored on the ν₀ seed when the order passes it. Orders within 1e-8 +// of a non-negative integer are served by the integer algorithm +// itself, which is the continuous limit there: the general route's +// normalisation would cancel against sin(πν) and lose every digit. +func BesselJRealOrder(nu, x float64) (float64, error) { + if math.IsNaN(nu) || math.IsNaN(x) { + return 0, errf("BesselJRealOrder: the order and the argument must be finite, got %g and %g", nu, x) + } + if nu < 0 { + return 0, errf("BesselJRealOrder: the order must be zero or greater, got %g", nu) + } + if x <= 0 { + return 0, errf("BesselJRealOrder: the argument must be positive, got %g", x) + } + if n := math.Round(nu); math.Abs(nu-n) <= 1e-8 { + return BesselJ(int(n), x), nil + } + if x < besselRealCrossover { + return besselJSeriesReal(nu, x), nil + } + whole := math.Floor(nu) + frac := nu - whole + if nu <= x { + // The climb: seed the fractional pair from the expansion and + // run the upward recurrence, the stable direction while the + // order stays at or below the argument. whole = 0 means the + // answer is the first seed itself. + jm, _ := besselAsymptotic(frac, x) + if whole == 0 { + return jm, nil + } + j, _ := besselAsymptotic(frac+1, x) + for k := 1; k < int(whole); k++ { + jm, j = j, 2*(frac+float64(k))/x*j-jm + } + return j, nil + } + // The downward Miller walk through one fractional residue class, + // anchored on the fractional seed the same way besselJMiller + // anchors its integer walk on J₀. The walk counts integer steps + // from the fractional floor instead of testing k against nu and + // frac: start descends in floats, and a capture missed by one ulp + // of the running k would answer from an unseeded slot. + steps := int(math.Floor(x)) + int(whole) + 40 + jp, j := 0.0, 1.0 // J_{k+1}, J_k, seeded at k = frac + steps + uNu := 0.0 + uFrac := 0.0 + for i := steps; i >= 0; i-- { + k := frac + float64(i) + if i == int(whole) { + uNu = j + } + if i == 0 { + uFrac = j + } + jp, j = j, 2*k/x*j-jp + if aj := math.Abs(j); aj > 1e200 { + jp /= aj + j /= aj + uNu /= aj + uFrac /= aj + } + } + anchor, _ := besselAsymptotic(frac, x) + return anchor * uNu / uFrac, nil +} + +// lnGammaReal returns lnΓ(z) for z > 0 at one point, the scalar +// companion the real-order series needs; the sign of Γ is positive on +// that domain, so only the logarithm comes back. +func lnGammaReal(z float64) float64 { + lg, _ := math.Lgamma(z) + return lg +} + +// besselRealCrossover splits the real-order series and recurrence +// regimes, lower than the integer one because the two sides' quality +// decides differently here: the series pays the same cancellation +// (about x·ln10/2 digits at x = 12, still leaving better than nine), +// while the expansion the seeds come from truncates to full precision +// for the half-integer orders and to ten-plus digits for the small +// fractional seeds the climb and the Miller walk lean on. +const besselRealCrossover = 12.0 + +// besselJSeriesReal evaluates J_ν(x) by the convergent Frobenius +// series for real ν, stepping the summand along +// tₖ = tₖ₋₁·(−(x/2)²)/(k(k+ν)). It is the real-order companion of +// besselJSeries, kept separate so the integer series keeps its +// exact LnFactorial opening and its recorded bits. +func besselJSeriesReal(nu, x float64) float64 { + half := 0.5 * x + // The opening term never overflows below the crossover: with + // half < 7.5 the exponent ν·log(half) − lnΓ(ν+1) turns downward + // past ν ≈ 20 and stays negative. + t := math.Exp(nu*math.Log(half) - lnGammaReal(nu+1)) + sum := t + q := half * half + for k := 1; k <= 400; k++ { + t *= -q / (float64(k) * (float64(k) + nu)) + sum += t + if math.Abs(t) <= 1e-17*math.Abs(sum) { + break + } + } + return sum +} diff --git a/internal/core/special2_test.go b/internal/core/special2_test.go new file mode 100644 index 0000000..4cd696b --- /dev/null +++ b/internal/core/special2_test.go @@ -0,0 +1,144 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "strings" + "testing" +) + +// TestBesselJTabulated pins J against independently computed values +// (mpmath, 30 significant digits) at points covering the power-series +// branch, the downward Miller branch and both parity laws. +func TestBesselJTabulated(t *testing.T) { + cases := []struct { + name string + n int + x float64 + want float64 + }{ + {"J_0(0)", 0, 0, 1}, + {"J_1(0)", 1, 0, 0}, + {"J_0(1)", 0, 1, 0.765197686557966551}, + {"J_1(1)", 1, 1, 0.440050585744933516}, + {"J_2(1)", 2, 1, 0.11490348493190048}, + {"J_5(1)", 5, 1, 0.000249757730211234431}, + {"J_7(2)", 7, 2, 0.000174944074868274169}, + {"J_0(5)", 0, 5, -0.177596771314338304}, + {"J_3(4)", 3, 4, 0.43017147387562194}, + {"J_2(10)", 2, 10, 0.254630313685120623}, + {"J_0(15.5)", 0, 15.5, -0.109230650900050168}, + {"J_2(15.5)", 2, 15.5, 0.130806545138985284}, + {"J_0(25)", 0, 25, 0.0962667832759581162}, + {"J_5(25)", 5, 25, -0.0660079953984229934}, + {"J_1(30)", 1, 30, -0.118751062616622937}, + {"J_1(-1)", 1, -1, -0.440050585744933516}, + {"J_2(-1)", 2, -1, 0.11490348493190048}, + {"J_-1(2)", -1, 2, -0.576724807756873387}, + {"J_-3(-2)", -3, -2, 0.128943249474402051}, + } + for _, c := range cases { + if got := BesselJ(c.n, c.x); math.Abs(got-c.want) > 1e-9*math.Abs(c.want) { + t.Errorf("%s = %.16g, want %.16g", c.name, got, c.want) + } + } +} + +// TestBesselYTabulated pins Y against independently computed values +// (mpmath, 30 significant digits): the series seeds at small x, the +// asymptotic seeds above the crossover and the upward recurrence +// between them, plus the negative-order parity law. +func TestBesselYTabulated(t *testing.T) { + cases := []struct { + name string + n int + x float64 + want float64 + }{ + {"Y_0(0.05)", 0, 0.05, -1.97931100081720967}, + {"Y_3(0.5)", 3, 0.5, -42.0594943047238827}, + {"Y_0(1)", 0, 1, 0.088256964215676958}, + {"Y_1(1)", 1, 1, -0.781212821300288717}, + {"Y_2(1)", 2, 1, -1.65068260681625439}, + {"Y_-2(1)", -2, 1, -1.65068260681625439}, + {"Y_3(3)", 3, 3, -0.538541616105031618}, + {"Y_5(2)", 5, 2, -9.93598912848197498}, + {"Y_7(2)", 7, 2, -271.54802536799367}, + {"Y_10(0.7)", 10, 0.7, -4244719426.07038669}, + {"Y_0(10)", 0, 10, 0.0556711672835993914}, + {"Y_0(15.5)", 0, 15.5, 0.170644911229434617}, + {"Y_2(15.5)", 2, 15.5, -0.155833796066422704}, + {"Y_5(15.5)", 5, 15.5, 0.204463657248615881}, + {"Y_0(20)", 0, 20, 0.0626405968093838312}, + {"Y_4(25)", 4, 25, -0.0910410709900928362}, + {"Y_1(30)", 1, 30, 0.0844255706617472349}, + {"Y_-1(3)", -1, 3, -0.324674424791799978}, + } + for _, c := range cases { + got, err := BesselY(c.n, c.x) + if err != nil { + t.Errorf("%s: %v", c.name, err) + continue + } + if math.Abs(got-c.want) > 1e-9*math.Abs(c.want) { + t.Errorf("%s = %.16g, want %.16g", c.name, got, c.want) + } + } +} + +// TestBesselJRecurrence checks the three-term recurrence +// J_{ν−1} + J_{ν+1} = 2ν/x·J_ν (ν = 4) across both evaluation +// branches, which any branch inconsistency would break. +func TestBesselJRecurrence(t *testing.T) { + for _, x := range []float64{1.5, 8, 15.5, 30} { + nu := 4 + jm1 := BesselJ(nu-1, x) + j0 := BesselJ(nu, x) + jp1 := BesselJ(nu+1, x) + got := jm1 + jp1 + want := 2 * float64(nu) / x * j0 + if math.Abs(got-want) > 1e-9*math.Abs(want) { + t.Errorf("x=%v: J_%d + J_%d = %.16g, want %.16g", x, nu-1, nu+1, got, want) + } + } +} + +// TestBesselYRecurrence checks the same recurrence for the second +// kind, tying the series seeds to the recurrence-climbed orders at +// both small and large arguments. +func TestBesselYRecurrence(t *testing.T) { + for _, x := range []float64{0.7, 2, 15.5} { + nu := 4 + at := func(k int) float64 { + v, err := BesselY(k, x) + if err != nil { + t.Fatalf("BesselY(%d, %v): %v", k, x, err) + } + return v + } + got := at(nu-1) + at(nu+1) + want := 2 * float64(nu) / x * at(nu) + if math.Abs(got-want) > 1e-9*math.Abs(want) { + t.Errorf("x=%v: Y_%d + Y_%d = %.16g, want %.16g", x, nu-1, nu+1, got, want) + } + } +} + +// TestBesselYRejects pins the domain contract: Yₙ is defined for +// x > 0 only and reports the violation as an error with the package +// prefix, never as a NaN. +func TestBesselYRejects(t *testing.T) { + for _, x := range []float64{0, -1, -1e-300, math.NaN()} { + v, err := BesselY(2, x) + if err == nil { + t.Errorf("BesselY(2, %g): expected an error, got %v", x, v) + } else if !strings.Contains(err.Error(), "tensor: BesselY") { + t.Errorf("BesselY(2, %g): error %q lacks the prefixed name", x, err) + } + } + if _, err := BesselY(2, 1); err != nil { + t.Errorf("BesselY(2, 1): %v", err) + } +} diff --git a/internal/core/special_edge_accuracy_test.go b/internal/core/special_edge_accuracy_test.go new file mode 100644 index 0000000..9d59ce8 --- /dev/null +++ b/internal/core/special_edge_accuracy_test.go @@ -0,0 +1,511 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "math/big" + "testing" +) + +// Edge-argument validation of the special-function surface against +// extended-precision references built on the helpers of the Bessel +// accuracy test: convergent series, asymptotics and recurrences summed +// in math/big far past the float64 grid. The float64 implementations +// under test share no arithmetic with the references. + +var ( + specSqrtPi = new(big.Float).SetPrec(besselRefPrec).Sqrt(bfPi()) + specTwoPi = bfMul(bf(2), bfPi()) +) + +// specSinCos returns sin(x), cos(x) for an exact big.Float argument, +// reduced mod 2π before the Taylor run. +func specSinCos(x *big.Float) (s, c *big.Float) { + n := new(big.Float).SetPrec(besselRefPrec).Quo(x, specTwoPi) + ni, _ := n.Int(nil) + r := bfSub(x, bfMul(new(big.Float).SetPrec(besselRefPrec).SetInt(ni), specTwoPi)) + halfPi := bfQuo(bfPi(), bf(2)) + if r.Cmp(halfPi) > 0 { + r = bfSub(r, specTwoPi) + } else if r.Cmp(bfSub(bf(0), halfPi)) < 0 { + r = bfAdd(r, specTwoPi) + } + if r.Sign() == 0 { + return bf(0), bf(1) + } + r2 := bfMul(r, r) + ts, tc := new(big.Float).SetPrec(besselRefPrec).Set(r), bf(1) + sumS, sumC := new(big.Float).SetPrec(besselRefPrec).Set(r), bf(1) + for j := int64(1); ; j++ { + ts = bfSub(bf(0), bfQuo(bfMul(ts, r2), bfFromInt((2*j)*(2*j+1)))) + sumS = bfAdd(sumS, ts) + tc = bfSub(bf(0), bfQuo(bfMul(tc, r2), bfFromInt((2*j-1)*(2*j)))) + sumC = bfAdd(sumC, tc) + if ts.Sign() == 0 || ts.MantExp(nil) < -int(besselRefPrec)-10 { + break + } + } + return sumS, sumC +} + +// specEulerGamma builds the Euler–Mascheroni constant from the +// harmonic-number asymptotics: γ = H_N − ln N − 1/(2N) + Σ B₂ₙ/(2n·N²ⁿ) +// with N = 100 and the exact Bernoulli coefficients the Stirling helper +// already carries. +func specEulerGamma() *big.Float { + const capN = 100 + h := bf(0) + for k := int64(1); k <= capN; k++ { + h = bfAdd(h, bfQuo(bf(1), bfFromInt(k))) + } + g := bfSub(h, bfLn(bf(capN))) + g = bfSub(g, bfQuo(bf(1), bfFromInt(2*capN))) + pow := bf(1) + sqN := bfMul(bf(capN), bf(capN)) + for n := 1; n <= len(besselRefStirling); n++ { + pow = bfMul(pow, sqN) + coef := new(big.Float).SetPrec(besselRefPrec).SetRat(besselRefStirling[n-1]) + term := bfQuo(bfMul(coef, bfFromInt(int64(2*n-1))), pow) + g = bfAdd(g, term) + } + return g +} + +var specGammaConst = specEulerGamma() + +// specErf evaluates erf(x) by the convergent series +// erf(x) = 2/√π · Σ (−1)^k x^{2k+1}/(k!(2k+1)). +func specErf(x float64) *big.Float { + xf := bf(x) + x2 := bfMul(xf, xf) + term := xf + sum := bf(0) + sum.Add(sum, term) + for k := int64(1); ; k++ { + term = bfSub(bf(0), bfQuo(bfMul(term, bfMul(x2, bfFromInt(2*k-1))), bfMul(bfFromInt(k), bfFromInt(2*k+1)))) + sum.Add(sum, term) + if term.Sign() == 0 || term.MantExp(nil) < sum.MantExp(nil)-int(besselRefPrec)-20 { + break + } + } + return bfQuo(bfMul(bf(2), sum), specSqrtPi) +} + +// TestErfErfcEdgeAccuracy holds the error-function pair against the +// extended series: the lower tail where erfc rounds through 1−erf, the +// mid range where both switch formulations, and the far tail where +// erfc(x) sits tens of orders below one. +func TestErfErfcEdgeAccuracy(t *testing.T) { + for _, x := range []float64{0.05, 0.5, 1, 2, 3, 4, 5, 5.5, 6} { + ref, _ := specErf(x).Float64() + got := math.Erf(x) + if d := math.Abs(got-ref) / math.Abs(ref); d > 1e-14 { + t.Fatalf("erf(%g) = %.17g, want %.17g (relative %.3g)", x, got, ref, d) + } + } + for _, x := range []float64{0.05, 0.5, 1, 2, 3, 4, 5, 6, 8, 10} { + e := specErf(x) + ref, _ := bfSub(bf(1), e).Float64() + got := math.Erfc(x) + // The tail value is compared absolutely against its own scale: + // erfc(10) ≈ 2e−45. + d := math.Abs(got - ref) + if d > math.Max(1e-15*math.Abs(ref), 1e-140) { + t.Fatalf("erfc(%g) = %.17g, want %.17g (absolute %.3g)", x, got, ref, d) + } + } +} + +// TestGammaLnGammaEdgeAccuracy holds the gamma pair against the Stirling +// reference: the pole neighbourhood, the unit neighbourhood where the +// reflection path of the standard library is least even, and the large +// arguments near the overflow frontier. +func TestGammaLnGammaEdgeAccuracy(t *testing.T) { + for _, x := range []float64{0.05, 0.5, 1, 1.5, 2, 5.5, 8, 20, 100, 170} { + want, _ := bfExp(lnGammaBig(x)).Float64() + got := math.Gamma(x) + if d := math.Abs(got-want) / math.Abs(want); d > 1e-14 { + t.Fatalf("Gamma(%g) = %.17g, want %.17g (relative %.3g)", x, got, want, d) + } + lg := lnGammaBig(x) + gl, _ := lg.Float64() + lgot, _ := math.Lgamma(x) + if d := math.Abs(lgot - gl); d > 1e-11*math.Max(1, math.Abs(gl)) { + t.Fatalf("Lgamma(%g) = %.17g, want %.17g (absolute %.3g)", x, lgot, gl, d) + } + } + // The sub-unit arguments: lnΓ diverges logarithmically while Γ + // explodes, both covered by the same reference. + for _, x := range []float64{1e-8, 1e-3} { + gl, _ := lnGammaBig(x).Float64() + lgot, _ := math.Lgamma(x) + if d := math.Abs(lgot - gl); d > 1e-12*math.Max(1, math.Abs(gl)) { + t.Fatalf("Lgamma(%g) = %.17g, want %.17g (absolute %.3g)", x, lgot, gl, d) + } + } +} + +// specE1Series is the small-argument reference, +// E1(x) = −γ − ln x − Σ (−x)^k/(k·k!), in extended precision. +func specE1Series(x float64) *big.Float { + term := bfSub(bf(0), bf(x)) // (−x)^k/k! at k = 1 + sum := bf(0) + for k := int64(1); ; k++ { + add := bfQuo(term, bfFromInt(k)) + sum = bfAdd(sum, add) + if add.Sign() == 0 || add.MantExp(nil) < sum.MantExp(nil)-int(besselRefPrec)-20 { + break + } + // (−x)^{k+1}/(k+1)! = (−x)^k/k! · (−x)/(k+1). + term = bfQuo(bfSub(bf(0), bfMul(term, bf(x))), bfFromInt(k+1)) + } + return bfSub(bfSub(bf(0), specGammaConst), bfAdd(bfLn(bf(x)), sum)) +} + +// specE1CF is the large-argument reference through the Gauss continued +// fraction, E1 = e^{−x}/f, evaluated bottom-up in extended precision. +func specE1CF(x float64) *big.Float { + const depth = 500 + f := bfAdd(bf(x), bfFromInt(2*depth+1)) + for i := depth; i >= 1; i-- { + f = bfSub(bfAdd(bf(x), bfFromInt(int64(2*i-1))), bfQuo(bfFromInt(int64(i)*int64(i)), f)) + } + return bfQuo(bfExp(bfSub(bf(0), bf(x))), f) +} + +// TestExpIntegralE1EdgeAccuracy holds E1 against the extended series +// below two and the extended continued fraction above, with the branch +// point itself covered from both sides. +func TestExpIntegralE1EdgeAccuracy(t *testing.T) { + for _, c := range []struct { + x float64 + ref func(float64) *big.Float + tol float64 + note string + }{ + {0.05, specE1Series, 1e-14, "series"}, + {0.5, specE1Series, 1e-14, "series"}, + {1, specE1Series, 1e-14, "series"}, + {1.999, specE1Series, 1e-14, "series"}, + {2.001, specE1CF, 1e-14, "fraction"}, + {3, specE1CF, 1e-14, "fraction"}, + {10, specE1CF, 1e-15, "fraction"}, + {50, specE1CF, 1e-15, "fraction"}, + {300, specE1CF, 1e-15, "fraction"}, + {700, specE1CF, 1e-14, "fraction"}, + } { + want, _ := c.ref(c.x).Float64() + got := expIntegralE1(c.x) + if d := math.Abs(got-want) / math.Abs(want); d > c.tol { + t.Fatalf("E1(%g) = %.17g, want %.17g (relative %.3g, %s reference)", c.x, got, want, d, c.note) + } + } + // The branch point from both sides: each side holds against its own + // extended reference at the same point. + for _, c := range []struct { + x float64 + ref func(float64) *big.Float + }{ + {1.9999999999, specE1Series}, + {2.0000000001, specE1CF}, + } { + want, _ := c.ref(c.x).Float64() + got := expIntegralE1(c.x) + if d := math.Abs(got-want) / math.Abs(want); d > 1e-14 { + t.Fatalf("E1(%v): %.17g against %.17g (relative %.3g)", c.x, got, want, d) + } + } +} + +// TestExpIntegralEiEdgeAccuracy holds Ei against the extended +// logarithmic series on both sides of the asymptotic switch at forty. +func TestExpIntegralEiEdgeAccuracy(t *testing.T) { + // Ei(x) = γ + ln x + Σ x^k/(k·k!), all terms positive. + ei := func(x float64) *big.Float { + term := bf(1) + sum := bf(0) + for k := int64(1); ; k++ { + term = bfQuo(bfMul(term, bf(x)), bfFromInt(k)) + add := bfQuo(term, bfFromInt(k)) + sum = bfAdd(sum, add) + if add.MantExp(nil) < sum.MantExp(nil)-int(besselRefPrec)-20 { + break + } + } + return bfAdd(bfAdd(specGammaConst, bfLn(bf(x))), sum) + } + for _, x := range []float64{0.5, 1, 5, 20, 39.9, 40.1, 41, 50, 100, 300} { + a, err := FromFloats([]float64{x}, 1) + if err != nil { + t.Fatal(err) + } + got, err := ExpIntegralEi(a) + if err != nil { + t.Fatal(err) + } + want, _ := ei(x).Float64() + // The convergent-series side meets the reference at the rounding + // floor; the asymptotic side truncates at its smallest term and + // carries a measured few parts in 1e14. + tol := 1e-14 + if x >= 40 { + tol = 1e-13 + } + if d := math.Abs(got.FloatAt(0)-want) / math.Abs(want); d > tol { + t.Fatalf("Ei(%g) = %.17g, want %.17g (relative %.3g)", x, got.FloatAt(0), want, d) + } + } +} + +// TestSincEdgeAccuracy holds the normalised sinc against the extended +// sin: at large arguments the float64 phase π·x carries an unavoidable +// rounding of about x·π·1e−16, which the bound admits, everything +// smaller must be exact. +func TestSincEdgeAccuracy(t *testing.T) { + for _, x := range []float64{1e-9, 1e-5, 0.25, 1, 3, 10, 100, 1000, 5e4} { + s, _ := specSinCos(bfMul(bfPi(), bf(x))) + den := bfMul(bfPi(), bf(x)) + want, _ := bfQuo(s, den).Float64() + got := math.Sin(math.Pi*x) / (math.Pi * x) + phase := math.Pi * math.Abs(x) * 1.1e-16 + if d := math.Abs(got - want); d > math.Max(1e-17, phase) { + t.Fatalf("sinc(%g) = %.17g, want %.17g (absolute %.3g, phase budget %.3g)", + x, got, want, d, phase) + } + } + if got := SincOne(); got != 1 { + t.Fatalf("sinc(0) = %.17g, want 1", got) + } +} + +// SincOne evaluates the package sinc at the origin through the array +// surface. +func SincOne() float64 { + a, err := FromFloats([]float64{0}, 1) + if err != nil { + return math.NaN() + } + out, err := Sinc(a) + if err != nil { + return math.NaN() + } + return out.FloatAt(0) +} + +// TestDigammaTrigammaEdgeAccuracy holds the polygamma pair against the +// lifted asymptotics in extended precision, positive arguments and the +// reflection branch on the negative side. +func TestDigammaTrigammaEdgeAccuracy(t *testing.T) { + // ψ(x): lift by ψ(x+1) = ψ(x) + 1/x, then the asymptotic series. + digammaBig := func(x float64) *big.Float { + acc := bf(0) + z := bf(x) + for z.MantExp(nil) < 7 { + acc = bfSub(acc, bfQuo(bf(1), z)) + z = bfAdd(z, bf(1)) + } + inv := bfQuo(bf(1), z) + res := bfSub(bfLn(z), bfQuo(inv, bf(2))) + pow := bfMul(z, z) + for n := 1; n <= len(besselRefStirling); n++ { + coef := new(big.Float).SetPrec(besselRefPrec).SetRat(besselRefStirling[n-1]) + res = bfSub(res, bfQuo(bfMul(coef, bfFromInt(int64(2*n-1))), pow)) + pow = bfMul(pow, bfMul(z, z)) + if coef.MantExp(nil) < -200 { + break + } + } + return bfAdd(acc, res) + } + // ψ′(x): the same lift with 1/x² terms, ψ′ ~ 1/x + 1/(2x²) + Σ B₂ₙ/x^{2n+1}. + trigammaBig := func(x float64) *big.Float { + acc := bf(0) + z := bf(x) + for z.MantExp(nil) < 7 { + acc = bfAdd(acc, bfQuo(bf(1), bfMul(z, z))) + z = bfAdd(z, bf(1)) + } + inv := bfQuo(bf(1), z) + res := bfAdd(inv, bfQuo(bfMul(inv, inv), bf(2))) + pow := bfMul(bfMul(z, z), z) + for n := 1; n <= len(besselRefStirling); n++ { + coef := new(big.Float).SetPrec(besselRefPrec).SetRat(besselRefStirling[n-1]) + // B₂ₙ/x^{2n+1} = cₙ·(2n)(2n−1)/x^{2n+1}. + factor := int64(2 * n * (2*n - 1)) + res = bfAdd(res, bfQuo(bfMul(coef, bfFromInt(factor)), pow)) + pow = bfMul(pow, bfMul(z, z)) + if coef.MantExp(nil) < -200 { + break + } + } + return bfAdd(acc, res) + } + positives := []float64{0.05, 0.5, 1, 2, 3.5, 19.9, 20.1, 50} + for _, x := range positives { + want, _ := digammaBig(x).Float64() + got := digamma(x) + if d := math.Abs(got-want) / math.Max(1, math.Abs(want)); d > 1e-14 { + t.Fatalf("digamma(%g) = %.17g, want %.17g (relative %.3g)", x, got, want, d) + } + twant, _ := trigammaBig(x).Float64() + tgot := trigamma(x) + if d := math.Abs(tgot-twant) / math.Max(1, math.Abs(twant)); d > 1e-14 { + t.Fatalf("trigamma(%g) = %.17g, want %.17g (relative %.3g)", x, tgot, twant, d) + } + } + // The reflection branch: ψ(x) = ψ(1−x) − π·cot(πx) and + // ψ′(x) = π²/sin²(πx) − ψ′(1−x), evaluated in extended precision. + for _, x := range []float64{-0.5, -1.5, -2.05, -7.3} { + sp, cp := specSinCos(bfMul(bfPi(), bf(x))) + other := digammaBig(1 - x) + want := bfSub(other, bfQuo(bfMul(bfPi(), cp), sp)) + wv, _ := want.Float64() + if got := digamma(x); math.Abs(got-wv) > 1e-13*math.Max(1, math.Abs(wv)) { + t.Fatalf("digamma(%g) = %.17g, want %.17g", x, got, wv) + } + tp, _ := specSinCos(bfMul(bfPi(), bf(x))) + pi2 := bfMul(bfPi(), bfPi()) + twant := bfSub(bfQuo(pi2, bfMul(tp, tp)), trigammaBig(1-x)) + twv, _ := twant.Float64() + if got := trigamma(x); math.Abs(got-twv) > 1e-13*math.Max(1, math.Abs(twv)) { + t.Fatalf("trigamma(%g) = %.17g, want %.17g", x, got, twv) + } + } +} + +// TestFresnelEdgeAccuracy holds the Fresnel pair against the extended +// series in the cancellation band and against the auxiliary asymptotics +// in the far band, both in extended precision. +func TestFresnelEdgeAccuracy(t *testing.T) { + // Series: C(x) = Σ (−1)^k (π/2)^{2k} x^{4k+1}/((2k)!(4k+1)). + fres := func(x float64, sine bool) *big.Float { + xf, halfPi := bf(x), bfQuo(bfPi(), bf(2)) + hh := bfMul(halfPi, halfPi) + x4 := bfMul(bfMul(xf, xf), bfMul(xf, xf)) + var term, sum *big.Float + if sine { + term = bfQuo(bfMul(bfMul(halfPi, xf), bfMul(xf, xf)), bfFromInt(3)) + } else { + term = bf(x) + } + sum = bf(0) + sum.Add(sum, term) + for k := int64(0); ; k++ { + an, d1, d2, d3 := 4*k+1, int64(2*k+1), int64(2*k+2), int64(4*k+5) + if sine { + an, d1, d2, d3 = 4*k+3, int64(2*k+2), int64(2*k+3), int64(4*k+7) + } + next := bfSub(bf(0), bfQuo(bfMul(bfMul(bfMul(term, hh), x4), bfFromInt(an)), + bfFromInt(d1*d2*d3))) + sum.Add(sum, next) + if next.Sign() == 0 || next.MantExp(nil) < sum.MantExp(nil)-int(besselRefPrec)-20 { + break + } + term = next + } + return sum + } + for _, x := range []float64{0.5, 2, 4, 8} { + wantC, _ := fres(x, false).Float64() + wantS, _ := fres(x, true).Float64() + if got := fresnelC(x); math.Abs(got-wantC) > 1e-14*math.Max(1, math.Abs(wantC)) { + t.Fatalf("FresnelC(%g) = %.17g, want %.17g", x, got, wantC) + } + if got := fresnelS(x); math.Abs(got-wantS) > 1e-14*math.Max(1, math.Abs(wantS)) { + t.Fatalf("FresnelS(%g) = %.17g, want %.17g", x, got, wantS) + } + } + // Far band: the auxiliary asymptotics in extended precision. + fresAsym := func(x float64, sine bool) *big.Float { + t := bfQuo(bf(1), bfMul(bfPi(), bfMul(bf(x), bf(x)))) + t2 := bfMul(t, t) + sumF, sumG := bf(1), new(big.Float).SetPrec(besselRefPrec).Set(t) + cf, cg := bf(1), new(big.Float).SetPrec(besselRefPrec).Set(t) + for k := range int64(40) { + cf = bfSub(bf(0), bfMul(bfMul(cf, bfFromInt(4*k+1)), bfMul(bfFromInt(4*k+3), t2))) + cg = bfSub(bf(0), bfMul(bfMul(cg, bfFromInt(4*k+3)), bfMul(bfFromInt(4*k+5), t2))) + sumF = bfAdd(sumF, cf) + sumG = bfAdd(sumG, cg) + if cf.MantExp(nil) < -int(besselRefPrec)-10 && cg.MantExp(nil) < -int(besselRefPrec)-10 { + break + } + } + inv := bfQuo(bf(1), bfMul(bfPi(), bf(x))) + f, g := bfMul(inv, sumF), bfMul(inv, sumG) + u := bfQuo(bfMul(bfPi(), bfMul(bf(x), bf(x))), bf(2)) + su, cu := specSinCos(u) + if sine { + return bfSub(bfSub(bfQuo(bf(1), bf(2)), bfMul(f, cu)), bfMul(g, su)) + } + return bfAdd(bfSub(bfQuo(bf(1), bf(2)), bfMul(g, cu)), bfMul(f, su)) + } + for _, x := range []float64{12, 20, 50} { + wantC, _ := fresAsym(x, false).Float64() + wantS, _ := fresAsym(x, true).Float64() + if got := fresnelC(x); math.Abs(got-wantC) > 1e-14 { + t.Fatalf("FresnelC(%g) = %.17g, want %.17g", x, got, wantC) + } + if got := fresnelS(x); math.Abs(got-wantS) > 1e-14 { + t.Fatalf("FresnelS(%g) = %.17g, want %.17g", x, got, wantS) + } + } +} + +// TestAiryEdgeAccuracy holds the Airy pair against the extended +// Maclaurin series across the oscillatory negative band and the +// cancellation-heavy positive band. +func TestAiryEdgeAccuracy(t *testing.T) { + c1, _, _ := big.ParseFloat(airyC1Digits, 10, besselRefPrec, big.ToNearestEven) + c2, _, _ := big.ParseFloat(airyC2Digits, 10, besselRefPrec, big.ToNearestEven) + root3 := new(big.Float).SetPrec(besselRefPrec).Sqrt(bf(3)) + airy := func(x float64) (ai, bi float64) { + xf := bf(x) + x3 := bfMul(bfMul(xf, xf), xf) + tf, tg := bf(1), xf + sumF, sumG := bf(1), xf + for k := int64(0); ; k++ { + tf = bfMul(tf, bfQuo(bfMul(bfFromInt(3), x3), bfFromInt((3*k+1)*(3*k+2)*(3*k+3)))) + tf = bfMul(tf, bfAdd(bfFromInt(k), bfQuo(bf(1), bfFromInt(3)))) + tg = bfMul(tg, bfQuo(bfMul(bfFromInt(3), x3), bfFromInt((3*k+2)*(3*k+3)*(3*k+4)))) + tg = bfMul(tg, bfAdd(bfFromInt(k), bfQuo(bf(2), bfFromInt(3)))) + sumF = bfAdd(sumF, tf) + sumG = bfAdd(sumG, tg) + if tf.MantExp(nil) < sumF.MantExp(nil)-int(besselRefPrec)-20 { + break + } + } + aiv := bfSub(bfMul(c1, sumF), bfMul(c2, sumG)) + biv := bfMul(root3, bfAdd(bfMul(c1, sumF), bfMul(c2, sumG))) + ai, _ = aiv.Float64() + bi, _ = biv.Float64() + return ai, bi + } + for _, x := range []float64{-8, -6.5, -4, -2, 0.5, 1.9, 2, 8} { + wantAI, wantBI := airy(x) + a, err := FromFloats([]float64{x}, 1) + if err != nil { + t.Fatal(err) + } + gotAI, gotBI, err := Airy(a) + if err != nil { + t.Fatal(err) + } + // On the oscillatory negative band both reconstructions cancel + // by up to e^{2|x|^{3/2}/3}, which caps the float64 series at a + // few parts in 1e11 at the |x| = 8 edge; the positive side runs + // through the extended-precision path and holds the rounding + // floor. + tol := 1e-13 + if x < 0 { + tol = 1e-10 + } + if d := math.Abs(gotAI.FloatAt(0)-wantAI) / math.Max(1e-300, math.Abs(wantAI)); d > tol { + t.Fatalf("Ai(%g) = %.17g, want %.17g (relative %.3g)", x, gotAI.FloatAt(0), wantAI, d) + } + if d := math.Abs(gotBI.FloatAt(0)-wantBI) / math.Abs(wantBI); d > tol { + t.Fatalf("Bi(%g) = %.17g, want %.17g (relative %.3g)", x, gotBI.FloatAt(0), wantBI, d) + } + } +} diff --git a/internal/core/special_test.go b/internal/core/special_test.go new file mode 100644 index 0000000..edf020e --- /dev/null +++ b/internal/core/special_test.go @@ -0,0 +1,383 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +func mustFloats(t *testing.T, vals []float64, shape ...int) *Array { + t.Helper() + if len(shape) == 0 { + shape = []int{len(vals)} + } + a, err := FromFloats(vals, shape...) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestGammaFamily checks Γ, ln|Γ| and B against exact values. +func TestGammaFamily(t *testing.T) { + x := mustFloats(t, []float64{1, 5, 0.5, -0.5}) + g, err := Gamma(x) + if err != nil { + t.Fatalf("Gamma: %v", err) + } + // Γ(1)=1, Γ(5)=24, Γ(½)=√π, Γ(−½)=−2√π. + want := []float64{1, 24, math.SqrtPi, -2 * math.SqrtPi} + for i := range want { + if math.Abs(g.FloatAt(i)-want[i]) > 1e-12*(1+math.Abs(want[i])) { + t.Fatalf("Gamma(%v) = %v, want %v", x.FloatAt(i), g.FloatAt(i), want[i]) + } + } + + lg, err := LnGamma(x) + if err != nil { + t.Fatalf("LnGamma: %v", err) + } + if math.Abs(lg.FloatAt(1)-math.Log(24)) > 1e-12 { + t.Fatalf("LnGamma(5) = %v, want ln 24", lg.FloatAt(1)) + } + + bx := mustFloats(t, []float64{1, 2}) + by := mustFloats(t, []float64{1, 3}) + b, err := Beta(bx, by) + if err != nil { + t.Fatalf("Beta: %v", err) + } + // B(1,1) = 1 and B(2,3) = Γ(2)Γ(3)/Γ(5) = 2/24 = 1/12. + if math.Abs(b.FloatAt(0)-1) > 1e-12 || math.Abs(b.FloatAt(1)-1.0/12) > 1e-12 { + t.Fatalf("Beta = [%v, %v], want [1, 1/12]", b.FloatAt(0), b.FloatAt(1)) + } + if _, err := Beta(mustFloats(t, []float64{1}), by); err == nil { + t.Fatal("expected a shape-mismatch error for Beta") + } +} + +// TestErrorFunctions checks the error function family through the +// defining identity erf(x) + erfc(x) = 1 and known values. +func TestErrorFunctions(t *testing.T) { + x := mustFloats(t, []float64{0, 0.5, 1, 3}) + e, err := Erf(x) + if err != nil { + t.Fatalf("Erf: %v", err) + } + c, err := Erfc(x) + if err != nil { + t.Fatalf("Erfc: %v", err) + } + for i := range x.Len() { + if math.Abs(e.FloatAt(i)+c.FloatAt(i)-1) > 1e-12 { + t.Fatalf("erf(%v)+erfc(%v) = %v, want 1", + x.FloatAt(i), x.FloatAt(i), e.FloatAt(i)+c.FloatAt(i)) + } + } + if e.FloatAt(0) != 0 { + t.Fatalf("erf(0) = %v, want 0", e.FloatAt(0)) + } + if math.Abs(e.FloatAt(2)-0.8427007929497149) > 1e-12 { + t.Fatalf("erf(1) = %v, want 0.8427007929497149", e.FloatAt(2)) + } +} + +// TestSinc checks the normalised sinc: value 1 at the origin, zeros at +// non-zero integers, 2/π at a half. +func TestSinc(t *testing.T) { + x := mustFloats(t, []float64{0, 0.5, 1, 2}) + s, err := Sinc(x) + if err != nil { + t.Fatalf("Sinc: %v", err) + } + if s.FloatAt(0) != 1 { + t.Fatalf("Sinc(0) = %v, want 1", s.FloatAt(0)) + } + if math.Abs(s.FloatAt(1)-2/math.Pi) > 1e-12 { + t.Fatalf("Sinc(0.5) = %v, want 2/π", s.FloatAt(1)) + } + for _, i := range []int{2, 3} { + if math.Abs(s.FloatAt(i)) > 1e-15 { + t.Fatalf("Sinc(%v) = %v, want 0", x.FloatAt(i), s.FloatAt(i)) + } + } +} + +// TestLegendrePolynomials checks P_l against closed forms and the +// Bonnet recurrence across degrees. +func TestLegendrePolynomials(t *testing.T) { + x := mustFloats(t, []float64{-0.8, -0.3, 0, 0.42, 0.9}) + closed := []struct { + l int + want func(v float64) float64 + }{ + {0, func(v float64) float64 { return 1 }}, + {1, func(v float64) float64 { return v }}, + {2, func(v float64) float64 { return 0.5 * (3*v*v - 1) }}, + {3, func(v float64) float64 { return 0.5 * (5*v*v*v - 3*v) }}, + // P_7's explicit expansion: (429v⁷ − 693v⁵ + 315v³ − 35v)/16. + {7, func(v float64) float64 { + return (429*math.Pow(v, 7) - 693*math.Pow(v, 5) + 315*v*v*v - 35*v) / 16 + }}, + } + for _, tc := range closed { + p, err := Legendre(tc.l, x) + if err != nil { + t.Fatalf("Legendre(%d): %v", tc.l, err) + } + for i := range x.Len() { + v := x.FloatAt(i) + want := tc.want(v) + if math.Abs(p.FloatAt(i)-want) > 1e-12*(1+math.Abs(want)) { + t.Fatalf("P_%d(%v) = %v, want %v", tc.l, v, p.FloatAt(i), want) + } + } + } + if _, err := Legendre(-1, x); err == nil { + t.Fatal("expected an error for a negative degree") + } +} + +// TestLegendreAssociated checks the associated functions against +// closed forms, the negative-order relation and orthogonality. +func TestLegendreAssociated(t *testing.T) { + x := mustFloats(t, []float64{-0.6, 0, 0.35, 0.77}) + // P_1^1 = −√(1−x²) with the Condon-Shortley phase. + p11, err := LegendreAssociated(1, 1, x) + if err != nil { + t.Fatalf("LegendreAssociated(1,1): %v", err) + } + for i := range x.Len() { + v := x.FloatAt(i) + want := -math.Sqrt(1 - v*v) + if math.Abs(p11.FloatAt(i)-want) > 1e-12 { + t.Fatalf("P_1^1(%v) = %v, want %v", v, p11.FloatAt(i), want) + } + } + // P_2^1 = −3x√(1−x²). + p21, err := LegendreAssociated(2, 1, x) + if err != nil { + t.Fatalf("LegendreAssociated(2,1): %v", err) + } + for i := range x.Len() { + v := x.FloatAt(i) + want := -3 * v * math.Sqrt(1-v*v) + if math.Abs(p21.FloatAt(i)-want) > 1e-12 { + t.Fatalf("P_2^1(%v) = %v, want %v", v, p21.FloatAt(i), want) + } + } + // P_1^0 must equal P_1. + p10, err := LegendreAssociated(1, 0, x) + if err != nil { + t.Fatalf("LegendreAssociated(1,0): %v", err) + } + p1, _ := Legendre(1, x) + for i := range x.Len() { + if p10.FloatAt(i) != p1.FloatAt(i) { + t.Fatalf("P_1^0(%v) = %v, want P_1 = %v", x.FloatAt(i), p10.FloatAt(i), p1.FloatAt(i)) + } + } + // Negative order: P_1^(−1) = −(0)!/(2)!·P_1^1 = −½ P_1^1. + pm1, err := LegendreAssociated(1, -1, x) + if err != nil { + t.Fatalf("LegendreAssociated(1,−1): %v", err) + } + for i := range x.Len() { + if math.Abs(pm1.FloatAt(i)+0.5*p11.FloatAt(i)) > 1e-12 { + t.Fatalf("P_1^−1(%v) = %v, want %v", x.FloatAt(i), pm1.FloatAt(i), -0.5*p11.FloatAt(i)) + } + } + // Orthogonality on a coarse Simpson grid: ∫₋₁¹ P_2 P_3 dx = 0. + const n = 2001 + grid := make([]float64, n) + for i := range n { + grid[i] = -1 + 2*float64(i)/float64(n-1) + } + g := mustFloats(t, grid, n) + p2, _ := Legendre(2, g) + p3, _ := Legendre(3, g) + sum := 0.0 + for i := range n { + w := 1.0 + if i == 0 || i == n-1 { + w = 1 + } else if i%2 == 1 { + w = 4 + } else { + w = 2 + } + sum += w * p2.FloatAt(i) * p3.FloatAt(i) + } + sum *= 2.0 / 3.0 / float64(n-1) + if math.Abs(sum) > 1e-6 { + t.Fatalf("∫P_2 P_3 = %v, want 0", sum) + } + if _, err := LegendreAssociated(2, 3, x); err == nil { + t.Fatal("expected an error for |m| > l") + } +} + +// TestSphericalHarmonics checks Y_lm against closed forms and the +// conjugation symmetry Y_{l,−m} = (−1)^m conj(Y_{l,m}). +func TestSphericalHarmonics(t *testing.T) { + theta := mustFloats(t, []float64{0.3, 1.1, 2.4}) + phi := mustFloats(t, []float64{0, 0.9, -1.3}) + + // Y_0^0 = 1/√(4π). + y00, err := SphericalHarmonic(0, 0, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonic(0,0): %v", err) + } + want := 1 / math.Sqrt(4*math.Pi) + for i := range theta.Len() { + if d := cmplxAbsDiff(y00.ComplexAt(i), complex(want, 0)); d > 1e-12 { + t.Fatalf("Y_0^0 = %v, want %v", y00.ComplexAt(i), want) + } + } + + // Y_1^0 = √(3/4π) cosθ. + y10, err := SphericalHarmonic(1, 0, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonic(1,0): %v", err) + } + for i := range theta.Len() { + w := math.Sqrt(3/(4*math.Pi)) * math.Cos(theta.FloatAt(i)) + if d := cmplxAbsDiff(y10.ComplexAt(i), complex(w, 0)); d > 1e-12 { + t.Fatalf("Y_1^0 = %v, want %v", y10.ComplexAt(i), w) + } + } + + // Y_1^1 = −√(3/8π) sinθ e^{iφ}. + y11, err := SphericalHarmonic(1, 1, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonic(1,1): %v", err) + } + for i := range theta.Len() { + mod := math.Sqrt(3/(8*math.Pi)) * math.Sin(theta.FloatAt(i)) + w := -complex(mod, 0) * cmplxExpStd(complex(0, phi.FloatAt(i))) + if d := cmplxAbsDiff(y11.ComplexAt(i), w); d > 1e-12 { + t.Fatalf("Y_1^1 = %v, want %v", y11.ComplexAt(i), w) + } + } + + // Conjugation symmetry. + ym1, err := SphericalHarmonic(1, -1, theta, phi) + if err != nil { + t.Fatalf("SphericalHarmonic(1,−1): %v", err) + } + for i := range theta.Len() { + if d := cmplxAbsDiff(ym1.ComplexAt(i), -cmplxConjStd(y11.ComplexAt(i))); d > 1e-12 { + t.Fatalf("Y_1^−1 symmetry broken: %v vs %v", + ym1.ComplexAt(i), -cmplxConjStd(y11.ComplexAt(i))) + } + } + + if _, err := SphericalHarmonic(1, 2, theta, phi); err == nil { + t.Fatal("expected an error for |m| > l") + } + if _, err := SphericalHarmonic(0, 0, mustFloats(t, []float64{1}), phi); err == nil { + t.Fatal("expected a shape-mismatch error") + } +} + +// cmplxAbsDiff returns |a − b| for two complex values. +func cmplxAbsDiff(a, b complex128) float64 { + d := real(a) - real(b) + e := imag(a) - imag(b) + return math.Sqrt(d*d + e*e) +} + +// cmplxExpStd is e^{i·z}. +func cmplxExpStd(z complex128) complex128 { + return complex(math.Cos(imag(z)), math.Sin(imag(z))) +} + +// cmplxConjStd conjugates z. +func cmplxConjStd(z complex128) complex128 { + return complex(real(z), -imag(z)) +} + +// TestSphericalBessel checks j_l and y_l against their closed forms +// for l ≤ 2, the recurrence linking three consecutive degrees, and +// the small-x limits of j. +func TestSphericalBessel(t *testing.T) { + x := mustFloats(t, []float64{0.05, 0.5, 1, 3.7, 20}) + + j0, err := SphericalBesselJ(0, x) + if err != nil { + t.Fatalf("SphericalBesselJ(0): %v", err) + } + y0, err := SphericalBesselY(0, x) + if err != nil { + t.Fatalf("SphericalBesselY(0): %v", err) + } + for i := range x.Len() { + xv := x.FloatAt(i) + if math.Abs(j0.FloatAt(i)-math.Sin(xv)/xv) > 1e-12 { + t.Fatalf("j_0(%v) = %v, want %v", xv, j0.FloatAt(i), math.Sin(xv)/xv) + } + if math.Abs(y0.FloatAt(i)-(-math.Cos(xv)/xv)) > 1e-12 { + t.Fatalf("y_0(%v) = %v, want %v", xv, y0.FloatAt(i), -math.Cos(xv)/xv) + } + } + + j1, _ := SphericalBesselJ(1, x) + j2, _ := SphericalBesselJ(2, x) + y1, _ := SphericalBesselY(1, x) + y2, _ := SphericalBesselY(2, x) + for i := range x.Len() { + xv := x.FloatAt(i) + // j_1 = sin/x² − cos/x and y_1 = −cos/x² − sin/x. + j1w := math.Sin(xv)/(xv*xv) - math.Cos(xv)/xv + y1w := -math.Cos(xv)/(xv*xv) - math.Sin(xv)/xv + if math.Abs(j1.FloatAt(i)-j1w) > 1e-11*(1+math.Abs(j1w)) { + t.Fatalf("j_1(%v) = %v, want %v", xv, j1.FloatAt(i), j1w) + } + if math.Abs(y1.FloatAt(i)-y1w) > 1e-11*(1+math.Abs(y1w)) { + t.Fatalf("y_1(%v) = %v, want %v", xv, y1.FloatAt(i), y1w) + } + // Three-term recurrence j_{l+1} = (2l+1)/x·j_l − j_{l−1} at l=1. + if d := j2.FloatAt(i) - (3/xv*j1.FloatAt(i) - j0.FloatAt(i)); math.Abs(d) > 1e-10*(1+math.Abs(j2.FloatAt(i))) { + t.Fatalf("j recurrence at %v broken: %v", xv, d) + } + if d := y2.FloatAt(i) - (3/xv*y1.FloatAt(i) - y0.FloatAt(i)); math.Abs(d) > 1e-10*(1+math.Abs(y2.FloatAt(i))) { + t.Fatalf("y recurrence at %v broken: %v", xv, d) + } + } + + // Small-x limits: j_l(0) = δ_{l0}. + atZero := mustFloats(t, []float64{0}) + for l := range 4 { + jl, err := SphericalBesselJ(l, atZero) + if err != nil { + t.Fatalf("SphericalBesselJ(%d, 0): %v", l, err) + } + want := 0.0 + if l == 0 { + want = 1 + } + if jl.FloatAt(0) != want { + t.Fatalf("j_%d(0) = %v, want %v", l, jl.FloatAt(0), want) + } + } + // j_l below its degree is tiny compared to j_0 at the same point: + // the downward recurrence must keep the small roots, not flood + // them with round-off. Reference: the small-x asymptote + // j_l(x) tends to x^l/(2l+1)!!; at x = 0.01 the first correction + // term x²/(2(2l+3)) is 6.4e-7 relative, well inside the tolerance. + small := mustFloats(t, []float64{0.01}) + j5, err := SphericalBesselJ(5, small) + if err != nil { + t.Fatalf("SphericalBesselJ(5): %v", err) + } + asymptote := math.Pow(0.01, 5) / 10395 // 11!! = 10395 + if d := math.Abs(j5.FloatAt(0) - asymptote); d > 1e-5*math.Abs(asymptote) { + t.Fatalf("j_5(0.01) = %v, want ≈ %v (downward recurrence lost the root)", + j5.FloatAt(0), asymptote) + } + if _, err := SphericalBesselJ(-1, small); err == nil { + t.Fatal("expected an error for a negative degree") + } +} diff --git a/internal/core/tensor.go b/internal/core/tensor.go new file mode 100644 index 0000000..0706914 --- /dev/null +++ b/internal/core/tensor.go @@ -0,0 +1,1246 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package tensor provides immutable, shape-checked numeric arrays for +// Go. +// +// tensor is a general-purpose scientific computing library: it adds no +// dependency at all. The founding rules: values are immutable, shapes never +// broadcast silently, and elements are int64 and float64 first. +package core + +import ( + "fmt" + "math" + "slices" + "strconv" + "strings" +) + +// Dtype is the element type of an Array. Mixing dtypes in one operation +// promotes along the ladder int to float16 to float32 to float64 to +// complex128. The constants' numeric values are +// part of the recorded oracle contract and stay fixed: float16 slots +// into the ladder at its own value, and promote runs through +// dtypeRank rather than through the ordinals. +type Dtype uint8 + +const ( + // Int is the int64 element type. + Int Dtype = iota + // Float32 is the float32 element type: the ML memory and bandwidth + // dtype. + Float32 + // Float is the float64 element type, the default real dtype. + Float + // Complex is the complex128 element type. + Complex + // Float16 is the IEEE 754 binary16 element type: the half payload + // holds uint16 bit patterns, widened exactly on every read. + // It sits between Int and Float32 on the promotion ladder (see + // dtypeRank), one step below float32 in both precision and range. + Float16 + // Bool is the boolean element type: logical vectors for masking, + // Where and masked reads. It carries no arithmetic of its own; a + // binary arithmetic operation on Bool operands is a loud error. + Bool + // Int8 is the int8 element type: signed bytes. + Int8 + // Uint8 is the uint8 element type: byte payloads for images, masks + // and the byte classes of the file formats. + Uint8 + // Int16 is the int16 element type. + Int16 + // Uint16 is the uint16 element type. + Uint16 + // Int32 is the int32 element type. + Int32 + // Uint32 is the uint32 element type. + Uint32 +) + +// String renders the dtype as it appears in diagnostics: "bool", +// "int", "int8", "uint8", "int16", "uint16", "int32", "uint32", +// "float16", "float32", "float", "complex". +func (d Dtype) String() string { + switch d { + case Bool: + return "bool" + case Float16: + return "float16" + case Float32: + return "float32" + case Float: + return "float" + case Complex: + return "complex" + case Int8: + return "int8" + case Uint8: + return "uint8" + case Int16: + return "int16" + case Uint16: + return "uint16" + case Int32: + return "int32" + case Uint32: + return "uint32" + default: + return "int" + } +} + +// Array is an immutable, shape-checked numeric array: every +// operation returns a new Array and writes neither its receiver nor its +// arguments. Elements are int64, float16 (uint16 half bit patterns), +// float32, float64 or complex128, any number of +// dimensions is allowed, and storage is row-major. +// +// A result may be a read-only view sharing another array's storage. +// Such a view's payload is rebased to the view's origin, so element 0 +// of the view is payload[0] and element i is payload[i]: the package +// never sets strides on an array it produces, which makes physIndex +// the identity for every array from a public constructor. The strides +// field is defensive machinery for strided arrays constructed directly +// in tests; kernels that read payloads raw rely on it staying nil. +type Array struct { + shape []int + dt Dtype + ints []int64 + halves []uint16 + floats32 []float32 + floats []float64 + complexes []complex128 + bools []bool + i8s []int8 + u8s []uint8 + i16s []int16 + u16s []uint16 + i32s []int32 + u32s []uint32 + // strides is nil for a contiguous array and non-nil only for a + // read-only view; it is never written through. + strides []int +} + +// Shape returns a copy of the array's dimensions. +func (a *Array) Shape() []int { + out := make([]int, len(a.shape)) + copy(out, a.shape) + return out +} + +// Dtype returns the element type. +func (a *Array) Dtype() Dtype { return a.dt } + +// Len returns the number of elements: the product of the shape, which +// for a view is shorter than the storage it aliases. +func (a *Array) Len() int { + n := 1 + for _, d := range a.shape { + n *= d + } + return n +} + +// isContiguous reports whether payload[i] is element i. Only a +// contiguous array may be handed to code that reads payload windows +// directly. +func (a *Array) isContiguous() bool { return a.strides == nil } + +// physIndex maps a logical row-major flat index to the payload index. +// A contiguous array maps identity, so payload[i] is element i; a view +// decomposes the index into coordinates and applies its strides. +func (a *Array) physIndex(flat int) int { + if a.strides == nil { + return flat + } + off := 0 + for d := a.NDim() - 1; d >= 0; d-- { + off += (flat % a.shape[d]) * a.strides[d] + flat /= a.shape[d] + } + return off +} + +// materialise returns an array owning a contiguous payload. A +// contiguous array is returned as is; a strided view is copied into a +// fresh dense array: the boundary a kernel calls when it needs raw +// payload windows rather than accessor reads. +func (a *Array) materialise() *Array { + if a.strides == nil { + return a + } + n := a.Len() + out := &Array{shape: a.Shape(), dt: a.dt} + out.alloc(n) + // The dtype dispatch sits outside the loop: the per-element work is + // the physIndex rebasing, not a switch re-evaluated n times. + switch a.dt { + case Int: + for i := range n { + out.ints[i] = a.ints[a.physIndex(i)] + } + case Float16: + for i := range n { + out.halves[i] = a.halves[a.physIndex(i)] + } + case Float32: + for i := range n { + out.floats32[i] = a.floats32[a.physIndex(i)] + } + case Float: + for i := range n { + out.floats[i] = a.floats[a.physIndex(i)] + } + case Complex: + for i := range n { + out.complexes[i] = a.complexes[a.physIndex(i)] + } + case Bool: + for i := range n { + out.bools[i] = a.bools[a.physIndex(i)] + } + case Int8: + for i := range n { + out.i8s[i] = a.i8s[a.physIndex(i)] + } + case Uint8: + for i := range n { + out.u8s[i] = a.u8s[a.physIndex(i)] + } + case Int16: + for i := range n { + out.i16s[i] = a.i16s[a.physIndex(i)] + } + case Uint16: + for i := range n { + out.u16s[i] = a.u16s[a.physIndex(i)] + } + case Int32: + for i := range n { + out.i32s[i] = a.i32s[a.physIndex(i)] + } + default: + for i := range n { + out.u32s[i] = a.u32s[a.physIndex(i)] + } + } + return out +} + +// NDim returns the number of dimensions. +func (a *Array) NDim() int { return len(a.shape) } + +// FromInts builds an int array of the given shape from vals, copying them: +// later changes to vals never reach the array. +func FromInts(vals []int64, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + ints := make([]int64, len(vals)) + copy(ints, vals) + return &Array{shape: sh, dt: Int, ints: ints}, nil +} + +// FromFloat32s builds a float32 array of the given shape from vals, +// copying them. +func FromFloat32s(vals []float32, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + floats32 := make([]float32, len(vals)) + copy(floats32, vals) + return &Array{shape: sh, dt: Float32, floats32: floats32}, nil +} + +// FromFloat32Slice aliases vals into a float32 array of the given +// shape WITHOUT copying. The returned array is valid only while the +// caller leaves vals untouched: reads reflect later writes, so the +// pattern fits build-then-consume windows (generation caches, staging +// slabs) and nothing else. Alias everything or nothing: no offset, +// no strides. +func FromFloat32Slice(vals []float32, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Float32, floats32: vals}, nil +} + +// FromFloatSlice aliases vals into a float array of the given shape +// without copying: the float64 twin of FromFloat32Slice with the same +// build-then-consume contract. +func FromFloatSlice(vals []float64, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Float, floats: vals}, nil +} + +// FromFloats builds a float array of the given shape from vals, copying +// them. +func FromFloats(vals []float64, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + floats := make([]float64, len(vals)) + copy(floats, vals) + return &Array{shape: sh, dt: Float, floats: floats}, nil +} + +// FromComplexes builds a complex array of the given shape from vals, +// copying them. +func FromComplexes(vals []complex128, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + complexes := make([]complex128, len(vals)) + copy(complexes, vals) + return &Array{shape: sh, dt: Complex, complexes: complexes}, nil +} + +// Zeros builds an array of the given dtype and shape filled with zeros. +// The dtype must be one of the five element types the package stores; +// anything else is an error, not a payload every diagnostic renders +// as int. +func Zeros(dt Dtype, shape ...int) (*Array, error) { + return filled(dt, shape, 0, 0, 0) +} + +// Ones builds an array of the given dtype and shape filled with ones. +func Ones(dt Dtype, shape ...int) (*Array, error) { + return filled(dt, shape, 1, 1, 1) +} + +// FullI builds an int array of the given shape filled with v. +func FullI(v int64, shape ...int) (*Array, error) { + return filled(Int, shape, v, float64(v), complex(float64(v), 0)) +} + +// FullF builds a float array of the given shape filled with v. +func FullF(v float64, shape ...int) (*Array, error) { + return filled(Float, shape, int64(v), v, complex(v, 0)) +} + +// FullF32s builds a float32 array of the given shape filled with v. +func FullF32s(v float32, shape ...int) (*Array, error) { + return filled(Float32, shape, int64(v), float64(v), complex(float64(v), 0)) +} + +// FullF16 builds a float16 array of the given shape filled with v, +// narrowed to the nearest half under the HalfFromFloat64 contract. +func FullF16(v float64, shape ...int) (*Array, error) { + return filled(Float16, shape, int64(v), v, complex(v, 0)) +} + +// FullC builds a complex array of the given shape filled with v. +func FullC(v complex128, shape ...int) (*Array, error) { + return filled(Complex, shape, int64(real(v)), real(v), v) +} + +func filled(dt Dtype, shape []int, iv int64, fv float64, cv complex128) (*Array, error) { + // filled is the throat every public filling constructor (Zeros, + // Ones, FullI, FullF, FullF16, FullF32s, FullC) and everything built + // on Zeros passes through, so the dtype check lives here once: a + // caller's Dtype(99) is an error, not a payload every diagnostic + // renders as int. The switch is over the element types the package + // stores, and the kernels that build arrays through alloc directly + // never pay it. + switch dt { + case Int, Float16, Float32, Float, Complex, Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32: + default: + return nil, errf("unknown element type %d", uint8(dt)) + } + total, sh, err := checkedDims(shape) + if err != nil { + return nil, err + } + out := &Array{shape: sh, dt: dt} + out.alloc(total) + switch dt { + case Int: + for i := range out.ints { + out.ints[i] = iv + } + case Float16: + hv := HalfFromFloat64(fv) + for i := range out.halves { + out.halves[i] = hv + } + case Float32: + f32v := float32(fv) + for i := range out.floats32 { + out.floats32[i] = f32v + } + case Float: + for i := range out.floats { + out.floats[i] = fv + } + case Complex: + for i := range out.complexes { + out.complexes[i] = cv + } + case Bool: + bv := iv != 0 + for i := range out.bools { + out.bools[i] = bv + } + case Int8: + v := int8(iv) + for i := range out.i8s { + out.i8s[i] = v + } + case Uint8: + v := uint8(iv) + for i := range out.u8s { + out.u8s[i] = v + } + case Int16: + v := int16(iv) + for i := range out.i16s { + out.i16s[i] = v + } + case Uint16: + v := uint16(iv) + for i := range out.u16s { + out.u16s[i] = v + } + case Int32: + v := int32(iv) + for i := range out.i32s { + out.i32s[i] = v + } + default: + v := uint32(iv) + for i := range out.u32s { + out.u32s[i] = v + } + } + return out, nil +} + +// Range builds the int array start, start+1, …, stop-1; start >= stop +// yields an empty array. +func Range(start, stop int64) (*Array, error) { + return RangeBy(start, stop, 1) +} + +// RangeBy builds the int array start, start+step, …, staying below stop +// for a positive step and above it for a negative one. A zero step is an +// error, and the walk stops early rather than looping when the next value +// would overflow int64. +func RangeBy(start, stop, step int64) (*Array, error) { + if step == 0 { + return nil, errf("RangeBy: step cannot be zero") + } + var vals []int64 + for v := start; (step > 0 && v < stop) || (step < 0 && v > stop); { + vals = append(vals, v) + next := v + step + if (step > 0 && next < v) || (step < 0 && next > v) { + break // the addition wrapped; nothing sane remains + } + v = next + } + return FromInts(vals, len(vals)) +} + +// Equal reports whether two arrays have the same dtype, shape and values. +// The dtype is part of the identity: an int 1 does not equal a float 1.0. +// Floats compare with ==, so arrays holding NaN are never equal. +func Equal(a, b *Array) bool { + if a == b { + return true + } + if a == nil || b == nil { + return false + } + if a.dt != b.dt || !sameShape(a.shape, b.shape) { + return false + } + // Compare exactly the arrays' own elements: a rebased view's payload + // may run past its element count, and those invisible tail slots + // must not influence equality. slices.Equal compares with ==, so + // NaN never equals NaN, the documented Equal semantics. A strided + // view's payload is not in element order, so both sides are + // materialised first; a contiguous array is returned unchanged. + a = a.materialise() + b = b.materialise() + n := a.Len() + switch a.dt { + case Int: + return slices.Equal(a.ints[:n], b.ints[:n]) + case Float16: + // Half payloads compare by value, not by bits: +0.0 and -0.0 + // compare equal the way every float dtype's == does, and two + // NaNs never compare equal. + return slices.EqualFunc(a.halves[:n], b.halves[:n], func(x, y uint16) bool { + return HalfToFloat64(x) == HalfToFloat64(y) + }) + case Float32: + return slices.Equal(a.floats32[:n], b.floats32[:n]) + case Float: + return slices.Equal(a.floats[:n], b.floats[:n]) + case Complex: + return slices.Equal(a.complexes[:n], b.complexes[:n]) + case Bool: + return slices.Equal(a.bools[:n], b.bools[:n]) + case Int8: + return slices.Equal(a.i8s[:n], b.i8s[:n]) + case Uint8: + return slices.Equal(a.u8s[:n], b.u8s[:n]) + case Int16: + return slices.Equal(a.i16s[:n], b.i16s[:n]) + case Uint16: + return slices.Equal(a.u16s[:n], b.u16s[:n]) + case Int32: + return slices.Equal(a.i32s[:n], b.i32s[:n]) + default: + return slices.Equal(a.u32s[:n], b.u32s[:n]) + } +} + +// String renders the dtype, the shape and the values, as in +// "int (2, 2) [1, 2, 3, 4]". Arrays of three or more dimensions wrap +// the values per trailing dimension so the structure is readable: +// "float (2, 2, 2) [[1, 2, 3, 4], [5, 6, 7, 8]]". It is a debugging +// aid, not a format. +func (a *Array) String() string { + var sb strings.Builder + fmt.Fprintf(&sb, "%s %s ", a.dt, shapeText(a.shape)) + if a.NDim() >= 3 { + // writeSlice emits the full bracket structure, including the + // outermost pair. + a.writeSlice(&sb, 0, 0) + } else { + sb.WriteByte('[') + a.writeValues(&sb) + sb.WriteByte(']') + } + return sb.String() +} + +// writeValues renders the flat values with brackets per dimension for +// arrays of rank ≥ 3; ranks 1 and 2 stay flat, matching the compact +// diagnostic format the tests and examples rely on. +func (a *Array) writeValues(sb *strings.Builder) { + nd := a.NDim() + if nd <= 2 { + for i := range a.Len() { + if i > 0 { + sb.WriteString(", ") + } + writeElem(sb, a, i) + } + return + } + // Recursively emit each slice along the leading dimension. + a.writeSlice(sb, 0, 0) +} + +// writeSlice emits the elements of a[coord...] with a bracket per +// remaining dimension. Dimensions of size 1 are transparent: a shape +// like (1, 1, 3) renders as [7, 9, 11], not [[[7, 9, 11]]]. +func (a *Array) writeSlice(sb *strings.Builder, dim, flat int) { + if a.shape[dim] == 1 && dim < a.NDim()-1 { + a.writeSlice(sb, dim+1, flat) + return + } + sb.WriteByte('[') + block := 1 + for d := dim + 1; d < a.NDim(); d++ { + block *= a.shape[d] + } + for i := range a.shape[dim] { + if i > 0 { + sb.WriteString(", ") + } + if dim == a.NDim()-1 { + writeElem(sb, a, flat+i) + } else { + a.writeSlice(sb, dim+1, flat+i*block) + } + } + sb.WriteByte(']') +} + +// writeElem renders one element in its dtype's format. +func writeElem(sb *strings.Builder, a *Array, i int) { + if a.strides != nil { + i = a.physIndex(i) + } + switch a.dt { + case Int: + fmt.Fprintf(sb, "%d", a.ints[i]) + case Float16: + // Printed through the exact float64 widening, the same 'g' + // shortest-round-trip formatting float32 and float use. + sb.WriteString(strconv.FormatFloat(HalfToFloat64(a.halves[i]), 'g', -1, 64)) + case Float32: + sb.WriteString(strconv.FormatFloat(float64(a.floats32[i]), 'g', -1, 32)) + case Float: + sb.WriteString(strconv.FormatFloat(a.floats[i], 'g', -1, 64)) + case Complex: + fmt.Fprintf(sb, "%v", a.complexes[i]) + case Bool: + sb.WriteString(strconv.FormatBool(a.bools[i])) + case Int8: + fmt.Fprintf(sb, "%d", a.i8s[i]) + case Uint8: + fmt.Fprintf(sb, "%d", a.u8s[i]) + case Int16: + fmt.Fprintf(sb, "%d", a.i16s[i]) + case Uint16: + fmt.Fprintf(sb, "%d", a.u16s[i]) + case Int32: + fmt.Fprintf(sb, "%d", a.i32s[i]) + default: + fmt.Fprintf(sb, "%d", a.u32s[i]) + } +} + +// checkedDims validates a shape: at least one dimension, none negative, +// and an element count that fits in an int. It returns the element count +// and a private copy of the shape. +func checkedDims(shape []int) (int, []int, error) { + if len(shape) == 0 { + return 0, nil, errf("an array needs at least one dimension") + } + total := 1 + for _, d := range shape { + if d < 0 { + return 0, nil, errf("dimensions must be zero or greater, got %d", d) + } + // Bound each factor before multiplying it in: a product that wraps + // would otherwise agree with a caller's own wrapped arithmetic and + // silently pair a huge declared shape with a tiny allocation. + if d != 0 && total > math.MaxInt/d { + return 0, nil, errf("the shape %s holds more elements than fit in an index", shapeText(shape)) + } + total *= d + } + sh := make([]int, len(shape)) + copy(sh, shape) + return total, sh, nil +} + +// shapeFor validates a shape for a payload of exactly n values and +// returns a private copy of it: the shared front half of every From… +// constructor. The shape must be well formed and hold all n values. +func shapeFor(shape []int, n int) ([]int, error) { + total, sh, err := checkedDims(shape) + if err != nil { + return nil, err + } + if total != n { + return nil, errf("%d values do not fill the shape %s", n, shapeText(sh)) + } + return sh, nil +} + +// sameShape reports whether two shapes are identical. +func sameShape(a, b []int) bool { return slices.Equal(a, b) } + +// floatAt returns element i as float64, widening every element type the +// package stores; each widening of an integer, a bool, a half or a +// float32 is exact. It must not be called on complex arrays; callers +// dispatch on the promoted dtype first. +func (a *Array) floatAt(i int) float64 { + if a.strides != nil { + i = a.physIndex(i) + } + switch a.dt { + case Float: + return a.floats[i] + case Float16: + return HalfToFloat64(a.halves[i]) + case Float32: + return float64(a.floats32[i]) + case Bool: + if a.bools[i] { + return 1 + } + return 0 + case Int8: + return float64(a.i8s[i]) + case Uint8: + return float64(a.u8s[i]) + case Int16: + return float64(a.i16s[i]) + case Uint16: + return float64(a.u16s[i]) + case Int32: + return float64(a.i32s[i]) + case Uint32: + return float64(a.u32s[i]) + default: + return float64(a.ints[i]) + } +} + +// float32At returns element i as float32, narrowing float64 elements, +// used where a Float32 result must be written from a wider computation. +func (a *Array) float32At(i int) float32 { + if a.dt == Float32 { + if a.strides != nil { + i = a.physIndex(i) + } + return a.floats32[i] + } + return float32(a.floatAt(i)) +} + +// complexAt returns element i as complex128, converting every real +// element exactly. +func (a *Array) complexAt(i int) complex128 { + if a.strides != nil { + i = a.physIndex(i) + } + switch a.dt { + case Complex: + return a.complexes[i] + case Float: + return complex(a.floats[i], 0) + case Float16: + return complex(HalfToFloat64(a.halves[i]), 0) + default: + return complex(a.floatAt(i), 0) + } +} + +// intAt returns element i as int64, resolving the stride table a +// test-side view carries. The integer-class dtypes widen exactly; bool +// reads 0 or 1; a float element casts the way an implicit store into an +// int destination has always cast. The elementwise fallback reads its +// operands through it, and setConverted reads narrow sources through it. +func (a *Array) intAt(i int) int64 { + if a.strides != nil { + i = a.physIndex(i) + } + switch a.dt { + case Int: + return a.ints[i] + case Bool: + if a.bools[i] { + return 1 + } + return 0 + case Int8: + return int64(a.i8s[i]) + case Uint8: + return int64(a.u8s[i]) + case Int16: + return int64(a.i16s[i]) + case Uint16: + return int64(a.u16s[i]) + case Int32: + return int64(a.i32s[i]) + case Uint32: + return int64(a.u32s[i]) + case Float16: + return int64(HalfToFloat64(a.halves[i])) + case Float32: + return int64(a.floats32[i]) + case Float: + return int64(a.floats[i]) + default: + return int64(real(a.complexes[i])) + } +} + +// cloneData returns a deep copy of the array's own elements for the +// five legacy payloads (int64, half, float32, float64, complex128): +// for a view that is the view's contents, not the storage it aliases. +// The narrow element types clone through cloneArray, which carries +// their payloads; a narrow array reaching this function has no payload +// field here, so the default arm copies from its nil complex payload +// and fabricates zero complex values, which is why every caller +// dispatches to cloneArray before it reaches this function. A +// contiguous payload copies with a single memcpy per dtype; only a +// strided view pays the per-element physIndex walk. Exactly one of the +// five slices is non-nil, matching the array's legacy dtype. +func (a *Array) cloneData() ([]int64, []uint16, []float32, []float64, []complex128) { + if a.strides == nil { + // A contiguous array's own elements are payload[:Len]; a + // rebased view's payload may run further, so the copy is + // bounded by n rather than trusting the payload length. + n := a.Len() + switch a.dt { + case Int: + ints := make([]int64, n) + copy(ints, a.ints) + return ints, nil, nil, nil, nil + case Float16: + halves := make([]uint16, n) + copy(halves, a.halves) + return nil, halves, nil, nil, nil + case Float32: + floats32 := make([]float32, n) + copy(floats32, a.floats32) + return nil, nil, floats32, nil, nil + case Float: + floats := make([]float64, n) + copy(floats, a.floats) + return nil, nil, nil, floats, nil + default: + complexes := make([]complex128, n) + copy(complexes, a.complexes) + return nil, nil, nil, nil, complexes + } + } + n := a.Len() + switch a.dt { + case Int: + ints := make([]int64, n) + for i := range n { + ints[i] = a.ints[a.physIndex(i)] + } + return ints, nil, nil, nil, nil + case Float16: + halves := make([]uint16, n) + for i := range n { + halves[i] = a.halves[a.physIndex(i)] + } + return nil, halves, nil, nil, nil + case Float32: + floats32 := make([]float32, n) + for i := range n { + floats32[i] = a.floats32[a.physIndex(i)] + } + return nil, nil, floats32, nil, nil + case Float: + floats := make([]float64, n) + for i := range n { + floats[i] = a.floats[a.physIndex(i)] + } + return nil, nil, nil, floats, nil + default: + complexes := make([]complex128, n) + for i := range n { + complexes[i] = a.complexes[a.physIndex(i)] + } + return nil, nil, nil, nil, complexes + } +} + +// alloc prepares the payload for total elements of the array's dtype. +func (a *Array) alloc(total int) { + switch a.dt { + case Int: + a.ints = make([]int64, total) + case Float16: + a.halves = make([]uint16, total) + case Float32: + a.floats32 = make([]float32, total) + case Float: + a.floats = make([]float64, total) + case Complex: + a.complexes = make([]complex128, total) + case Bool: + a.bools = make([]bool, total) + case Int8: + a.i8s = make([]int8, total) + case Uint8: + a.u8s = make([]uint8, total) + case Int16: + a.i16s = make([]int16, total) + case Uint16: + a.u16s = make([]uint16, total) + case Int32: + a.i32s = make([]int32, total) + default: + a.u32s = make([]uint32, total) + } +} + +// setFrom copies element srcFlat of src into element dstFlat of a. The +// dtypes must match; the destination is always a freshly allocated, +// contiguous array, while the source may be a view. +func (a *Array) setFrom(dstFlat int, src *Array, srcFlat int) { + srcFlat = src.physIndex(srcFlat) + switch a.dt { + case Int: + a.ints[dstFlat] = src.ints[srcFlat] + case Float16: + a.halves[dstFlat] = src.halves[srcFlat] + case Float32: + a.floats32[dstFlat] = src.floats32[srcFlat] + case Float: + a.floats[dstFlat] = src.floats[srcFlat] + case Complex: + a.complexes[dstFlat] = src.complexes[srcFlat] + case Bool: + a.bools[dstFlat] = src.bools[srcFlat] + case Int8: + a.i8s[dstFlat] = src.i8s[srcFlat] + case Uint8: + a.u8s[dstFlat] = src.u8s[srcFlat] + case Int16: + a.i16s[dstFlat] = src.i16s[srcFlat] + case Uint16: + a.u16s[dstFlat] = src.u16s[srcFlat] + case Int32: + a.i32s[dstFlat] = src.i32s[srcFlat] + default: + a.u32s[dstFlat] = src.u32s[srcFlat] + } +} + +// canStore reports whether an element of dtype src may be stored in a +// destination of dtype dst. The ladder runs up (int to float16 to +// float32 to float to complex) and the descending directions the +// library performs are float to int, float to float32, float to float16 +// and complex to float (the real part), all as Astype documents; +// narrowing complex into int, float16 or float32 is the pair Astype +// rejects, and every store respects that. +func canStore(dst, src Dtype) bool { + if src == Complex { + // Narrowing complex into any integer destination is the pair + // Astype rejects, and every store respects that; complex to + // float keeps the real part, and complex to bool reads against + // zero, the same test Astype's bool target applies. + return dst != Int && dst != Float16 && dst != Float32 && + dst != Int8 && dst != Uint8 && + dst != Int16 && dst != Uint16 && dst != Int32 && dst != Uint32 + } + return true +} + +// setConverted stores element srcFlat of src in a's dtype, converting +// along the ladder, used when Concat and Stack promote their operands +// and when Scatter mixes dtypes. The read goes through the source's own +// accessor, so a source above the destination on the ladder never reads +// a payload the source does not have; the rejected narrowings are +// decided by canStore before the walk. +func (a *Array) setConverted(dstFlat int, src *Array, srcFlat int) { + switch a.dt { + case Int: + if src.dt == Int { + // Exact: the float64 detour rounds above 2^53. + a.ints[dstFlat] = src.ints[src.physIndex(srcFlat)] + return + } + a.ints[dstFlat] = int64(src.floatAt(srcFlat)) + case Float16: + a.halves[dstFlat] = HalfFromFloat64(src.floatAt(srcFlat)) + case Float32: + a.floats32[dstFlat] = src.float32At(srcFlat) + case Float: + if src.dt == Complex { + // Complex to float keeps the real part, as Astype documents. + a.floats[dstFlat] = real(src.complexAt(srcFlat)) + return + } + a.floats[dstFlat] = src.floatAt(srcFlat) + case Complex: + a.complexes[dstFlat] = src.complexAt(srcFlat) + case Bool: + a.bools[dstFlat] = src.boolAt(srcFlat) + case Int8: + a.i8s[dstFlat] = int8(src.intAt(srcFlat)) + case Uint8: + a.u8s[dstFlat] = uint8(src.intAt(srcFlat)) + case Int16: + a.i16s[dstFlat] = int16(src.intAt(srcFlat)) + case Uint16: + a.u16s[dstFlat] = uint16(src.intAt(srcFlat)) + case Int32: + a.i32s[dstFlat] = int32(src.intAt(srcFlat)) + default: + // Uint32 and any unlisted ordinal: the widened read casts down, + // the implicit-store semantics every promoted Concat and Scatter + // target has always carried. + a.u32s[dstFlat] = uint32(src.intAt(srcFlat)) + } +} + +// dtypeRank returns a dtype's position on the numeric tower: bool +// below the integer widths below int64 below float16 below float32 +// below float64 below complex128. Mixed signedness integer pairs are +// not resolved by rank alone; promote resolves the whole integer class +// through intPromote in narrowdt.go, and only cross-class pairs walk +// this ladder. +func dtypeRank(d Dtype) int { + switch d { + case Bool: + return 0 + case Int8, Uint8: + return 1 + case Int16, Uint16: + return 2 + case Int32, Uint32: + return 3 + case Float16: + return 5 + case Float32: + return 6 + case Float: + return 7 + case Complex: + return 8 + default: + return 4 + } +} + +// promote returns the common dtype under the library's numeric tower. +// Across classes the higher ladder position decides, the rule the +// int-to-float-to-complex promotions have always walked; inside the +// integer class the result is the smallest integer dtype whose value +// range contains both operands, so a mixed signedness pair widens +// instead of losing its negative half (int8 with uint8 answers int16). +func promote(a, b Dtype) Dtype { + if a == b { + return a + } + if intClass(a) && intClass(b) { + return intPromote[intClassIndex(a)][intClassIndex(b)] + } + if dtypeRank(a) >= dtypeRank(b) { + return a + } + return b +} + +// Strided reports whether the array carries a non-trivial stride +// layout. Dense arrays answer false. +func (a *Array) Strided() bool { return a.strides != nil } + +// RawFloats returns the array's float64 payload directly: element i +// of the array sits at payload index i, views included, because the +// package never sets strides. Treat the slice as read-only; only +// freshly allocated arrays an owner writes through it are safe to +// mutate. +func (a *Array) RawFloats() []float64 { return a.floats } + +// RawFloat32s returns the float32 payload with the RawFloats +// contract. +func (a *Array) RawFloat32s() []float32 { return a.floats32 } + +// RawInts returns the int64 payload with the RawFloats contract. +func (a *Array) RawInts() []int64 { return a.ints } + +// RawComplexes returns the complex128 payload with the RawFloats +// contract. +func (a *Array) RawComplexes() []complex128 { return a.complexes } + +// FloatsFromArray builds an array that takes ownership of vals: no +// copy is made, so the caller must not touch the slice afterwards. +// The value count must fill the shape exactly. +func FloatsFromArray(vals []float64, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Float, floats: vals}, nil +} + +// ComplexFromArray builds an array that takes ownership of vals, +// with the FloatsFromArray contract. +func ComplexFromArray(vals []complex128, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Complex, complexes: vals}, nil +} + +// IntsFromArray builds an array that takes ownership of vals, with +// the FloatsFromArray contract. +func IntsFromArray(vals []int64, shape ...int) (*Array, error) { + sh, err := shapeFor(shape, len(vals)) + if err != nil { + return nil, err + } + return &Array{shape: sh, dt: Int, ints: vals}, nil +} + +// New allocates a zeroed array of the given dtype and shape, the +// no-error constructor for internally derived shapes: the count and +// shape come from existing arrays, so an invalid argument is a bug in +// the caller, answered by a nil array. +func New(dt Dtype, shape ...int) *Array { + a, err := Zeros(dt, shape...) + if err != nil { + return nil + } + return a +} + +// ComplexValues returns the array's elements as complex values, +// reading a complex payload directly and converting everything else. +// A contiguous complex array shares its payload, so the slice is a +// read-only alias; a strided view is copied out element by element. +func (a *Array) ComplexValues(name string) ([]complex128, error) { + if a.dt == Complex { + if a.strides == nil { + // A rebased view's payload runs past its element count, and + // only the prefix bounded by Len is the view's own data. + return a.complexes[:a.Len()], nil + } + out := make([]complex128, a.Len()) + for i := range out { + out[i] = a.complexes[a.physIndex(i)] + } + return out, nil + } + if a.NDim() != 1 { + return nil, errf("%s: needs a 1-D array, got shape %s", name, shapeText(a.shape)) + } + if a.Len() == 0 { + return nil, errf("%s: an empty array has no transform", name) + } + out := make([]complex128, a.Len()) + for i := range out { + out[i] = a.complexAt(i) + } + return out, nil +} + +// ComplexAt returns element i as complex128, converting real +// elements; every real dtype widens exactly, bool reads 0/1. +func (a *Array) ComplexAt(i int) complex128 { + if a.strides != nil { + i = a.physIndex(i) + } + switch a.dt { + case Complex: + return a.complexes[i] + case Float: + return complex(a.floats[i], 0) + case Float16: + return complex(HalfToFloat64(a.halves[i]), 0) + case Float32: + return complex(float64(a.floats32[i]), 0) + default: + return complex(a.floatAt(i), 0) + } +} + +// FloatAt returns element i widened to float64. It is the numeric +// read primitive the derivative packages build on; for typed access +// prefer IntAt/FloatAt-by-name variants and Elements[E]. +func (a *Array) FloatAt(i int) float64 { return a.floatAt(i) } + +// SetFloatAt sets element i from v, converting to the array's dtype, +// the numeric write primitive complementing FloatAt. +func (a *Array) SetFloatAt(i int, v float64) { a.setFromValue(i, v) } + +// CopyRows returns a new 2+-D array holding the rows of a at the given +// leading-dimension indices, preserving every other dimension and the +// element type. An index outside the leading dimension is an error, as +// in every other selecting entry. +func (a *Array) CopyRows(idx []int) (*Array, error) { + for _, r := range idx { + if r < 0 || r >= a.shape[0] { + return nil, errf("CopyRows: row %d out of range for %d rows", r, a.shape[0]) + } + } + rowLen := 1 + for _, d := range a.shape[1:] { + rowLen *= d + } + out := &Array{shape: append([]int{len(idx)}, a.shape[1:]...), dt: a.dt} + out.alloc(len(idx) * rowLen) + if a.strides == nil { + // Contiguous rows copy whole: one memcpy per selected row, no + // per-element dispatch. Views fall through to setFrom, which + // rebases each flat index through the strides. + switch a.dt { + case Int: + for i, r := range idx { + copy(out.ints[i*rowLen:(i+1)*rowLen], a.ints[r*rowLen:(r+1)*rowLen]) + } + case Float16: + for i, r := range idx { + copy(out.halves[i*rowLen:(i+1)*rowLen], a.halves[r*rowLen:(r+1)*rowLen]) + } + case Float32: + for i, r := range idx { + copy(out.floats32[i*rowLen:(i+1)*rowLen], a.floats32[r*rowLen:(r+1)*rowLen]) + } + case Float: + for i, r := range idx { + copy(out.floats[i*rowLen:(i+1)*rowLen], a.floats[r*rowLen:(r+1)*rowLen]) + } + case Complex: + for i, r := range idx { + copy(out.complexes[i*rowLen:(i+1)*rowLen], a.complexes[r*rowLen:(r+1)*rowLen]) + } + case Bool: + for i, r := range idx { + copy(out.bools[i*rowLen:(i+1)*rowLen], a.bools[r*rowLen:(r+1)*rowLen]) + } + case Int8: + for i, r := range idx { + copy(out.i8s[i*rowLen:(i+1)*rowLen], a.i8s[r*rowLen:(r+1)*rowLen]) + } + case Uint8: + for i, r := range idx { + copy(out.u8s[i*rowLen:(i+1)*rowLen], a.u8s[r*rowLen:(r+1)*rowLen]) + } + case Int16: + for i, r := range idx { + copy(out.i16s[i*rowLen:(i+1)*rowLen], a.i16s[r*rowLen:(r+1)*rowLen]) + } + case Uint16: + for i, r := range idx { + copy(out.u16s[i*rowLen:(i+1)*rowLen], a.u16s[r*rowLen:(r+1)*rowLen]) + } + case Int32: + for i, r := range idx { + copy(out.i32s[i*rowLen:(i+1)*rowLen], a.i32s[r*rowLen:(r+1)*rowLen]) + } + default: + for i, r := range idx { + copy(out.u32s[i*rowLen:(i+1)*rowLen], a.u32s[r*rowLen:(r+1)*rowLen]) + } + } + return out, nil + } + for i, r := range idx { + for j := range rowLen { + out.setFrom(i*rowLen+j, a, r*rowLen+j) + } + } + return out, nil +} + +// Bytes returns the payload interpreted as raw bytes, used by the +// archive writers to embed non-numeric payloads. Only int arrays carry +// byte payloads; other element types return nil. +func (a *Array) Bytes() []byte { + if a.dt != Int { + return nil + } + out := make([]byte, a.Len()) + if a.strides == nil { + // Contiguous payload: the first n slots are the array's own + // elements (a rebased view's payload may run further, so the + // walk is bounded by n, not the payload length). + for i := range out { + out[i] = byte(a.ints[i]) + } + return out + } + for i := range out { + out[i] = byte(a.ints[a.physIndex(i)]) + } + return out +} + +// FromBytes wraps raw bytes as an int64 array, the inverse of Bytes. +func FromBytes(b []byte) (*Array, error) { + vals := make([]int64, len(b)) + for i, c := range b { + vals[i] = int64(c) + } + return FromInts(vals, len(b)) +} diff --git a/internal/core/tensor_bench_test.go b/internal/core/tensor_bench_test.go new file mode 100644 index 0000000..511f031 --- /dev/null +++ b/internal/core/tensor_bench_test.go @@ -0,0 +1,32 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// Benchmarks for the tensor.go core: the payload clone behind Reshape, +// the CopyRows row gather and the view materialisation boundary. Run +// with -bench before and after any change to those paths. + +func BenchmarkReshapeContiguous(b *testing.B) { + a, _ := FromFloats(make([]float64, 1<<20), 1<<20) + b.ReportAllocs() + for b.Loop() { + if _, err := Reshape(a, 1024, 1024); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCopyRows(b *testing.B) { + a, _ := FromFloats(make([]float64, 512*512), 512, 512) + idx := make([]int, 256) + for i := range idx { + idx[i] = (i * 7) % 512 + } + b.ReportAllocs() + for b.Loop() { + a.CopyRows(idx) + } +} diff --git a/internal/core/tensor_test.go b/internal/core/tensor_test.go new file mode 100644 index 0000000..e07c0db --- /dev/null +++ b/internal/core/tensor_test.go @@ -0,0 +1,193 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "strings" + "testing" +) + +func mustFromInts(t *testing.T, vals []int64, shape ...int) *Array { + t.Helper() + a, err := FromInts(vals, shape...) + if err != nil { + t.Fatalf("FromInts(%v, %v): %v", vals, shape, err) + } + return a +} + +func mustFromFloats(t *testing.T, vals []float64, shape ...int) *Array { + t.Helper() + a, err := FromFloats(vals, shape...) + if err != nil { + t.Fatalf("FromFloats(%v, %v): %v", vals, shape, err) + } + return a +} + +func TestConstructors(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + if a.NDim() != 2 || a.Len() != 6 || a.Dtype() != Int { + t.Fatalf("accessors: ndim %d len %d dtype %s", a.NDim(), a.Len(), a.Dtype()) + } + shape := a.Shape() + if len(shape) != 2 || shape[0] != 2 || shape[1] != 3 { + t.Fatalf("shape: %v", shape) + } + + f := mustFromFloats(t, []float64{1.5, 2.5}, 2) + if f.Dtype() != Float || f.Len() != 2 { + t.Fatalf("float array: %s len %d", f.Dtype(), f.Len()) + } +} + +func TestConstructorErrors(t *testing.T) { + if _, err := FromInts(nil); err == nil || !strings.Contains(err.Error(), "at least one dimension") { + t.Fatalf("no shape: %v", err) + } + if _, err := FromInts([]int64{1}, 2); err == nil || !strings.Contains(err.Error(), "do not fill the shape") { + t.Fatalf("wrong fill: %v", err) + } + if _, err := FromInts([]int64{1}, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") { + t.Fatalf("negative dim: %v", err) + } +} + +func TestFilledConstructors(t *testing.T) { + z, err := Zeros(Int, 2, 2) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + if v, _ := IntAt(z, 1, 1); v != 0 { + t.Fatalf("Zeros value: %d", v) + } + + o, err := Ones(Float, 3) + if err != nil { + t.Fatalf("Ones: %v", err) + } + if v, _ := FloatAt(o, 2); v != 1.0 { + t.Fatalf("Ones value: %v", v) + } + + fi, err := FullI(7, 2) + if err != nil { + t.Fatalf("FullI: %v", err) + } + if v, _ := IntAt(fi, 0); v != 7 { + t.Fatalf("FullI value: %d", v) + } + + ff, err := FullF(2.5, 2, 2) + if err != nil { + t.Fatalf("FullF: %v", err) + } + if v, _ := FloatAt(ff, 1, 0); v != 2.5 { + t.Fatalf("FullF value: %v", v) + } +} + +func TestRange(t *testing.T) { + r := mustFromInts(t, []int64{0, 1, 2, 3}, 4) + got, err := Range(0, 4) + if err != nil { + t.Fatalf("Range: %v", err) + } + if !Equal(r, got) { + t.Fatalf("Range: %s", got) + } + + by, err := RangeBy(10, 0, -3) + if err != nil { + t.Fatalf("RangeBy: %v", err) + } + want := mustFromInts(t, []int64{10, 7, 4, 1}, 4) + if !Equal(want, by) { + t.Fatalf("RangeBy negative: %s", by) + } + + empty, err := Range(5, 5) + if err != nil { + t.Fatalf("Range empty: %v", err) + } + if empty.Len() != 0 { + t.Fatalf("Range empty len: %d", empty.Len()) + } + + if _, err := RangeBy(0, 5, 0); err == nil || !strings.Contains(err.Error(), "step cannot be zero") { + t.Fatalf("RangeBy zero step: %v", err) + } + + // The walk stops instead of looping when the next value would wrap. + big, err := RangeBy(9223372036854775805, 9223372036854775807, 2) + if err != nil { + t.Fatalf("RangeBy wrap: %v", err) + } + if big.Len() != 1 { + t.Fatalf("RangeBy wrap len: %d", big.Len()) + } +} + +func TestEqual(t *testing.T) { + a := mustFromInts(t, []int64{1, 2}, 2) + b := mustFromInts(t, []int64{1, 2}, 2) + c := mustFromInts(t, []int64{1, 2, 3}, 3) + f := mustFromFloats(t, []float64{1, 2}, 2) + + if !Equal(a, b) { + t.Fatalf("equal arrays must compare equal") + } + if Equal(a, c) { + t.Fatalf("different shapes must not compare equal") + } + if Equal(a, f) { + t.Fatalf("int 1 must not equal float 1.0") + } + if Equal(a, nil) { + t.Fatalf("an array never equals nil") + } +} + +func TestImmutability(t *testing.T) { + vals := []int64{1, 2, 3} + a := mustFromInts(t, vals, 3) + + // The constructor copies: later changes to vals never reach the array. + vals[0] = 99 + if v, _ := IntAt(a, 0); v != 1 { + t.Fatalf("FromInts must copy: got %d", v) + } + + // Shape returns a copy, so mutating it cannot corrupt the array. + shape := a.Shape() + shape[0] = 100 + if a.Shape()[0] != 3 { + t.Fatalf("Shape must copy: got %d", a.Shape()[0]) + } +} + +func TestString(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + if got := a.String(); got != "int (2, 2) [1, 2, 3, 4]" { + t.Fatalf("String: %q", got) + } + f := mustFromFloats(t, []float64{1.5, 2}, 2) + if got := f.String(); got != "float (2) [1.5, 2]" { + t.Fatalf("String float: %q", got) + } +} + +func TestStringMultidim(t *testing.T) { + // Ranks 1 and 2 stay flat; rank 3+ wraps per trailing dimension, + // with size-1 dimensions transparent. + c, _ := FromInts([]int64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2) + want := "int (2, 2, 2) [[[1, 2], [3, 4]], [[5, 6], [7, 8]]]" + if got := c.String(); got != want { + t.Errorf("3-D:\n got %q\nwant %q", got, want) + } + s, _ := FromFloats([]float64{7, 9, 11}, 1, 1, 3) + if got := s.String(); got != "float (1, 1, 3) [7, 9, 11]" { + t.Errorf("size-1 dims: %q", got) + } +} diff --git a/internal/core/validation_pins_test.go b/internal/core/validation_pins_test.go new file mode 100644 index 0000000..0be75b0 --- /dev/null +++ b/internal/core/validation_pins_test.go @@ -0,0 +1,492 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Regression pins for input validation in internal/core: one test per +// defect, named after what it pins; the radix chunk split also carries +// a deterministic invariant test. + +package core + +import ( + "fmt" + "math" + "slices" + "strings" + "sync" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// callNoPanic runs fn on the test goroutine and reports a panic as a +// test failure: a validation gap must surface as an error the caller can +// handle, never as a crash inside the library. +func callNoPanic(t *testing.T, what string, fn func() error) error { + t.Helper() + var err error + func() { + defer func() { + if r := recover(); r != nil { + t.Fatalf("%s panicked instead of returning an error: %v", what, r) + } + }() + err = fn() + }() + return err +} + +// Pad used to accept negative pad values, drive the new shape +// negative and die in alloc with "makeslice: len out of range". Pad +// documents an error for a malformed pad argument, so the negatives have +// to be refused by name. +func TestPadRejectsNegativePadValues(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 3) + err := callNoPanic(t, "Pad(-2, -2)", func() error { + _, err := Pad(a, []int{-2, -2}, "constant", 0) + return err + }) + if err == nil { + t.Fatal("Pad accepted the negative pad pair (-2, -2)") + } + if !strings.Contains(err.Error(), "non-negative") { + t.Fatalf("Pad error does not name the negative pads: %v", err) + } + + // One negative side of a pair is just as malformed. + err = callNoPanic(t, "Pad(1, -1)", func() error { + _, err := Pad(a, []int{1, -1}, "constant", 0) + return err + }) + if err == nil { + t.Fatal("Pad accepted the mixed pad pair (1, -1)") + } + + // The valid path is untouched. + ok, err := Pad(a, []int{1, 1}, "constant", 0) + if err != nil { + t.Fatalf("Pad valid pair: %v", err) + } + if !slices.Equal(ok.Shape(), []int{5}) { + t.Fatalf("Pad valid pair shape %v, want [5]", ok.Shape()) + } +} + +// TruncatedNormal with std = NaN used to pass the `std <= 0` +// guard, leave the rejection window NaN and spin in the retry loop +// forever. The draw must fall back to the degenerate all-zero result +// promptly, on this goroutine: the watchdog keeps a regression from +// hanging the suite. +func TestTruncatedNormalNaNStdReturnsPromptly(t *testing.T) { + done := make(chan *Array, 1) + go func() { + done <- TruncatedNormal(NewGenerator(1), []int{3}, 0, math.NaN()) + }() + select { + case got := <-done: + if got == nil { + t.Fatal("TruncatedNormal(std=NaN) returned nil") + } + for i, v := range got.RawFloat32s() { + if v != 0 { + t.Fatalf("TruncatedNormal(std=NaN) value %d = %v, want the degenerate 0", i, v) + } + } + case <-time.After(5 * time.Second): + t.Fatal("TruncatedNormal(std=NaN) still running after 5s: the rejection loop cannot exit") + } + + // The documented degenerate path (std <= 0) still returns zeros, and + // a valid std still draws inside the window. + zero := TruncatedNormal(NewGenerator(2), []int{4}, 0, 0) + if zero == nil { + t.Fatal("TruncatedNormal(std=0) returned nil") + } + for i, v := range zero.RawFloat32s() { + if v != 0 { + t.Fatalf("TruncatedNormal(std=0) value %d = %v, want 0", i, v) + } + } + drawn := TruncatedNormal(NewGenerator(3), []int{64}, 0, 1) + if drawn == nil { + t.Fatal("TruncatedNormal(std=1) returned nil") + } + for i, v := range drawn.RawFloat32s() { + if v < -2 || v > 2 { + t.Fatalf("TruncatedNormal(std=1) value %d = %v outside the +-2 sigma window", i, v) + } + } +} + +// Normal let a NaN std through (`std < 0` is false for NaN) and +// returned an array of NaNs silently. This is the feeder of the TruncatedNormal case above, so it +// has to be a loud error. +func TestNormalRejectsNaNStd(t *testing.T) { + g := NewGenerator(1) + if _, err := Normal(g, 3, 0, math.NaN()); err == nil { + t.Fatal("Normal accepted a NaN std and drew silently wrong values") + } + if _, err := Normal(g, 3, 0, -1); err == nil { + t.Fatal("Normal accepted a negative std") + } + arr, err := Normal(g, 3, 0, 1) + if err != nil { + t.Fatalf("Normal std=1: %v", err) + } + for i, v := range arr.RawFloats() { + if math.IsNaN(v) { + t.Fatalf("Normal std=1 value %d is NaN", i) + } + } +} + +// InterpolateGrid with a NaN query used to pass both clamps, +// convert to the platform's indefinite integer and panic inside +// FloatAt; with a stride above 1 it silently returned a value read from +// a wrapped index instead. The documented clamp must hold or the call +// must be refused. +Inf and -Inf keep clamping, as documented. +func TestInterpolateGridRejectsNaNQuery(t *testing.T) { + grid := mustFloats(t, []float64{0, 10, 20, 30}, 4) + origins := []float64{0} + steps := []float64{1} + nan := mustFloats(t, []float64{math.NaN()}, 1, 1) + err := callNoPanic(t, "InterpolateGrid(NaN)", func() error { + _, err := InterpolateGrid(grid, origins, steps, nan) + return err + }) + if err == nil { + t.Fatal("InterpolateGrid accepted a NaN query") + } + if !strings.Contains(err.Error(), "NaN") { + t.Fatalf("InterpolateGrid error does not name the NaN position: %v", err) + } + + // The infinite queries keep their clamp: -Inf to the first sample, + // +Inf to the last. + inf := mustFloats(t, []float64{math.Inf(-1), math.Inf(1)}, 2, 1) + out, err := InterpolateGrid(grid, origins, steps, inf) + if err != nil { + t.Fatalf("InterpolateGrid with infinite queries: %v", err) + } + if got := out.FloatAt(0); got != 0 { + t.Fatalf("-Inf query = %v, want the first sample 0", got) + } + if got := out.FloatAt(1); got != 30 { + t.Fatalf("+Inf query = %v, want the last sample 30", got) + } +} + +// HaltonPoints had no skip + n bound, so an index near MaxInt +// wrapped the int arithmetic and every point collapsed to the origin +// with no error. The constructor now enforces the same 2^32 index +// budget SobolPoints does. +func TestHaltonPointsRejectsSkipOverflow(t *testing.T) { + if _, err := HaltonPoints(2, 2, math.MaxInt); err == nil { + t.Fatal("HaltonPoints accepted skip = MaxInt: the points collapse to the origin silently") + } + // The boundary the twin constructor enforces: n + skip must stay + // below 2^32, so the last acceptable index is 2^32 - 1. + if _, err := HaltonPoints(2, 2, 1<<32-2); err == nil { + t.Fatal("HaltonPoints accepted n + skip = 2^32") + } + if _, err := HaltonPoints(2, 2, 1<<32-3); err != nil { + t.Fatalf("HaltonPoints rejected n + skip = 2^32 - 1: %v", err) + } + if _, err := HaltonPoints(1<<32, 2, 0); err == nil { + t.Fatal("HaltonPoints accepted n = 2^32") + } + // The ordinary path is untouched. + if _, err := HaltonPoints(4, 2, 3); err != nil { + t.Fatalf("HaltonPoints valid skip: %v", err) + } +} + +// Unsqueeze(-1) appends the new axis at the end (the convention +// the code has always followed and the frozen API keeps); the doc +// comment claimed it inserts before the last dimension. The behaviour +// is pinned here so the comment and the code cannot drift apart again +// in either direction. +func TestUnsqueezeNegativeDimAppendsAxis(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + got, err := Unsqueeze(a, -1) + if err != nil { + t.Fatalf("Unsqueeze(-1): %v", err) + } + if !slices.Equal(got.Shape(), []int{2, 3, 1}) { + t.Fatalf("Unsqueeze((2,3), -1) shape %v, want the appended [2 3 1]", got.Shape()) + } + // Counting from the end of the result rank: -2 lands one axis earlier. + got, err = Unsqueeze(a, -2) + if err != nil { + t.Fatalf("Unsqueeze(-2): %v", err) + } + if !slices.Equal(got.Shape(), []int{2, 1, 3}) { + t.Fatalf("Unsqueeze((2,3), -2) shape %v, want [2 1 3]", got.Shape()) + } + // One step past the rank is refused. + if _, err := Unsqueeze(a, -4); err == nil { + t.Fatal("Unsqueeze(-4) accepted for a rank-2 input") + } +} + +// seekCoord was dead code with no caller in the module. Removing +// it must not move a single element, so BroadcastTo is pinned against +// the naive per-element odometer reference the run-filled fast path +// replaced. +func TestBroadcastToMatchesOdometerReference(t *testing.T) { + cases := []struct { + name string + vals []int64 + srcSh []int + target []int + }{ + {"column", []int64{1, 2, 3}, []int{3, 1}, []int{3, 2}}, + {"prepend", []int64{1, 2, 3}, []int{3, 1}, []int{2, 3, 1}}, + {"scalar", []int64{7}, []int{1}, []int{2, 3, 4}}, + {"bias", []int64{1, 2, 3, 4, 5, 6}, []int{1, 3, 2}, []int{4, 3, 2}}, + {"same-shape", []int64{9, 8, 7, 6}, []int{2, 2}, []int{2, 2}}, + {"trailing-run", []int64{1, 2}, []int{2, 1}, []int{2, 3}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + src := mustFromInts(t, tc.vals, tc.srcSh...) + got, err := BroadcastTo(src, tc.target...) + if err != nil { + t.Fatalf("BroadcastTo: %v", err) + } + want := broadcastOdometerReference(t, src, tc.target) + if !slices.Equal(want, got.RawInts()) { + t.Fatalf("BroadcastTo diverged from the odometer reference:\n got %v\nwant %v", + got.RawInts(), want) + } + }) + } +} + +// broadcastOdometerReference expands src to target one element at a +// time, coordinates recomputed from scratch: the slow path BroadcastTo's +// run fill replaced. +func broadcastOdometerReference(t *testing.T, src *Array, target []int) []int64 { + t.Helper() + total := 1 + for _, d := range target { + total *= d + } + sh := src.Shape() + out := make([]int64, total) + coord := make([]int, len(target)) + srcCoord := make([]int, len(sh)) + off := len(target) - len(sh) + for i := range total { + for d := range sh { + c := 0 + if sh[d] != 1 { + c = coord[off+d] + } + srcCoord[d] = c + } + v, err := IntAt(src, srcCoord...) + if err != nil { + t.Fatalf("IntAt(%v): %v", srcCoord, err) + } + out[i] = v + advanceOdometer(coord, target) + } + return out +} + +// Radix chunk split, invariant half: the parallel radix derives its histogram rows +// from its own split of [0, n) instead of re-reading the global worker +// count the way engine.ParallelMin did, so a concurrent SetNumCPU can +// move work between goroutines but never two live chunks onto one row. +// This test reproduces the disagreement deterministically and with no +// reliance on scheduling: the stale-snapshot case raises the live worker +// count while the chunk width handed to the split stays the one a caller +// would have snapshotted earlier, which is exactly the window SetNumCPU +// opens (the raising read is what used to collapse two live chunks onto +// one row). Every row must be handed out once with the range it owns, +// start = row*chunk, and the ranges must partition [0, n); the remaining +// cases pin that contract either side of the radixParMin spawn floor. +func TestRadixChunksRowMatchesOwnedRange(t *testing.T) { + cases := []struct { + name string + n int + workers int + chunk int // the caller's snapshot width, stale where noted + }{ + {"exact-fit", 40_000, 2, 20_000}, + {"ragged", 40_001, 3, 13_334}, + {"one-chunk", 40_000, 1, 40_000}, + {"floor-exact", 40_000, 1, radixParMin}, + {"stale-snapshot", 40_000, 8, 20_000}, // width from a 2-worker snapshot + {"below-floor", 40_000, 8, radixParMin - 1}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + prev := engine.SetNumWorkers(tc.workers) + defer engine.SetNumWorkers(prev) + + parallel := tc.chunk >= radixParMin && tc.chunk < tc.n + nchunks := 1 + if parallel { + nchunks = (tc.n + tc.chunk - 1) / tc.chunk + } + // The callback may run on a spawned goroutine, so failures + // are collected and reported on the test goroutine. + var mu sync.Mutex + var problems []string + touched := make([]int, tc.n) + rows := make([]bool, nchunks) + forEachRadixChunk(tc.n, tc.chunk, func(row, start, end int) { + var local []string + if row < 0 || row >= nchunks { + local = append(local, fmt.Sprintf("row %d out of range for %d chunks", row, nchunks)) + } else if rows[row] { + local = append(local, fmt.Sprintf("row %d handed out twice", row)) + } else { + rows[row] = true + } + // The row identity is the chunk position: the range a + // worker owns and the histogram row it writes are the + // same object, so they cannot disagree. + wantStart, wantEnd := 0, tc.n + if parallel { + wantStart, wantEnd = row*tc.chunk, min(row*tc.chunk+tc.chunk, tc.n) + } + if start != wantStart || end != wantEnd { + local = append(local, fmt.Sprintf("row %d owns [%d, %d), want [%d, %d)", + row, start, end, wantStart, wantEnd)) + if start < 0 || end > tc.n || start >= end { + local = append(local, fmt.Sprintf("row %d owns the illegal range [%d, %d)", row, start, end)) + } + } + if start >= 0 && end <= tc.n && start < end { + for i := start; i < end; i++ { + touched[i]++ + } + } + mu.Lock() + problems = append(problems, local...) + mu.Unlock() + }) + if len(problems) > 0 { + t.Fatalf("chunk split broken: %s", strings.Join(problems, "; ")) + } + for i, c := range touched { + if c != 1 { + t.Fatalf("index %d visited %d times", i, c) + } + } + for row, seen := range rows { + if !seen { + t.Fatalf("row %d never ran", row) + } + } + }) + } +} + +// flipWorkers runs fn at least rounds times while a second goroutine +// flips the global worker count between low and high. SetNumCPU is +// documented as safe to call at any time, so a kernel that derives its +// chunking from the global count must still finish correctly; before +// Radix chunk split, growth inside the kernel's own window collapsed two live chunks +// onto one histogram row and corrupted the output. +func flipWorkers(t *testing.T, low, high, rounds int, fn func(round int)) { + t.Helper() + prev := engine.SetNumWorkers(low) + defer engine.SetNumWorkers(prev) + + stop := make(chan struct{}) + var flipper sync.WaitGroup + flipper.Go(func() { + w := low + for { + select { + case <-stop: + return + default: + } + if w == low { + w = high + } else { + w = low + } + engine.SetNumWorkers(w) + } + }) + defer func() { + close(stop) + flipper.Wait() + }() + + for round := range rounds { + fn(round) + } +} + +// Radix chunk split, value-radix half: Sort of a large float64 payload must stay an +// ascending permutation while a concurrent SetNumCPU moves the worker +// count from 2 to 8 and back. The output is checked against the sorted +// order itself, so a lost histogram update (rows colliding), a +// mis-ordered scatter (bit-identity broken across chunks) or an +// out-of-range write all fail here. +func TestSortRadixStableUnderConcurrentSetNumCPU(t *testing.T) { + const n = 200_000 + g := NewGenerator(7) + src, err := Floats(g, n) + if err != nil { + t.Fatalf("Floats: %v", err) + } + template := slices.Clone(src.RawFloats()) + // Ties force the scatter to be exercised across chunk boundaries. + for i := range template { + if i%3 == 0 { + template[i] = float64(i % 17) + } + } + reference := sortReference(template) + + flipWorkers(t, 2, 8, 8, func(round int) { + a := mustFromFloats(t, slices.Clone(template), n) + got, err := Sort(a) + if err != nil { + t.Fatalf("round %d: Sort: %v", round, err) + } + vals := got.RawFloats() + if !equalSortedFloats(reference, vals) { + t.Fatalf("round %d: Sort corrupted the parallel radix output", round) + } + }) +} + +// Radix chunk split, permutation-radix half: ArgSort must keep returning the same +// stable permutation the serial radix returns while the worker count +// changes concurrently. A row collision shows up either as a repeated +// or missing index or as a change of the permutation itself. +func TestArgSortRadixStableUnderConcurrentSetNumCPU(t *testing.T) { + const n = 200_000 + g := NewGenerator(11) + src, err := Floats(g, n) + if err != nil { + t.Fatalf("Floats: %v", err) + } + vals := slices.Clone(src.RawFloats()) + for i := range vals { + if i%3 == 0 { + vals[i] = float64(i % 17) + } + } + reference := argSortReference(vals) + + flipWorkers(t, 2, 8, 8, func(round int) { + a := mustFromFloats(t, slices.Clone(vals), n) + got, err := ArgSort(a) + if err != nil { + t.Fatalf("round %d: ArgSort: %v", round, err) + } + if !slices.Equal(reference, got.RawInts()) { + t.Fatalf("round %d: ArgSort diverged from the stable serial permutation", round) + } + }) +} diff --git a/internal/core/view_payload_reads_test.go b/internal/core/view_payload_reads_test.go new file mode 100644 index 0000000..9df9cd5 --- /dev/null +++ b/internal/core/view_payload_reads_test.go @@ -0,0 +1,150 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// Regression tests: entry points that read a +// payload past the logical extent of a view, which panicked on any +// array that a Slice had rebased. + +// mustSlice returns a view of a ranked by dim, the constructor the +// callers below all use. +func mustSlice(t *testing.T, a *Array, dim, start, stop int) *Array { + t.Helper() + v, err := Slice(a, dim, start, stop) + if err != nil { + t.Fatalf("Slice: %v", err) + } + return v +} + +// TestPowIOnViews pins PowI for every dtype against a rebased view: +// the payload of a view is longer than its length, and the kernel must +// stay inside the extent. +func TestPowIOnViews(t *testing.T) { + t.Run("float64", func(t *testing.T) { + a := mustFloats(t, []float64{1, 2, 3, 4, 5, 6}, 6) + v := mustSlice(t, a, 0, 2, 5) + out, err := PowI(v, 3) + if err != nil { + t.Fatalf("PowI: %v", err) + } + want := []float64{27, 64, 125} + for i, w := range want { + if got := out.FloatAt(i); got != w { + t.Fatalf("PowI[%d] = %v, want %v", i, got, w) + } + } + }) + t.Run("int64", func(t *testing.T) { + a, err := FromInts([]int64{1, 2, 3, 4, 5, 6}, 6) + if err != nil { + t.Fatal(err) + } + v := mustSlice(t, a, 0, 1, 3) + out, err := PowI(v, 2) + if err != nil { + t.Fatalf("PowI: %v", err) + } + want := []int64{4, 9} + for i, w := range want { + if got := out.RawInts()[i]; got != w { + t.Fatalf("PowI[%d] = %v, want %v", i, got, w) + } + } + }) + t.Run("float32", func(t *testing.T) { + a, err := FromFloat32s([]float32{1, 2, 3, 4, 5, 6}, 6) + if err != nil { + t.Fatal(err) + } + v := mustSlice(t, a, 0, 1, 3) + out, err := PowI(v, 2) + if err != nil { + t.Fatalf("PowI: %v", err) + } + want := []float32{4, 9} + for i, w := range want { + if got := out.RawFloat32s()[i]; got != w { + t.Fatalf("PowI[%d] = %v, want %v", i, got, w) + } + } + }) + t.Run("complex128", func(t *testing.T) { + a, err := FromComplexes([]complex128{1i, 2i, 3i, 4i}, 4) + if err != nil { + t.Fatal(err) + } + v := mustSlice(t, a, 0, 1, 3) + out, err := PowI(v, 2) + if err != nil { + t.Fatalf("PowI: %v", err) + } + // (2i)² = −4, (3i)² = −9. + want := []complex128{-4, -9} + for i, w := range want { + if got := out.RawComplexes()[i]; got != w { + t.Fatalf("PowI[%d] = %v, want %v", i, got, w) + } + } + }) +} + +// TestInterpolate2DOnWidenedView pins the query arrays of Interpolate2D +// when they are float32 views: widening walks the element count, not +// the payload. +func TestInterpolate2DOnWidenedView(t *testing.T) { + grid := mustFloats(t, []float64{0, 1, 2, 0, 1, 2}, 2, 3) + xs, err := FromFloat32s([]float32{0, 1, 2, 0.5, 1.5}, 5) + if err != nil { + t.Fatal(err) + } + ys, err := FromFloat32s([]float32{0, 1, 0, 0.5, 1.5}, 5) + if err != nil { + t.Fatal(err) + } + xv := mustSlice(t, xs, 0, 0, 3) + yv := mustSlice(t, ys, 0, 0, 3) + out, err := Interpolate2D(grid, xv, yv, 0, 0, 1, 1) + if err != nil { + t.Fatalf("Interpolate2D: %v", err) + } + want := []float64{0, 1, 2} + if out.Len() != len(want) { + t.Fatalf("Interpolate2D returned %d values, want %d", out.Len(), len(want)) + } + for i, w := range want { + if got := out.FloatAt(i); math.Abs(got-w) > 1e-12 { + t.Fatalf("value %d = %v, want %v", i, got, w) + } + } +} + +// TestInterpolate2DOnIntView covers the int widening branch too. +func TestInterpolate2DOnIntView(t *testing.T) { + grid := mustFloats(t, []float64{0, 1, 2, 0, 1, 2}, 2, 3) + xs, err := FromInts([]int64{0, 1, 2, 0, 2}, 5) + if err != nil { + t.Fatal(err) + } + ys, err := FromInts([]int64{0, 1, 0, 0, 1}, 5) + if err != nil { + t.Fatal(err) + } + xv := mustSlice(t, xs, 0, 0, 3) + yv := mustSlice(t, ys, 0, 0, 3) + out, err := Interpolate2D(grid, xv, yv, 0, 0, 1, 1) + if err != nil { + t.Fatalf("Interpolate2D: %v", err) + } + for i, w := range []float64{0, 1, 2} { + if got := out.FloatAt(i); got != w { + t.Fatalf("value %d = %v, want %v", i, got, w) + } + } +} diff --git a/internal/core/view_test.go b/internal/core/view_test.go new file mode 100644 index 0000000..f6285ea --- /dev/null +++ b/internal/core/view_test.go @@ -0,0 +1,347 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "testing" +) + +// The staged views design (docs/ARCHITECTURE.md) starts with the +// representation: an array may alias another's storage, either as a +// contiguous region (payload rebased, strides nil; invisible to every +// kernel) or as a strided view (payload rebased plus explicit strides). +// These tests pin the machinery before Slice starts producing views. + +// stridedColumnViewOf builds the column slice [c0, c1) of a 2-D array +// as a strided view, the shape a non-contiguous Slice will produce. +func stridedColumnViewOf(a *Array, c0, c1 int) *Array { + rows, cols := a.Shape()[0], a.Shape()[1] + view := &Array{ + shape: []int{rows, c1 - c0}, + dt: a.Dtype(), + floats: a.RawFloats()[c0:], + strides: []int{cols, 1}, + } + return view +} + +// TestStridedViewAccessors pins the accessor contract on a strided +// view: element i must come from payload[Σ i_d·strides[d]], never from +// payload[i]. +func TestStridedViewAccessors(t *testing.T) { + base, err := FromFloats([]float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + }, 3, 4) + if err != nil { + t.Fatal(err) + } + view := stridedColumnViewOf(base, 1, 3) // columns 1..2 + + if view.Len() != 6 { + t.Fatalf("view len %d, want 6", view.Len()) + } + if view.isContiguous() { + t.Fatal("strided view reported contiguous") + } + // Row-major over (3, 2): [[2,3],[6,7],[10,11]]. + want := []float64{2, 3, 6, 7, 10, 11} + for i, w := range want { + if got := view.FloatAt(i); got != w { + t.Errorf("view.FloatAt(%d) = %v, want %v", i, got, w) + } + } + // The multi-dimensional accessor agrees. + for r := range 3 { + for c := range 2 { + if got, _ := FloatAt(view, r, c); got != want[r*2+c] { + t.Errorf("FloatAt(view, %d, %d) = %v, want %v", r, c, got, want[r*2+c]) + } + } + } + if got := view.String(); got != "float (3, 2) [2, 3, 6, 7, 10, 11]" { + t.Errorf("view.String() = %q", got) + } +} + +// TestStridedViewMaterialise pins that a view copies out its own +// elements, densely and in view order, and that the copy is independent +// of the aliased storage. +func TestStridedViewMaterialise(t *testing.T) { + base, err := FromFloats([]float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + }, 3, 4) + if err != nil { + t.Fatal(err) + } + view := stridedColumnViewOf(base, 1, 3) + + dense := view.materialise() + if !dense.isContiguous() { + t.Fatal("materialised copy still carries strides") + } + if dense.Len() != 6 || dense.Shape()[0] != 3 || dense.Shape()[1] != 2 { + t.Fatalf("materialised shape %s len %d", shapeText(dense.Shape()), dense.Len()) + } + want := []float64{2, 3, 6, 7, 10, 11} + for i, w := range want { + if got := dense.FloatAt(i); got != w { + t.Errorf("materialised[%d] = %v, want %v", i, got, w) + } + } + // The copy owns its payload: the view is still the base's columns. + if base.FloatAt(1) != 2 { + t.Fatalf("base mutated by materialise: %v", base.FloatAt(1)) + } +} + +// TestMaterialiseContiguousIsIdentity pins the fast path: a contiguous +// array is returned untouched, so callers can materialise unconditionally +// without paying a copy. +func TestMaterialiseContiguousIsIdentity(t *testing.T) { + a, err := FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatal(err) + } + if !a.isContiguous() { + t.Fatal("fresh array is not contiguous") + } + if a.materialise() != a { + t.Fatal("contiguous materialise copied instead of returning the receiver") + } +} + +// TestContiguousRebasedViewIsTransparent pins the property the kernels +// lean on: a contiguous region becomes a view by rebasing the payload +// with nil strides, so payload[i] is element i exactly as before and no +// kernel needs to know a view exists. +func TestContiguousRebasedViewIsTransparent(t *testing.T) { + base, err := FromInts([]int64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + }, 3, 4) + if err != nil { + t.Fatal(err) + } + // Row 1 of the base, as a (1, 4) 2-D view. + row := &Array{ + shape: []int{1, 4}, + dt: Int, + ints: base.RawInts()[4:8], + } + if !row.isContiguous() { + t.Fatal("rebased row view is not contiguous") + } + if row.Len() != 4 { + t.Fatalf("row view len %d, want 4", row.Len()) + } + for i, w := range []int64{5, 6, 7, 8} { + if got, _ := IntAt(row, 0, i); got != w { + t.Errorf("IntAt(row, 0, %d) = %d, want %d", i, got, w) + } + } + // Direct payload reads are valid for a contiguous view: this is the + // guarantee that keeps every kernel working unchanged. + if row.RawInts()[0] != 5 || row.RawInts()[3] != 8 { + t.Fatalf("direct payload read on a contiguous view is wrong: %v", row.RawInts()[:4]) + } + // A copy of a view carries only the view's elements. + cp := Copy(row) + if cp.Len() != 4 || len(cp.RawInts()) != 4 { + t.Fatalf("Copy(view) len %d payload %d, want 4", cp.Len(), len(cp.RawInts())) + } +} + +// TestPhysIndexMapping pins the index mapping directly, including the +// contiguous fast path. +func TestPhysIndexMapping(t *testing.T) { + contiguous := &Array{shape: []int{2, 3}, dt: Int, ints: make([]int64, 6)} + for i := range 6 { + if got := contiguous.physIndex(i); got != i { + t.Errorf("contiguous physIndex(%d) = %d, want identity", i, got) + } + } + + // (2, 2) view over a 3-wide base starting at column 1: + // element (r, c) is at payload[r*3 + c]. + view := &Array{shape: []int{2, 2}, dt: Int, ints: make([]int64, 6), strides: []int{3, 1}} + want := []int{0, 1, 3, 4} + for i, w := range want { + if got := view.physIndex(i); got != w { + t.Errorf("view.physIndex(%d) = %d, want %d", i, got, w) + } + } +} + +// TestStridedViewComplexAccessors pins complex and int views too, since +// each dtype has its own payload field. +func TestStridedViewComplexAccessors(t *testing.T) { + base, err := FromComplexes([]complex128{ + 1 + 1i, 2 + 2i, + 3 + 3i, 4 + 4i, + }, 2, 2) + if err != nil { + t.Fatal(err) + } + // Column 1 view of a 2x2 base. + view := &Array{ + shape: []int{2, 1}, + dt: Complex, + complexes: base.RawComplexes()[1:], + strides: []int{2, 1}, + } + if got := view.ComplexAt(0); got != 2+2i { + t.Errorf("complexAt(0) = %v, want (2+2i)", got) + } + if got := view.ComplexAt(1); got != 4+4i { + t.Errorf("complexAt(1) = %v, want (4+4i)", got) + } + dense := view.materialise() + if dense.Len() != 2 || dense.ComplexAt(1) != 4+4i { + t.Errorf("materialised complex view: %v", dense) + } +} + +// TestStridedViewFloat32Accessors covers the float32 payload path. +func TestStridedViewFloat32Accessors(t *testing.T) { + base, err := FromFloat32s([]float32{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + // Columns 1..2 of a 3-wide base. + view := &Array{ + shape: []int{2, 2}, + dt: Float32, + floats32: base.RawFloat32s()[1:], + strides: []int{3, 1}, + } + want := []float32{2, 3, 5, 6} + for i, w := range want { + if got := view.float32At(i); got != w { + t.Errorf("float32At(%d) = %v, want %v", i, got, w) + } + if got := view.FloatAt(i); math.Abs(got-float64(w)) > 1e-12 { + t.Errorf("floatAt(%d) = %v, want %v", i, got, w) + } + } +} + +// TestSliceReturnsContiguousView pins stage (b): a leading-axis slice is +// a view sharing the source's storage, not a copy. +func TestSliceReturnsContiguousView(t *testing.T) { + base, err := FromInts([]int64{ + 1, 2, 3, + 4, 5, 6, + 7, 8, 9, + }, 3, 3) + if err != nil { + t.Fatal(err) + } + view, err := Slice(base, 0, 1, 3) + if err != nil { + t.Fatal(err) + } + if view.Len() != 6 || view.Shape()[0] != 2 { + t.Fatalf("slice shape %s len %d", shapeText(view.Shape()), view.Len()) + } + want := []int64{4, 5, 6, 7, 8, 9} + for i, w := range want { + if got := view.RawInts()[i]; got != w { + t.Errorf("view payload[%d] = %d, want %d", i, got, w) + } + } + // Aliasing proof: a write into the base's storage (which only this + // test does; the library never writes through an array it did not + // allocate) is visible through the view. The view starts at row 1, + // so its payload is base.RawInts()[3:]. + base.RawInts()[3] = 111 + if got, _ := IntAt(view, 0, 0); got != 111 { + t.Fatalf("view does not alias the base: got %d", got) + } + // And the view is indistinguishable from a dense array to a kernel. + if !view.isContiguous() { + t.Fatal("slice view is not contiguous") + } +} + +// TestSliceFullDimensionReturnsView pins the second contiguous case: a +// slice that takes a dimension in full is a view too. +func TestSliceFullDimensionReturnsView(t *testing.T) { + base, err := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + view, err := Slice(base, 1, 0, 3) // the whole of dim 1 + if err != nil { + t.Fatal(err) + } + for i := range 6 { + if got := view.RawFloats()[i]; got != float64(i+1) { + t.Fatalf("view[%d] = %v, want %d", i, got, i+1) + } + } + base.RawFloats()[2] = 42 + if got := view.RawFloats()[2]; got != 42 { + t.Fatalf("full-dimension slice does not alias: %v", got) + } +} + +// TestSliceInteriorRangeCopies pins the deliberate boundary: an interior +// range is not a contiguous region, so Slice copies it: a strided view +// there would need the kernel-boundary materialisation audit. +func TestSliceInteriorRangeCopies(t *testing.T) { + base, err := FromInts([]int64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + }, 2, 4) + if err != nil { + t.Fatal(err) + } + out, err := Slice(base, 1, 1, 3) + if err != nil { + t.Fatal(err) + } + if !out.isContiguous() { + t.Fatal("interior slice is not contiguous") + } + // Columns 1..2: [[2,3],[6,7]]. + want := []int64{2, 3, 6, 7} + for i, w := range want { + if got := out.RawInts()[i]; got != w { + t.Errorf("interior slice[%d] = %d, want %d", i, got, w) + } + } + base.RawInts()[1] = 99 + if got := out.RawInts()[0]; got != 2 { + t.Fatalf("copy aliases the base: %d", got) + } +} + +// TestSliceOfViewStaysCorrect pins that slicing a view materialises the +// source first, so the copy path never reads the wrong payload slots. +func TestSliceOfViewStaysCorrect(t *testing.T) { + base, err := FromInts([]int64{1, 2, 3, 4, 5, 6}, 3, 2) + if err != nil { + t.Fatal(err) + } + view, err := Slice(base, 0, 1, 3) // rows 1..2, a view + if err != nil { + t.Fatal(err) + } + inner, err := Slice(view, 1, 0, 1) // first column of the view + if err != nil { + t.Fatal(err) + } + want := []int64{3, 5} + for i, w := range want { + if got := inner.RawInts()[i]; got != w { + t.Errorf("slice of view[%d] = %d, want %d", i, got, w) + } + } +} diff --git a/internal/core/walks_pin_test.go b/internal/core/walks_pin_test.go new file mode 100644 index 0000000..c2d1d2a --- /dev/null +++ b/internal/core/walks_pin_test.go @@ -0,0 +1,180 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import ( + "math" + "slices" + "strings" + "testing" +) + +// Pins for the parallel payload walks: strided operands must be read in +// logical order through the accessors, Take must report the smallest +// offending position, and the radix ArgSort permutation must equal the +// stable comparator permutation on tie-heavy inputs. + +func pinFloats(t *testing.T, vals []float64, shape ...int) *Array { + t.Helper() + a, err := FromFloats(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +func pinInts(t *testing.T, vals []int64, shape ...int) *Array { + t.Helper() + a := New(Int, shape...) + copy(a.RawInts(), vals) + return a +} + +// pinStridedFloat builds a test-mechanics strided array: logical (r, c) +// reads payload[r*rowStride + c]. The payload carries an invisible tail +// element at an unaddressed slot, the state a raw payload walk gets +// wrong. +func pinStridedFloat(t *testing.T, payload []float64, shape []int, strides []int) *Array { + t.Helper() + return &Array{shape: shape, dt: Float, floats: payload, strides: strides} +} + +func pinStridedInt(t *testing.T, payload []int64, shape []int, strides []int) *Array { + t.Helper() + return &Array{shape: shape, dt: Int, ints: payload, strides: strides} +} + +func TestArgwhereNonzeroStridedUseLogicalElements(t *testing.T) { + // Logical window [[1, 0], [0, 5]] at strides [3, 1]: the elements + // sit at payload slots 0, 1, 3, 4, and slot 2 holds an invisible + // 99. A payload walk at the logical position reads slot 3 for the + // bottom-right element and calls the 5 a zero. + stridedF := pinStridedFloat(t, []float64{1, 0, 99, 0, 5, 7, 7, 7}, []int{2, 2}, []int{3, 1}) + twinF := pinFloats(t, []float64{1, 0, 0, 5}, 2, 2) + gotF, err := Argwhere(stridedF) + if err != nil { + t.Fatalf("Argwhere(strided float): %v", err) + } + wantF, err := Argwhere(twinF) + if err != nil { + t.Fatalf("Argwhere(twin float): %v", err) + } + if !Equal(gotF, wantF) { + t.Fatalf("Argwhere(strided float) = %v, want %v", gotF, wantF) + } + nzS, err := Nonzero(stridedF) + if err != nil { + t.Fatalf("Nonzero(strided float): %v", err) + } + nzT, err := Nonzero(twinF) + if err != nil { + t.Fatalf("Nonzero(twin float): %v", err) + } + if !slices.Equal(nzS[0], nzT[0]) || !slices.Equal(nzS[1], nzT[1]) { + t.Fatalf("Nonzero(strided float) = %v, want %v", nzS, nzT) + } + + stridedI := pinStridedInt(t, []int64{1, 0, 99, 0, 5, 7, 7, 7}, []int{2, 2}, []int{3, 1}) + twinI := pinInts(t, []int64{1, 0, 0, 5}, 2, 2) + gotI, err := Argwhere(stridedI) + if err != nil { + t.Fatalf("Argwhere(strided int): %v", err) + } + wantI, err := Argwhere(twinI) + if err != nil { + t.Fatalf("Argwhere(twin int): %v", err) + } + if !Equal(gotI, wantI) { + t.Fatalf("Argwhere(strided int) = %v, want %v", gotI, wantI) + } + nzSI, err := Nonzero(stridedI) + if err != nil { + t.Fatalf("Nonzero(strided int): %v", err) + } + nzTI, err := Nonzero(twinI) + if err != nil { + t.Fatalf("Nonzero(twin int): %v", err) + } + if !slices.Equal(nzSI[0], nzTI[0]) || !slices.Equal(nzSI[1], nzTI[1]) { + t.Fatalf("Nonzero(strided int) = %v, want %v", nzSI, nzTI) + } +} + +func TestTakeReportsSmallestOffender(t *testing.T) { + src := pinFloats(t, []float64{10, 11, 12, 13, 14}, 5) + idx := pinInts(t, []int64{99, 88, 0}, 3) + _, err := Take(src, idx) + if err == nil { + t.Fatal("Take accepted out-of-range indices") + } + if !strings.Contains(err.Error(), "position 0") || strings.Contains(err.Error(), "position 1") { + t.Fatalf("Take offender report = %q, want the smallest offending position 0", err.Error()) + } +} + +func TestArgSortTiePermutationsMatchStableReference(t *testing.T) { + for _, n := range []int{64, 1024, 4096} { + vals := make([]float64, n) + ivals := make([]int64, n) + state := uint64(0x9E3779B97F4A7C15 + uint64(n)) + for i := range vals { + state = state*6364136223846793005 + 1442695040888963407 + // A small level set, so ties dominate the permutation. + level := int64(state>>60) % 5 + ivals[i] = level - 2 + vals[i] = float64(level - 2) + } + ref := make([]int, n) + for i := range ref { + ref[i] = i + } + // The stable comparator permutation: ties keep index order, the + // contract a stable digit scatter must reproduce. + slices.SortStableFunc(ref, func(x, y int) int { + switch { + case vals[x] < vals[y]: + return -1 + case vals[x] > vals[y]: + return 1 + } + return 0 + }) + gotF, err := ArgSort(pinFloats(t, vals, n)) + if err != nil { + t.Fatalf("n=%d ArgSort(float): %v", n, err) + } + gotI, err := ArgSort(pinInts(t, ivals, n)) + if err != nil { + t.Fatalf("n=%d ArgSort(int): %v", n, err) + } + // ArgSort answers the permutation as an int array whatever the + // sorted dtype was. + rf, ri := gotF.RawInts(), gotI.RawInts() + for i := range ref { + if int(rf[i]) != ref[i] { + t.Fatalf("n=%d float permutation at %d = %d, want %d (values %v)", n, i, int(rf[i]), ref[i], vals) + } + if int(ri[i]) != ref[i] { + t.Fatalf("n=%d int permutation at %d = %d, want %d (values %v)", n, i, int(ri[i]), ref[i], ivals) + } + } + } + // The NaN and signed-zero contract on the same walk: NaN sorts + // last, -0 folds onto +0 in the value order but keeps its index + // order among the zeros. + withSpecial := pinFloats(t, []float64{2, math.NaN(), math.Copysign(0, -1), -1, 0, math.NaN(), 3}, 7) + got, err := ArgSort(withSpecial) + if err != nil { + t.Fatalf("ArgSort(special): %v", err) + } + perm := got.RawInts() + // Sorted values: -1, then the zeros at indices 2 and 4 in index + // order, then 2, 3, then the NaNs at indices 1 and 5 in index order. + want := []int{3, 2, 4, 0, 6, 1, 5} + for i := range want { + if int(perm[i]) != want[i] { + t.Fatalf("special permutation at %d = %d, want %d (full %v)", i, int(perm[i]), want[i], perm) + } + } +} diff --git a/internal/core/where_half_test.go b/internal/core/where_half_test.go new file mode 100644 index 0000000..3367f8d --- /dev/null +++ b/internal/core/where_half_test.go @@ -0,0 +1,34 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package core + +import "testing" + +// TestWhereHalfCanonicalisesNaN pins the float16 fast path of Where: +// every selected half value narrows through HalfFromFloat64 on the way +// out, so a NaN loses its payload and lands as the canonical 0x7E00 with +// its sign kept, while every finite value survives the round trip bit +// for bit. A payload copy would leak the incoming patterns 0x7C01 and +// 0x7D00 instead. +func TestWhereHalfCanonicalisesNaN(t *testing.T) { + cond := mustFromInts(t, []int64{1, 0, 1, 0}, 4) + // x: a signalling NaN with payload, 2, a negative signalling NaN, and + // a slot the condition does not pick. y: 4, 3, 4 and -1. + x := mustFromHalves(t, []uint16{0x7C01, 0x4000, 0xFE01, 0x3C00}, 4) + y := mustFromHalves(t, []uint16{0x4400, 0x4200, 0x4400, 0xBC00}, 4) + out, err := Where(cond, x, y) + if err != nil { + t.Fatalf("Where: %v", err) + } + if out.Dtype() != Float16 { + t.Fatalf("Where answered dtype %s, want float16", out.Dtype()) + } + want := []uint16{0x7E00, 0x4200, 0xFE00, 0xBC00} + got := out.RawHalves() + for i, w := range want { + if got[i] != w { + t.Fatalf("element %d = %#04x, want %#04x", i, got[i], w) + } + } +} diff --git a/internal/engine/engine.go b/internal/engine/engine.go new file mode 100644 index 0000000..3cd6b4f --- /dev/null +++ b/internal/engine/engine.go @@ -0,0 +1,125 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package engine hosts the private machinery shared by the tensor +// packages: the parallel scheduling primitive, the worker-count policy +// and pooled scratch buffers. It is under internal/: the compiler +// keeps it invisible outside this module. +package engine + +import ( + "runtime" + "sync" +) + +var ( + mu sync.RWMutex + numWorkers = runtime.NumCPU() +) + +// SetNumWorkers sets the number of goroutines the parallel kernels may +// use and returns the previous value. Values below 1 reset to NumCPU. +func SetNumWorkers(n int) int { + mu.Lock() + defer mu.Unlock() + prev := numWorkers + if n < 1 { + n = runtime.NumCPU() + } + numWorkers = n + return prev +} + +// NumWorkers returns the current worker count. +func NumWorkers() int { + mu.RLock() + defer mu.RUnlock() + return numWorkers +} + +// WorkersFor returns the number of goroutines to use for a workload of +// n independent items, bounded by both the worker count and n. +func WorkersFor(n int) int { return max(min(NumWorkers(), n), 1) } + +// Parallel splits the [0, n) index range into chunks and runs fn on +// each chunk in its own goroutine. Fixed chunk size, no per-item +// channel traffic; every chunk owns a disjoint slice of the output, so +// kernels need no locks. A workload the worker policy collapses to a +// single worker (n of 1, or the worker count pinned to 1) runs inline +// on the calling goroutine; any other workload spawns one goroutine +// per worker, so use ParallelMin for a real per-worker floor. +func Parallel(n int, fn func(start, end int)) { ParallelMin(n, 1, fn) } + +// ParallelMin splits the [0, n) index range into chunks and runs fn on +// each chunk in its own goroutine, exactly like Parallel, with one +// extra constraint: while the per-worker chunk would fall below +// minPerWorker, fn runs whole as fn(0, n) on the calling goroutine. A +// worker whose chunk is below the floor costs more to create and +// schedule than the work it carries, so parallelising that workload +// only adds latency; the caller's goroutine is already warm and pays +// nothing to start. Chunk boundaries and the worker choice are +// computed exactly as Parallel computes them, so minPerWorker values +// below 2 reproduce Parallel bit for bit. +func ParallelMin(n, minPerWorker int, fn func(start, end int)) { + w := WorkersFor(n) + if w == 1 { + fn(0, n) + return + } + chunk := (n + w - 1) / w + if chunk < minPerWorker { + fn(0, n) + return + } + var wg sync.WaitGroup + for start := 0; start < n; start += chunk { + end := min(start+chunk, n) + wg.Go(func() { + fn(start, end) + }) + } + wg.Wait() +} + +var float64Pool = sync.Pool{New: func() any { return make([]float64, 0, 1024) }} + +// maxPooledFloat64 caps what the scratch pool keeps. A pooled buffer is +// retained per processor until the next garbage collection, so a kernel +// that borrows hundreds of megabytes would pin that much memory times +// the processor count. Buffers above the cap are dropped on return and +// re-allocated by the next borrower, one allocation per deep chunk; +// every buffer at or below it still round-trips. +const maxPooledFloat64 = 1 << 20 // elements, 8 MiB + +// keepPooled reports whether a returned buffer of the given capacity is +// worth retaining. +func keepPooled(capacity int) bool { return capacity <= maxPooledFloat64 } + +// GetFloat64Buf borrows a float64 buffer of exactly n elements with +// capacity for at least that many. The buffer may be recycled from an +// earlier borrower, so it is cleared before it leaves the pool: every +// slot arrives zero and stays zero until the borrower writes it. That +// makes accumulation kernels safe by construction: stale sums can +// never leak into a result, whichever path the buffer took through the +// pool or the garbage collector. +func GetFloat64Buf(n int) []float64 { + b := float64Pool.Get().([]float64) + if cap(b) < n { + return make([]float64, n) // freshly allocated: already zero + } + b = b[:n] + clear(b) // pooled buffers come back dirty; never hand that on + return b +} + +// PutFloat64Buf returns a borrowed buffer. The backing array is offered +// to the next caller, although sync.Pool may drop it at any garbage +// collection; reuse is opportunistic, never guaranteed. A buffer larger +// than maxPooledFloat64 is dropped outright so that one deep kernel +// cannot pin its scratch memory on every processor. +func PutFloat64Buf(b []float64) { + if !keepPooled(cap(b)) { + return + } + float64Pool.Put(b[:0]) +} diff --git a/internal/engine/engine_bench_test.go b/internal/engine/engine_bench_test.go new file mode 100644 index 0000000..56e8c24 --- /dev/null +++ b/internal/engine/engine_bench_test.go @@ -0,0 +1,25 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package engine + +import ( + "fmt" + "testing" +) + +// BenchmarkPoolRoundTrip measures the full borrow/return cycle at the +// sizes the kernels actually request, including the clear-on-borrow +// cost the pool contract pays. Run before and after touching the pool. +func BenchmarkPoolRoundTrip(b *testing.B) { + for _, n := range []int{64, 4096, 262144} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + buf := GetFloat64Buf(n) + buf[0] = 1 // touch one slot: prove the buffer is writable + PutFloat64Buf(buf) + } + }) + } +} diff --git a/internal/engine/engine_test.go b/internal/engine/engine_test.go new file mode 100644 index 0000000..c518df1 --- /dev/null +++ b/internal/engine/engine_test.go @@ -0,0 +1,171 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package engine + +import ( + "sync" + "testing" +) + +// TestParallelCoversEveryIndexExactlyOnce is the core invariant: the +// chunks partition [0, n), whatever the worker count does to their +// boundaries. +func TestParallelCoversEveryIndexExactlyOnce(t *testing.T) { + for _, n := range []int{0, 1, 2, 7, 33, 100, 1024} { + touched := make([]int, n) + // The chunk check records violations instead of calling + // Fatalf from inside the worker goroutines: FailNow is defined + // for the test's own goroutine only. The append sits under a + // mutex because the chunks run concurrently. + var mu sync.Mutex + illegal := make([][3]int, 0, 4) + Parallel(n, func(start, end int) { + if start < 0 || end > n || start > end { + mu.Lock() + illegal = append(illegal, [3]int{n, start, end}) + mu.Unlock() + } + for i := start; i < end; i++ { + touched[i]++ + } + }) + if len(illegal) > 0 { + t.Fatalf("illegal chunks: %v", illegal) + } + for i, c := range touched { + if c != 1 && n > 0 { + t.Fatalf("n=%d: index %d visited %d times", n, i, c) + } + } + } +} + +// TestParallelSmallWorkloadRunsInline pins the small-workload rule: +// with a single effective worker the callback runs on the caller's +// goroutine before Parallel returns: no goroutine churn for tiny +// kernels. A panicked chunk therefore crashes this test instead of +// hiding behind the WaitGroup. +func TestParallelSmallWorkloadRunsInline(t *testing.T) { + prev := SetNumWorkers(1) + defer SetNumWorkers(prev) + + called := false + Parallel(4, func(start, end int) { + called = true + if start != 0 || end != 4 { + t.Fatalf("single-worker chunk [%d, %d), want [0, 4)", start, end) + } + }) + if !called { + t.Fatal("callback never ran") + } +} + +// TestWorkersForBounds checks both ceilings and the floor. +func TestWorkersForBounds(t *testing.T) { + prev := SetNumWorkers(8) + defer SetNumWorkers(prev) + + for _, tc := range []struct{ n, want int }{ + {0, 1}, {1, 1}, {3, 3}, {8, 8}, {500, 8}, + } { + if got := WorkersFor(tc.n); got != tc.want { + t.Errorf("WorkersFor(%d) = %d, want %d", tc.n, got, tc.want) + } + } + + if got := SetNumWorkers(0); got != 8 { + t.Errorf("SetNumWorkers(0) reported previous %d, want 8", got) + } + if NumWorkers() < 1 { + t.Error("reset landed on an unusable worker count") + } +} + +// TestFloat64PoolRoundTripKeepsLengthAndCapacity pins the borrow +// contract: every GetFloat64Buf returns exactly the requested length +// with capacity for at least that many, buffers survive a Put/Get round +// trip as usable memory, and a buffer handed back dirty arrives cleared. +// The pool enforces the zero-on-borrow guarantee itself, so no kernel +// can leak an earlier borrower's sums into its result (the regression +// TestMatMul2DFloat32ScratchCleaned pins end-to-end in the tensor +// package). +func TestFloat64PoolRoundTripKeepsLengthAndCapacity(t *testing.T) { + buf := GetFloat64Buf(64) + if len(buf) != 64 { + t.Fatalf("borrowed len %d, want 64", len(buf)) + } + if cap(buf) < 64 { + t.Fatalf("borrowed cap %d, want at least 64", cap(buf)) + } + for i := range buf { + buf[i] = float64(i) // fill the whole window: prove it is writable + } + PutFloat64Buf(buf) + + again := GetFloat64Buf(32) + if len(again) != 32 { + t.Fatalf("re-borrowed len %d, want 32", len(again)) + } + if cap(again) < 32 { + t.Fatalf("re-borrowed capacity %d, want at least 32", cap(again)) + } + // The dirty residue the test just put back must never surface: the + // pool clears on borrow, so the window arrives all zeros. (sync.Pool + // may also drop the buffer at any GC, in which case a fresh (and + // therefore zeroed) allocation takes its place; the guarantee holds + // on both paths.) + for i := range again { + if again[i] != 0 { + t.Fatalf("slot %d = %v on arrival, want 0", i, again[i]) + } + again[i] = float64(i) + if again[i] != float64(i) { + t.Fatalf("slot %d = %v after write, want %v", i, again[i], float64(i)) + } + } + PutFloat64Buf(again) +} + +// TestFloat64PoolGrowsForLargerBorrow checks the grow path returns a +// slice of exactly the requested length even when the pooled buffer +// must be reallocated. +func TestFloat64PoolGrowsForLargerBorrow(t *testing.T) { + small := GetFloat64Buf(4) + PutFloat64Buf(small) + + big := GetFloat64Buf(4096) + if len(big) != 4096 { + t.Fatalf("grown len %d, want 4096", len(big)) + } + big[4095] = 1 // writable end to end + PutFloat64Buf(big) +} + +// TestFloat64PoolRetentionCap pins the size rule: the cap itself +// round-trips, anything above it is dropped rather than retained per +// processor, and the drop path leaves the pool usable. +func TestFloat64PoolRetentionCap(t *testing.T) { + if !keepPooled(maxPooledFloat64) { + t.Fatalf("a buffer of the cap (%d) must be retained", maxPooledFloat64) + } + if keepPooled(maxPooledFloat64 + 1) { + t.Fatalf("a buffer above the cap (%d) must be dropped", maxPooledFloat64+1) + } + + oversized := GetFloat64Buf(4 * maxPooledFloat64) + oversized[len(oversized)-1] = 1 + PutFloat64Buf(oversized) // dropped: must not corrupt the pool + + next := GetFloat64Buf(16) + if len(next) != 16 { + t.Fatalf("borrowed after a dropped buffer: len %d, want 16", len(next)) + } + for i := range next { + if next[i] != 0 { + t.Fatalf("slot %d = %v after a dropped buffer, want 0", i, next[i]) + } + } + PutFloat64Buf(next) +} diff --git a/io/bench_test.go b/io/bench_test.go new file mode 100644 index 0000000..0a3d1a1 --- /dev/null +++ b/io/bench_test.go @@ -0,0 +1,353 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +// Benchmarks for the read and write paths that carry the per-value +// work: a table decode, a CSV parse, the HDF5 and NetCDF writers and +// readers. Every input is deterministic and written once, before the +// measured loop. + +import ( + "encoding/binary" + "math" + "os" + "path/filepath" + "strconv" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Sizes: several thousand rows for the tables, and a few hundred +// thousand values for the array formats, which is large enough for the +// per-cell work to dominate the fixed cost of each entry point. +const ( + benchRows = 4000 + benchCols = 24 + benchSide = 512 +) + +// benchFloats builds a deterministic float64 array. +func benchFloats(shape ...int) *core.Array { + a := core.New(core.Float, shape...) + raw := a.RawFloats() + for i := range raw { + raw[i] = math.Sin(float64(i)*0.03125)*1000 + float64(i%97)*0.5 + } + return a +} + +// benchInts builds a deterministic int array. +func benchInts(shape ...int) *core.Array { + a := core.New(core.Int, shape...) + raw := a.RawInts() + for i := range raw { + raw[i] = int64(i)*7 - 3 + } + return a +} + +// benchBinaryTableFile writes a BINTABLE holding one column of every +// numeric form the reader decodes plus a character column, and returns +// its path. The package's own writer emits a subset of those forms +// (B, I and J arrive from other writers), so the table is laid out +// here. +func benchBinaryTableFile(b testing.TB, rows int) string { + b.Helper() + forms := []string{"K", "D", "E", "J", "I", "B", "L", "8A"} + widths := []int{8, 8, 4, 4, 2, 1, 1, 8} + rowBytes := 0 + for _, w := range widths { + rowBytes += w + } + cards := []string{ + fitsStringCardRaw("XTENSION", "BINTABLE"), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 2), + fitsIntCard("NAXIS1", rowBytes), + fitsIntCard("NAXIS2", rows), + fitsIntCard("PCOUNT", 0), + fitsIntCard("GCOUNT", 1), + fitsIntCard("TFIELDS", len(forms)), + } + for i, form := range forms { + n := strconv.Itoa(i + 1) + cards = append(cards, + fitsStringCardRaw("TTYPE"+n, "COL"+n), + fitsStringCardRaw("TFORM"+n, form)) + } + cards = append(cards, fitsEndCard()) + out := fitsAppendCards(nil, []string{ + fitsBoolCard("SIMPLE", true), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 0), + fitsBoolCard("EXTEND", true), + fitsEndCard(), + }) + out = fitsAppendCards(out, cards) + body := make([]byte, rows*rowBytes) + star := []byte("star ") + for r := range rows { + p := r * rowBytes + binary.BigEndian.PutUint64(body[p:], uint64(1000+r)) + p += 8 + binary.BigEndian.PutUint64(body[p:], math.Float64bits(float64(r)*0.25-1)) + p += 8 + binary.BigEndian.PutUint32(body[p:], math.Float32bits(float32(r)*0.5)) + p += 4 + binary.BigEndian.PutUint32(body[p:], uint32(int32(-r))) + p += 4 + binary.BigEndian.PutUint16(body[p:], uint16(int16(r))) + p += 2 + body[p] = byte(r) + p++ + if r%2 == 0 { + body[p] = 'T' + } else { + body[p] = 'F' + } + p++ + copy(body[p:], star) + } + out = append(out, body...) + out = fitsAppendZeroPad(out) + path := filepath.Join(b.TempDir(), "binary.fits") + if err := os.WriteFile(path, out, 0o644); err != nil { + b.Fatal(err) + } + return path +} + +// benchASCIITableFile writes an ASCII table with a character, an +// integer and a float column, and returns its path. +func benchASCIITableFile(b testing.TB, rows int) string { + b.Helper() + text := make([]string, rows) + ints := core.New(core.Int, rows) + floats := core.New(core.Float, rows) + for i := range rows { + text[i] = "star-" + strconv.Itoa(i) + ints.RawInts()[i] = int64(1000 + i) + floats.RawFloats()[i] = float64(i)*0.125 - 42 + } + cols := []FITSTableColumn{ + {Name: "STAR", Form: "12A", Text: text}, + {Name: "ID", Form: "I10", Data: ints}, + {Name: "MAG", Form: "D20.12", Data: floats}, + } + path := filepath.Join(b.TempDir(), "ascii.fits") + if err := SaveFITSTable(path, true, cols, nil); err != nil { + b.Fatal(err) + } + return path +} + +// benchCSVFile writes a float64 matrix as CSV and returns its path. +func benchCSVFile(b testing.TB, rows, cols int) string { + b.Helper() + path := filepath.Join(b.TempDir(), "bench.csv") + if err := SaveCSV(path, benchFloats(rows, cols)); err != nil { + b.Fatal(err) + } + return path +} + +// benchHDF5File writes a classic HDF5 file of one float64 and one +// int64 dataset, optionally through the deflate and shuffle filters, +// and returns its path. +func benchHDF5File(b testing.TB, filtered bool) string { + b.Helper() + sets := []HDF5Dataset{ + {Path: "/field", Values: benchFloats(benchSide, benchSide)}, + {Path: "/ids", Values: benchInts(benchSide, benchSide)}, + } + attrs := map[string]map[string]string{"/": {"origin": "benchmark"}} + opts := HDF5WriteOptions{} + if filtered { + opts = HDF5WriteOptions{Gzip: 6, Shuffle: true} + } + path := filepath.Join(b.TempDir(), "bench.h5") + if err := SaveHDF5(path, sets, attrs, opts); err != nil { + b.Fatal(err) + } + return path +} + +// benchNetCDFFile writes a classic NetCDF file of one float64 variable +// in two dimensions and returns its path. +func benchNetCDFFile(b testing.TB) string { + b.Helper() + dims := []NetCDFDim{{Name: "row", Length: benchSide}, {Name: "col", Length: benchSide}} + vars := []NetCDFVar{{Name: "field", Dims: []string{"row", "col"}, Values: benchFloats(benchSide, benchSide)}} + path := filepath.Join(b.TempDir(), "bench.nc") + if err := SaveNetCDF(path, dims, vars, map[string]string{"title": "benchmark"}); err != nil { + b.Fatal(err) + } + return path +} + +func BenchmarkLoadFITSTableBinary(b *testing.B) { + path := benchBinaryTableFile(b, benchRows) + b.ReportAllocs() + for b.Loop() { + table, err := LoadFITSTable(path) + if err != nil { + b.Fatal(err) + } + if table.Rows != benchRows || len(table.Columns) != 8 || table.Text[7] == nil { + b.Fatalf("table came back as %s with %d columns", table.Kind, len(table.Columns)) + } + } +} + +func BenchmarkLoadFITSTableASCII(b *testing.B) { + path := benchASCIITableFile(b, benchRows) + b.ReportAllocs() + for b.Loop() { + table, err := LoadFITSTable(path) + if err != nil { + b.Fatal(err) + } + if table.Rows != benchRows || len(table.Columns) != 3 { + b.Fatalf("table came back with %d rows and %d columns", table.Rows, len(table.Columns)) + } + } +} + +func BenchmarkLoadCSV(b *testing.B) { + path := benchCSVFile(b, benchRows, benchCols) + b.ReportAllocs() + for b.Loop() { + a, err := LoadCSV(path, false) + if err != nil { + b.Fatal(err) + } + if a.Shape()[0] != benchRows || a.Shape()[1] != benchCols { + b.Fatalf("array came back with shape %v", a.Shape()) + } + } +} + +func BenchmarkSaveHDF5(b *testing.B) { + sets := []HDF5Dataset{ + {Path: "/field", Values: benchFloats(benchSide, benchSide)}, + {Path: "/ids", Values: benchInts(benchSide, benchSide)}, + } + attrs := map[string]map[string]string{"/": {"origin": "benchmark"}} + path := filepath.Join(b.TempDir(), "bench.h5") + b.ReportAllocs() + b.SetBytes(int64(2 * benchSide * benchSide * 8)) + for b.Loop() { + if err := SaveHDF5(path, sets, attrs); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSaveHDF5Large writes one 8 MiB float64 field, the size +// where moving the payload around shows above the per-value work. +func BenchmarkSaveHDF5Large(b *testing.B) { + sets := []HDF5Dataset{{Path: "/field", Values: benchFloats(1024, 1024)}} + path := filepath.Join(b.TempDir(), "bench.h5") + b.ReportAllocs() + b.SetBytes(8 << 20) + for b.Loop() { + if err := SaveHDF5(path, sets, nil); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSaveHDF5Filtered writes the two benchmark datasets through +// the shuffle and deflate filters, the chunked path. +func BenchmarkSaveHDF5Filtered(b *testing.B) { + sets := []HDF5Dataset{ + {Path: "/field", Values: benchFloats(benchSide, benchSide)}, + {Path: "/ids", Values: benchInts(benchSide, benchSide)}, + } + path := filepath.Join(b.TempDir(), "bench.h5") + b.ReportAllocs() + b.SetBytes(int64(2 * benchSide * benchSide * 8)) + for b.Loop() { + if err := SaveHDF5(path, sets, nil, HDF5WriteOptions{Gzip: 6, Shuffle: true}); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSaveHDF5Text writes one string dataset, the fixed-width +// text path whose elements are padded to the longest of the file. +func BenchmarkSaveHDF5Text(b *testing.B) { + const side = 512 + text := make([]string, side*side) + width := 0 + for i := range text { + text[i] = "row" + strconv.Itoa(i) + width = max(width, len(text[i])) + } + sets := []HDF5TextDataset{{Path: "/labels", Shape: []int{side, side}, Text: text}} + path := filepath.Join(b.TempDir(), "bench.h5") + b.ReportAllocs() + b.SetBytes(int64(side * side * width)) + for b.Loop() { + if err := SaveHDF5Text(path, sets); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLoadHDF5(b *testing.B) { + path := benchHDF5File(b, false) + b.ReportAllocs() + for b.Loop() { + sets, err := LoadHDF5(path) + if err != nil { + b.Fatal(err) + } + if len(sets) != 2 { + b.Fatalf("file came back with %d datasets", len(sets)) + } + } +} + +func BenchmarkLoadHDF5Filtered(b *testing.B) { + path := benchHDF5File(b, true) + b.ReportAllocs() + for b.Loop() { + sets, err := LoadHDF5(path) + if err != nil { + b.Fatal(err) + } + if len(sets) != 2 { + b.Fatalf("file came back with %d datasets", len(sets)) + } + } +} + +func BenchmarkSaveNetCDF(b *testing.B) { + dims := []NetCDFDim{{Name: "row", Length: benchSide}, {Name: "col", Length: benchSide}} + vars := []NetCDFVar{{Name: "field", Dims: []string{"row", "col"}, Values: benchFloats(benchSide, benchSide)}} + attrs := map[string]string{"title": "benchmark"} + path := filepath.Join(b.TempDir(), "bench.nc") + b.ReportAllocs() + for b.Loop() { + if err := SaveNetCDF(path, dims, vars, attrs); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLoadNetCDF(b *testing.B) { + path := benchNetCDFFile(b) + b.ReportAllocs() + for b.Loop() { + _, vars, _, err := LoadNetCDF(path) + if err != nil { + b.Fatal(err) + } + if len(vars) != 1 { + b.Fatalf("file came back with %d variables", len(vars)) + } + } +} diff --git a/io/csv.go b/io/csv.go new file mode 100644 index 0000000..0c6fc26 --- /dev/null +++ b/io/csv.go @@ -0,0 +1,471 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "bufio" + "bytes" + "encoding/csv" + "io" + "os" + "strconv" + "unsafe" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// CSV IO. SaveCSV writes a 2-D array as comma-separated text; +// LoadCSV reads it back. CSV is the interchange format for tabular +// data with spreadsheets and statistical packages, and the plain-text +// counterpart to the binary FITS format for data pipelines. Only 2-D +// arrays are supported (rows by columns), matching the tabular model. +// The writer formats every dtype the core carries: the integer class +// as exact integer text, booleans as 0 and 1, floats as float text; +// the reader parses everything back as float64. + +// SaveCSV writes a 2-D array to path as comma-separated values. The +// result is loadable by any spreadsheet or data tool. The close error +// is part of the write: os.Create's descriptor buffers nothing itself, +// but the operating system still reports a failed write through Close +// on some filesystems, so a discarded Close error would report a save +// that never reached the disk. +func SaveCSV(path string, a *core.Array) error { + if a.NDim() != 2 { + return base.Errf("SaveCSV: needs a 2-D array, got shape %s", base.ShapeText(a.Shape())) + } + f, err := os.Create(path) + if err != nil { + return base.Errf("SaveCSV: %w", err) + } + werr := SaveCSVWriter(f, a) + if cerr := f.Close(); werr != nil { + return base.Errf("SaveCSV: %w", werr) + } else if cerr != nil { + return base.Errf("SaveCSV: %w", cerr) + } + return nil +} + +// SaveCSVWriter writes a 2-D array as CSV to w. Every dtype the core +// carries has a text form: Bool as 0 and 1, the integer dtypes as +// exact decimal widenings of their payload, Float16 through its exact +// float64 widening, and Float32 and Float as 'g' float text. Complex +// arrays are refused: CSV carries plain numeric text, and a pair of +// raw halves would read back as two unrelated columns. An array with +// no columns is refused for the same reason: one empty record per row +// is what the writer would emit, and CSV readers treat blank lines as +// no records at all, so the shape would come back as (0,0). +func SaveCSVWriter(w io.Writer, a *core.Array) error { + if a.NDim() != 2 { + return base.Errf("SaveCSV: needs a 2-D array, got shape %s", base.ShapeText(a.Shape())) + } + if a.Dtype() == core.Complex { + return base.Errf("SaveCSV: complex arrays are not supported") + } + cw := csv.NewWriter(w) + rows, cols := a.Shape()[0], a.Shape()[1] + if cols == 0 { + // No column count can follow from CSV text that has no record. + return base.Errf("SaveCSV: an array with no columns has no CSV form, got shape %s", base.ShapeText(a.Shape())) + } + // The package's arrays, views included, always keep payload[i] at + // element i (views rebase the payload, never stride it), so the + // row walk indexes the payload directly instead of going through + // the per-element accessor. + rec := make([]string, cols) + for r := range rows { + baseIdx := r * cols + switch a.Dtype() { + case core.Float: + for c := range cols { + rec[c] = strconv.FormatFloat(a.RawFloats()[baseIdx+c], 'g', -1, 64) + } + case core.Float32: + for c := range cols { + rec[c] = strconv.FormatFloat(float64(a.RawFloat32s()[baseIdx+c]), 'g', -1, 64) + } + case core.Float16: + // The half payload widens to float64 exactly, and the + // widened value formats the way every float path formats. + vals := a.RawHalves() + for c := range cols { + rec[c] = strconv.FormatFloat(core.HalfToFloat64(vals[baseIdx+c]), 'g', -1, 64) + } + case core.Bool: + // Bool writes as 0 and 1: the numeric column semantics a + // CSV column carries, the two values the payload stores. + vals := a.RawBools() + for c := range cols { + if vals[baseIdx+c] { + rec[c] = "1" + } else { + rec[c] = "0" + } + } + case core.Int8: + vals := a.RawInt8s() + for c := range cols { + rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) + } + case core.Uint8: + vals := a.RawUint8s() + for c := range cols { + rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) + } + case core.Int16: + vals := a.RawInt16s() + for c := range cols { + rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) + } + case core.Uint16: + vals := a.RawUint16s() + for c := range cols { + rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) + } + case core.Int32: + vals := a.RawInt32s() + for c := range cols { + rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) + } + case core.Uint32: + vals := a.RawUint32s() + for c := range cols { + rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) + } + case core.Int: + // The %d contract: an int64 past 2^53 keeps every digit, + // where a float64 detour would round it away. + vals := a.RawInts() + for c := range cols { + rec[c] = strconv.FormatInt(vals[baseIdx+c], 10) + } + default: + // Unreachable: Complex is refused above and every other + // dtype has its case; the writer answers an error rather + // than panic on a payload it cannot name. + return base.Errf("SaveCSV: dtype %s has no CSV form", a.Dtype()) + } + if err := cw.Write(rec); err != nil { + return base.Errf("SaveCSV: %w", err) + } + } + cw.Flush() + if err := cw.Error(); err != nil { + return base.Errf("SaveCSV: %w", err) + } + return nil +} + +// LoadCSV reads a 2-D float array from a CSV file. Every row must have +// the same number of fields; a header row is treated as data unless +// skipHeader is true. +func LoadCSV(path string, skipHeader bool) (*core.Array, error) { + f, err := os.Open(path) + if err != nil { + return nil, base.Errf("LoadCSV: %w", err) + } + defer f.Close() + return LoadCSVReader(f, skipHeader) +} + +// LoadCSVReader reads a 2-D float array from a CSV stream. A leading +// UTF-8 byte order mark is skipped: spreadsheets write one, and the +// parser would otherwise glue it onto the first field. +// +// The records stream: each one is parsed and its values are copied out +// as it arrives, so no row of text survives the row that produced it +// and the whole file is never held as strings. A malformed record +// stops the read and is reported first, then the first row whose field +// count differs from the first row's, then the first value that is not +// a number, the order a whole-file parse reports them in. +// +// The record parser is the hand-rolled tokenizer below, tuned for the +// numeric tables this entry point serves: it keeps the record +// semantics of encoding/csv for this caller's configuration, quotes +// and error reports included, while skipping the per-record string the +// standard parser builds. +func LoadCSVReader(r io.Reader, skipHeader bool) (*core.Array, error) { + br := bufio.NewReaderSize(r, csvReadBuffer) + prefix, _ := br.Peek(3) + if len(prefix) == 3 && prefix[0] == 0xEF && prefix[1] == 0xBB && prefix[2] == 0xBF { + _, _, _ = br.ReadRune() // consume the mark + } + t := &csvTokenizer{br: br} + vals := make([]float64, 0) + rows, cols := 0, 0 + header := skipHeader + var ragged, badValue error + for { + rec, err := t.nextRecord() + if err == io.EOF { + break + } + if err != nil { + return nil, base.Errf("LoadCSV: %w", err) + } + if header { + header = false + continue + } + if rows == 0 { + cols = len(rec) + } else if len(rec) != cols { + // The read carries on to the end even after a defect, so a + // later record's syntax error still outranks an earlier + // ragged row, as the whole-file parse decides it. The row + // counter keeps advancing for the same reason. + if ragged == nil { + ragged = base.Errf("LoadCSV: row %d has %d fields, want %d", rows, len(rec), cols) + } + rows++ + continue + } + if badValue == nil { + // The buffer doubles as the file arrives: append alone grows a + // large float slice by a quarter and rewrites it several times + // over, while doubling copies less than the final size once, + // and the first record sizes the buffer exactly. + if need := len(vals) + len(rec); need > cap(vals) { + grown := make([]float64, len(vals), max(2*cap(vals), need)) + copy(grown, vals) + vals = grown + } + for j, f := range rec { + v, perr := strconv.ParseFloat(csvFieldText(f), 64) + if perr != nil { + if ne, ok := perr.(*strconv.NumError); ok { + // The message carries the field text: rebuild + // it over a private copy so it does not hang + // off the row buffer this read goes on + // overwriting. + perr = &strconv.NumError{Func: ne.Func, Num: string(f), Err: ne.Err} + } + badValue = base.Errf("LoadCSV: row %d col %d: %w", rows, j, perr) + break + } + vals = append(vals, v) + } + } + rows++ + } + if ragged != nil { + return nil, ragged + } + if badValue != nil { + return nil, badValue + } + return core.FromFloats(vals, rows, cols) +} + +// csvReadBuffer sizes the read buffer. A numeric table's rows run to a +// few hundred bytes, so one refill covers thousands of them and the +// reads stop being part of the cost. +const csvReadBuffer = 1 << 16 + +// csvFieldText views a field's bytes as the string strconv.ParseFloat +// parses, with no copy. +// +// SAFETY: the bytes live in the tokenizer's line or record buffer and +// stay untouched until the next record is read: ParseFloat reads the +// string only within the call, and on failure the NumError is rebuilt +// over a private copy before the loop can move on to the next record. +// On success no reference to the string escapes. +func csvFieldText(f []byte) string { + return unsafe.String(unsafe.SliceData(f), len(f)) +} + +// csvTokenizer is the record reader LoadCSVReader runs. It reproduces +// what encoding/csv's Reader does for the configuration this package +// reads with (comma separator, no comment character, no lazy quotes, +// no leading-space trimming, a variable field count): blank lines are +// skipped, a quote opens a quoted field, "" escapes one quote, \r\n is +// normalised to \n everywhere, interior newlines of a quoted field +// included, the last record may lack its newline, a trailing \r is +// dropped before EOF, and a malformed record comes back as a +// *csv.ParseError naming the same line and column the standard parser +// names. What it does not do is build the per-record string: the +// fields are byte ranges over the tokenizer's own buffers, valid until +// the next call. +type csvTokenizer struct { + br *bufio.Reader + + numLine int // the line the reader sits on, counted from one + + lineBuf []byte // assembles a record longer than the read buffer + recordBuf []byte // the record's unescaped fields, back to back + fields [][]byte // the current record's fields, reused +} + +// lengthNL reports the number of bytes for the trailing \n. +func lengthNL(b []byte) int { + if len(b) > 0 && b[len(b)-1] == '\n' { + return 1 + } + return 0 +} + +// readLine returns the next line with its end mark. A trailing \r\n is +// normalised to \n in place, a trailing \r is dropped before EOF, and +// every line read counts toward the line numbers the parse errors +// carry. The result is only valid until the next call. +func (t *csvTokenizer) readLine() ([]byte, error) { + line, err := t.br.ReadSlice('\n') + if err == bufio.ErrBufferFull { + t.lineBuf = append(t.lineBuf[:0], line...) + for err == bufio.ErrBufferFull { + line, err = t.br.ReadSlice('\n') + t.lineBuf = append(t.lineBuf, line...) + } + line = t.lineBuf + } + readSize := len(line) + if readSize > 0 && err == io.EOF { + err = nil + // For backwards compatibility, drop a trailing \r before EOF. + if line[readSize-1] == '\r' { + line = line[:readSize-1] + } + } + t.numLine++ + // Normalise \r\n to \n on every input line. + if n := len(line); n >= 2 && line[n-2] == '\r' && line[n-1] == '\n' { + line[n-2] = '\n' + line = line[:n-1] + } + return line, err +} + +// nextRecord parses the next record. The returned slices share the +// tokenizer's buffers and are only valid until the next call. +func (t *csvTokenizer) nextRecord() ([][]byte, error) { + // Read the record's first line, skipping the blank ones. + var line []byte + var errRead error + for errRead == nil { + line, errRead = t.readLine() + if errRead == nil && len(line) == lengthNL(line) { + continue // a blank line carries no record + } + break + } + if errRead == io.EOF { + return nil, errRead + } + + // Fast path: a record with no quote anywhere splits on the commas + // of its one line. A quote is the only construct that carries a + // record across lines, escapes a field or raises a parse error, + // so the quoted path below owns every case beyond the split. + if bytes.IndexByte(line, '"') < 0 { + t.fields = t.fields[:0] + for { + i := bytes.IndexByte(line, ',') + if i < 0 { + t.fields = append(t.fields, line[:len(line)-lengthNL(line)]) + return t.fields, errRead + } + t.fields = append(t.fields, line[:i]) + line = line[i+1:] + } + } + + recLine := t.numLine // the line the record starts on + t.recordBuf = t.recordBuf[:0] + t.fields = t.fields[:0] + posLine, posCol := t.numLine, 1 + var parseErr error +parseField: + for { + if len(line) == 0 || line[0] != '"' { + // Non-quoted field: everything up to the comma, with + // the end mark stripped. + i := bytes.IndexByte(line, ',') + field := line + if i >= 0 { + field = field[:i] + } else { + field = field[:len(field)-lengthNL(field)] + } + // A quote may not appear in a non-quoted field. + if j := bytes.IndexByte(field, '"'); j >= 0 { + parseErr = &csv.ParseError{StartLine: recLine, Line: t.numLine, Column: posCol + j, Err: csv.ErrBareQuote} + break parseField + } + start := len(t.recordBuf) + t.recordBuf = append(t.recordBuf, field...) + t.fields = append(t.fields, t.recordBuf[start:len(t.recordBuf)]) + if i >= 0 { + line = line[i+1:] + posCol += i + 1 + continue parseField + } + break parseField + } + // Quoted field: the opening quote is consumed, the rest + // accumulates until the closing one. + fieldStart := len(t.recordBuf) + line = line[1:] + posCol++ + for { + i := bytes.IndexByte(line, '"') + if i >= 0 { + t.recordBuf = append(t.recordBuf, line[:i]...) + line = line[i+1:] + posCol += i + 1 + switch { + case len(line) > 0 && line[0] == '"': + // "" escapes one quote. + t.recordBuf = append(t.recordBuf, '"') + line = line[1:] + posCol++ + case len(line) > 0 && line[0] == ',': + // ", closes the field. + line = line[1:] + posCol++ + t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)]) + continue parseField + case len(line) == 0 || (len(line) == 1 && line[0] == '\n'): + // A closing quote at the end of the line closes + // the field and the record with it. + t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)]) + break parseField + default: + // Anything after a closing quote that is neither + // comma nor end of line. + parseErr = &csv.ParseError{StartLine: recLine, Line: t.numLine, Column: posCol - 1, Err: csv.ErrQuote} + break parseField + } + } else if len(line) > 0 { + // End of line inside the field: the whole line, end + // mark included, belongs to the field. + t.recordBuf = append(t.recordBuf, line...) + if errRead != nil { + break parseField + } + posCol += len(line) + line, errRead = t.readLine() + if len(line) > 0 { + posLine++ + posCol = 1 + } + if errRead == io.EOF { + errRead = nil + } + } else { + // End of input inside the field. + if errRead == nil { + parseErr = &csv.ParseError{StartLine: recLine, Line: posLine, Column: posCol, Err: csv.ErrQuote} + break parseField + } + t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)]) + break parseField + } + } + } + if parseErr == nil { + parseErr = errRead + } + return t.fields, parseErr +} diff --git a/io/csv_bench_test.go b/io/csv_bench_test.go new file mode 100644 index 0000000..c4fe4a0 --- /dev/null +++ b/io/csv_bench_test.go @@ -0,0 +1,65 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +// The A/B benchmark for the CSV read path: the legacy encoding/csv +// reader and the hand-rolled tokenizer side by side, in one process. +// The two sub-benchmarks alternate within every -count round, so both +// see the same machine, and the decision reads the medians with 1 to 2 +// percent treated as noise. + +import ( + "bytes" + "io" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The A/B table: a hundred thousand rows of sixteen float64 values, in +// the text shape SaveCSV emits, sized so the per-value work dominates +// everything else. +const ( + benchCSVABRows = 100000 + benchCSVABCols = 16 +) + +// benchCSVABData renders the A/B table as CSV bytes, once, before the +// measured loops. +func benchCSVABData(b *testing.B) []byte { + b.Helper() + var buf bytes.Buffer + if err := SaveCSVWriter(&buf, benchFloats(benchCSVABRows, benchCSVABCols)); err != nil { + b.Fatal(err) + } + return buf.Bytes() +} + +// BenchmarkLoadCSVReader races the two readers over the same table. +func BenchmarkLoadCSVReader(b *testing.B) { + data := benchCSVABData(b) + impls := []struct { + name string + load func(io.Reader, bool) (*core.Array, error) + }{ + {"old", loadCSVReaderLegacy}, + {"new", LoadCSVReader}, + } + for _, impl := range impls { + b.Run(impl.name, func(b *testing.B) { + r := bytes.NewReader(data) + b.ReportAllocs() + for b.Loop() { + r.Reset(data) + a, err := impl.load(r, false) + if err != nil { + b.Fatal(err) + } + if a.Shape()[0] != benchCSVABRows || a.Shape()[1] != benchCSVABCols { + b.Fatalf("array came back with shape %v", a.Shape()) + } + } + }) + } +} diff --git a/io/csv_parity_test.go b/io/csv_parity_test.go new file mode 100644 index 0000000..1e25a16 --- /dev/null +++ b/io/csv_parity_test.go @@ -0,0 +1,242 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +// Parity between the hand-rolled CSV tokenizer and the encoding/csv +// implementation it replaced. The legacy reader below is the oracle: +// the differential test walks a corpus of inputs through both and +// demands the same error text, the same shape and the same float bits, +// and the fuzz target hunts for an input where they part ways. + +import ( + "bufio" + "encoding/csv" + "io" + "math" + "slices" + "strconv" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// loadCSVReaderLegacy is the encoding/csv reader the tokenizer +// replaced, kept verbatim, byte order mark, buffer growth included. +func loadCSVReaderLegacy(r io.Reader, skipHeader bool) (*core.Array, error) { + br := bufio.NewReader(r) + prefix, _ := br.Peek(3) + if len(prefix) == 3 && prefix[0] == 0xEF && prefix[1] == 0xBB && prefix[2] == 0xBF { + _, _, _ = br.ReadRune() // consume the mark + } + cr := csv.NewReader(br) + cr.FieldsPerRecord = -1 // allow variable; validated manually + // The record is consumed before the next Read, so the reader may + // reuse its slice: one row buffer serves the whole file. + cr.ReuseRecord = true + vals := make([]float64, 0) + rows, cols := 0, 0 + header := skipHeader + var ragged, badValue error + for { + rec, err := cr.Read() + if err == io.EOF { + break + } + if err != nil { + return nil, base.Errf("LoadCSV: %w", err) + } + if header { + header = false + continue + } + if rows == 0 { + cols = len(rec) + } else if len(rec) != cols { + if ragged == nil { + ragged = base.Errf("LoadCSV: row %d has %d fields, want %d", rows, len(rec), cols) + } + rows++ + continue + } + if badValue == nil { + if need := len(vals) + len(rec); need > cap(vals) { + grown := make([]float64, len(vals), max(2*cap(vals), need)) + copy(grown, vals) + vals = grown + } + for j, f := range rec { + v, perr := strconv.ParseFloat(f, 64) + if perr != nil { + badValue = base.Errf("LoadCSV: row %d col %d: %w", rows, j, perr) + break + } + vals = append(vals, v) + } + } + rows++ + } + if ragged != nil { + return nil, ragged + } + if badValue != nil { + return nil, badValue + } + return core.FromFloats(vals, rows, cols) +} + +// csvParityCases collects the inputs whose treatment the tokenizer has +// to match byte for byte: the shapes a numeric table takes, the quote +// grammar, the line-ending conventions and every defect the reader +// reports. +var csvParityCases = []struct { + name string + in string + header bool +}{ + {"plain table", "1.5,2.25\n3.125,4\n", false}, + {"last row without a newline", "1,2\n3,4", false}, + {"single row without a newline", "1,2", false}, + {"crlf", "1,2\r\n3,4\r\n", false}, + {"mixed line endings", "1,2\r\n3,4\n5,6\r\n", false}, + {"blank lines between records", "1,2\n\n3,4\n\n", false}, + {"blank line only", "\n\n\n", false}, + {"empty input", "", false}, + {"empty field", "1,,3\n4,5,6\n", false}, + {"trailing empty field", "1,2,\n3,4,5\n", false}, + {"single comma", ",", false}, + {"spaces stay in the field", " 1 , 2 \n3,4\n", false}, + {"lone carriage return inside a line", "1,2\r3,4\n", false}, + {"carriage return before eof", "1,2\r", false}, + {"carriage return blank line", "1,2\n\r\n3,4\n", false}, + {"bom then table", "\xEF\xBB\xBF1.5,2.5\n3,4\n", false}, + {"quoted values", "1,\"2\"\n\"3\",4\n", false}, + {"quoted header", "\"a\",\"b\"\n1,2\n", true}, + {"quoted comma", "\"1,2\",3\n4,5,6\n", false}, + {"escaped quote", "\"a\"\"b\",2\n", false}, + {"quote run at field end", "\"a\"\"\",2\n", false}, + {"empty quoted fields", "\"\",\"\",\"\"\n", false}, + {"multiline quoted field", "\"1\n2\",3\n4,5,6\n", false}, + {"multiline quoted field with crlf", "\"1\r\n2\",3\n4,5,6\n", false}, + {"blank line inside a quoted field", "\"a\n\nb\",1\n", false}, + {"quote only field mid table", "1,2\n\"3\",4\n", false}, + {"header skip", "a,b\n1,2\n", true}, + {"header skip with blank line", "a,b\n\n1,2\n", true}, + {"ragged rows", "1,2\n3\n4,5,6\n", false}, + {"ragged rows with header", "a,b\n1,2\n3\n", true}, + {"two bad values", "1,x\n4,y\n", false}, + {"bad value after ragged row", "1,2\n3\n4,x\n", false}, + {"bare quote", "a\"b,1\n", false}, + {"bare quote in a later field", "1,b\"c,2\n", false}, + {"text after a closing quote", "1,\"a\"b,2\n", false}, + {"text after a closing quote on a later line", "1\n\"a\"b,2\n", false}, + {"unterminated quote at eof", "1,2\n3,\"ab", false}, + {"unterminated quote at end of line", "1,2\n3,\"ab\n", false}, + {"unterminated quote after a multiline field", "1,\"a\nb", false}, + {"unterminated quote in the header", "\"a,b\n1,2\n", true}, + {"only a header", "a,b\n", true}, + {"only a header without skip", "a,b\n", false}, + {"range overflow", "1e309,2\n", false}, + {"range underflow", "1,1e-400\n", false}, + {"hex float", "0x1p-2,2\n", false}, + {"infinity spelling", "Inf,+Inf,-inf\n", false}, + {"nan spelling", "NaN,-nan\n", false}, + {"utf-8 in a field", "λ,2\n", false}, + {"wide utf-8 in a field", "1,𝄞\n", false}, +} + +// csvParityExtra builds the generated cases the fixed list cannot +// spell: records longer than the read buffer, quoted or not, and a +// generated table in the shape SaveCSV emits. +func csvParityExtra(t *testing.T) []struct { + name string + in string + header bool +} { + t.Helper() + var sb strings.Builder + for j := range 40000 { + if j > 0 { + sb.WriteByte(',') + } + sb.WriteString(strconv.Itoa(j)) + } + longLine := sb.String() + var qb strings.Builder + qb.WriteString("1,\"") + qb.WriteString(strings.Repeat("234567890\n", 10000)) + qb.WriteString("\",2\n") + return []struct { + name string + in string + header bool + }{ + {"record longer than the read buffer", longLine + "\n1,2\n", false}, + {"long record without a newline", longLine, false}, + {"long quoted field across the buffer", qb.String(), false}, + {"ragged long record", longLine + "\n1\n", false}, + {"long last field without a newline", "1," + strings.Repeat("2", 70000), false}, + } +} + +// checkCSVParity runs one input through both readers and demands the +// same outcome: the same error text, or the same shape and the same +// float bits. +func checkCSVParity(t *testing.T, name, in string, skipHeader bool) { + t.Helper() + want, wantErr := loadCSVReaderLegacy(strings.NewReader(in), skipHeader) + got, gotErr := LoadCSVReader(strings.NewReader(in), skipHeader) + switch { + case wantErr != nil && gotErr != nil: + if wantErr.Error() != gotErr.Error() { + t.Fatalf("%s: error %q, want %q", name, gotErr, wantErr) + } + return + case wantErr != nil || gotErr != nil: + t.Fatalf("%s: error mismatch: legacy %v, tokenizer %v", name, wantErr, gotErr) + } + if !slices.Equal(want.Shape(), got.Shape()) { + t.Fatalf("%s: shape %v, want %v", name, got.Shape(), want.Shape()) + } + wv, gv := want.RawFloats(), got.RawFloats() + for i := range wv { + if math.Float64bits(wv[i]) != math.Float64bits(gv[i]) { + t.Fatalf("%s: value %d is %v (%#x), want %v (%#x)", name, i, gv[i], math.Float64bits(gv[i]), wv[i], math.Float64bits(wv[i])) + } + } +} + +// TestLoadCSVParity pins the tokenizer to the encoding/csv reader it +// replaced, one input at a time. +func TestLoadCSVParity(t *testing.T) { + for _, tc := range csvParityCases { + checkCSVParity(t, tc.name, tc.in, tc.header) + checkCSVParity(t, tc.name+" with header flag", tc.in, !tc.header) + } + for _, tc := range csvParityExtra(t) { + checkCSVParity(t, tc.name, tc.in, tc.header) + checkCSVParity(t, tc.name+" with header flag", tc.in, !tc.header) + } + // A generated table, the shape the reader is built for. + var sb strings.Builder + a, _ := core.FromFloats([]float64{1.5, -2.25, 3.125, 4, 5.5, -6.75, 1e-9, -0, 0.1, 2, 3, 4}, 4, 3) + if err := SaveCSVWriter(&sb, a); err != nil { + t.Fatal(err) + } + checkCSVParity(t, "saved table", sb.String(), false) + checkCSVParity(t, "saved table with header", "c0,c1,c2\n"+sb.String(), true) +} + +// FuzzLoadCSVParity hunts for the input where the tokenizer and the +// encoding/csv reader part ways. +func FuzzLoadCSVParity(f *testing.F) { + for _, tc := range csvParityCases { + f.Add(tc.in, tc.header) + } + f.Add(strings.Repeat("1.5,", 300)+"1.5\n", false) + f.Fuzz(func(t *testing.T, in string, skipHeader bool) { + checkCSVParity(t, "fuzz", in, skipHeader) + }) +} diff --git a/io/csv_test.go b/io/csv_test.go new file mode 100644 index 0000000..d98b808 --- /dev/null +++ b/io/csv_test.go @@ -0,0 +1,222 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "path/filepath" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +func TestCSVRoundTrip(t *testing.T) { + a, _ := core.FromFloats([]float64{1.5, -2.25, 3.125, 4.0, 5.5, -6.75}, 2, 3) + dir := t.TempDir() + path := filepath.Join(dir, "data.csv") + if err := SaveCSV(path, a); err != nil { + t.Fatal(err) + } + back, err := LoadCSV(path, false) + if err != nil { + t.Fatal(err) + } + if !core.Equal(a, back) { + t.Errorf("CSV round-trip: got %s, want %s", back, a) + } +} + +func TestCSVWriterReader(t *testing.T) { + a, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + var sb strings.Builder + if err := SaveCSVWriter(&sb, a); err != nil { + t.Fatal(err) + } + if sb.String() != "1\n2\n3\n4\n" && sb.String() != "1,2\n3,4\n" { + // csv.Writer separates with commas by default. + if !strings.Contains(sb.String(), ",") { + t.Fatalf("unexpected CSV: %q", sb.String()) + } + } + back, err := LoadCSVReader(strings.NewReader(sb.String()), false) + if err != nil { + t.Fatal(err) + } + if !core.Equal(a, back) { + t.Errorf("reader round-trip: got %s, want %s", back, a) + } +} + +func TestCSVSkipHeader(t *testing.T) { + data := "a,b,c\n1,2,3\n4,5,6\n" + back, err := LoadCSVReader(strings.NewReader(data), true) + if err != nil { + t.Fatal(err) + } + want, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if !core.Equal(back, want) { + t.Errorf("skip header: got %s, want %s", back, want) + } +} + +func TestCSVErrors(t *testing.T) { + // Non-2-D array. + a, _ := core.FromFloats([]float64{1, 2, 3}, 3) + if err := SaveCSV("/tmp/x.csv", a); err == nil { + t.Error("SaveCSV: expected error for 1-D array") + } + // Ragged rows. + if _, err := LoadCSVReader(strings.NewReader("1,2\n3\n"), false); err == nil { + t.Error("LoadCSV: expected error for ragged rows") + } + // Non-numeric. + if _, err := LoadCSVReader(strings.NewReader("1,x\n"), false); err == nil { + t.Error("LoadCSV: expected error for non-numeric cell") + } +} + +// TestCSVComplexRefused pins the dtype contract: CSV carries plain +// float text, so a complex array is refused with an error instead of +// panicking on its nil float payload. +func TestCSVComplexRefused(t *testing.T) { + a, err := core.FromComplexes([]complex128{complex(1, 1), complex(2, 2)}, 2, 1) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + var sb strings.Builder + if err := SaveCSVWriter(&sb, a); err == nil { + t.Error("SaveCSVWriter: expected an error for a complex array") + } + dir := t.TempDir() + if err := SaveCSV(filepath.Join(dir, "c.csv"), a); err == nil { + t.Error("SaveCSV: expected an error for a complex array") + } +} + +// TestLoadCSVErrorPrecedence pins which defect a malformed file +// reports: a record the csv parser refuses stops the read and is named +// first, then the first row whose field count differs from the first +// row's, then the first cell that is not a number, the order a +// whole-file parse reports them in. The row and cell named are the +// first offenders, never the last. +func TestLoadCSVErrorPrecedence(t *testing.T) { + cases := []struct { + name string + in string + header bool + want string + absent string + }{ + {"a ragged row is named", "1,2\n3\n", false, "row 1 has 1 fields, want 2", ""}, + {"the first ragged row is named", "1,2\n3\n4,5,6\n", false, "row 1 has 1 fields, want 2", "row 2"}, + {"the first bad cell is named", "1,2\n3,x\n4,y\n", false, "row 1 col 1", "row 2"}, + {"a malformed record outranks a ragged row", "1,2\n3\n\"unterminated\n", false, "parse error", "fields, want"}, + {"a ragged row outranks an earlier bad cell", "1,2\nx,2\n3\n", false, "row 2 has 1 fields, want 2", "col"}, + {"the header is not a data row", "a,b\n1,2\n3\n", true, "row 1 has 1 fields, want 2", ""}, + } + for _, tc := range cases { + _, err := LoadCSVReader(strings.NewReader(tc.in), tc.header) + if err == nil { + t.Fatalf("%s: %q parsed without an error", tc.name, tc.in) + } + if !strings.Contains(err.Error(), tc.want) { + t.Fatalf("%s: %v, want the message to carry %q", tc.name, err, tc.want) + } + if tc.absent != "" && strings.Contains(err.Error(), tc.absent) { + t.Fatalf("%s: %v, want no mention of %q", tc.name, err, tc.absent) + } + } +} + +// TestCSVWriterDtypes pins the writer's text form for every dtype the +// core carries: the integer class through exact decimal widenings (an +// int64 past 2^53 keeps every digit, which a float64 detour would +// round away), Bool as 0 and 1, the numeric column semantics CSV +// carries, and Float16 through its exact float64 widening in the same +// 'g' form the float paths use. +func TestCSVWriterDtypes(t *testing.T) { + text := func(t *testing.T, a *core.Array) string { + t.Helper() + var sb strings.Builder + if err := SaveCSVWriter(&sb, a); err != nil { + t.Fatalf("SaveCSVWriter(%s): %v", a.Dtype(), err) + } + return sb.String() + } + bools, err := core.FromBools([]bool{true, false, false, true}, 2, 2) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + if got, want := text(t, bools), "1,0\n0,1\n"; got != want { + t.Fatalf("bool wrote %q, want %q", got, want) + } + i8, err := core.FromInt8s([]int8{-128, 127, 1, -1}, 2, 2) + if err != nil { + t.Fatalf("FromInt8s: %v", err) + } + if got, want := text(t, i8), "-128,127\n1,-1\n"; got != want { + t.Fatalf("int8 wrote %q, want %q", got, want) + } + u8, err := core.FromUint8s([]uint8{0, 255, 10, 42}, 2, 2) + if err != nil { + t.Fatalf("FromUint8s: %v", err) + } + if got, want := text(t, u8), "0,255\n10,42\n"; got != want { + t.Fatalf("uint8 wrote %q, want %q", got, want) + } + i16, err := core.FromInt16s([]int16{-32768, 32767, 0, -7}, 2, 2) + if err != nil { + t.Fatalf("FromInt16s: %v", err) + } + if got, want := text(t, i16), "-32768,32767\n0,-7\n"; got != want { + t.Fatalf("int16 wrote %q, want %q", got, want) + } + u16, err := core.FromUint16s([]uint16{0, 65535, 5, 4096}, 2, 2) + if err != nil { + t.Fatalf("FromUint16s: %v", err) + } + if got, want := text(t, u16), "0,65535\n5,4096\n"; got != want { + t.Fatalf("uint16 wrote %q, want %q", got, want) + } + i32, err := core.FromInt32s([]int32{-2147483648, 2147483647, 0, -1}, 2, 2) + if err != nil { + t.Fatalf("FromInt32s: %v", err) + } + if got, want := text(t, i32), "-2147483648,2147483647\n0,-1\n"; got != want { + t.Fatalf("int32 wrote %q, want %q", got, want) + } + u32, err := core.FromUint32s([]uint32{0, 4294967295, 7, 65536}, 2, 2) + if err != nil { + t.Fatalf("FromUint32s: %v", err) + } + if got, want := text(t, u32), "0,4294967295\n7,65536\n"; got != want { + t.Fatalf("uint32 wrote %q, want %q", got, want) + } + // The %d contract in its element: 2^53+1 cannot survive a float64 + // detour, and the integer path must not take one. + big, err := core.FromInts([]int64{9007199254740993, -9007199254740993, 0, 2}, 2, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + if got, want := text(t, big), "9007199254740993,-9007199254740993\n0,2\n"; got != want { + t.Fatalf("int wrote %q, want %q", got, want) + } + half, err := core.FromFloat16s([]float64{0.5, -2, 1, 4}, 2, 2) + if err != nil { + t.Fatalf("FromFloat16s: %v", err) + } + if got, want := text(t, half), "0.5,-2\n1,4\n"; got != want { + t.Fatalf("float16 wrote %q, want %q", got, want) + } + f32, err := core.FromFloat32s([]float32{1.5, -2.25, 0, 3.25}, 2, 2) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + if got, want := text(t, f32), "1.5,-2.25\n0,3.25\n"; got != want { + t.Fatalf("float32 wrote %q, want %q", got, want) + } + f64 := mustFloats(t, []float64{1.5, -2.25, 0, 4}, 2, 2) + if got, want := text(t, f64), "1.5,-2.25\n0,4\n"; got != want { + t.Fatalf("float wrote %q, want %q", got, want) + } +} diff --git a/io/doc.go b/io/doc.go new file mode 100644 index 0000000..1db04fa --- /dev/null +++ b/io/doc.go @@ -0,0 +1,61 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package io reads and writes the formats scientific data arrives in: +// comma-separated text, FITS images and tables, HDF5, NetCDF classic +// and native-endian memory maps. Every loader returns the library's +// own Array, so a file read is an ordinary value the rest of the +// library operates on, and every writer takes one. +// +// CSV is the interchange format for tabular data with spreadsheets +// and statistical packages: LoadCSV, LoadCSVReader, SaveCSV and +// SaveCSVWriter handle 2-D arrays, with or without a header row. The +// writer formats every numeric dtype the core carries: the integer +// class as exact integer text, booleans as 0 and 1, floats as float +// text; complex is refused, because CSV carries plain numeric text +// and a pair of raw halves would read back as two unrelated columns. +// The readers parse everything back as float64. +// +// FITS is astronomy's archival format. LoadFITS reads a primary image +// (BITPIX -64 and -32) with its header cards and applies the +// BSCALE/BZERO scaling; SaveFITS writes one. LoadFITSTable and +// SaveFITSTable cover the binary and ASCII table extensions, where +// catalogues and observation logs live. +// +// HDF5 is read by LoadHDF5, which returns every dataset of a file by +// path, with the attributes of the groups it sits in merged into it. +// Superblocks 0 to 3, object headers of version 1 and 2, symbol-table +// and link-message groups, contiguous, compact and chunked storage and +// the deflate, shuffle and fletcher32 filters are supported. Dataset +// values land the core dtype their datatype declares: fixed-point data +// by stored width and signedness, the boolean enumeration convention +// as Bool, floating-point data as float64 or float32 by width. What is +// not supported, among it dense groups, the version 2 chunk B-tree, +// string datasets, every big-endian datatype, bit fields, non-boolean +// enumerations and unsigned 64-bit integers, is refused with an error +// naming it. SaveHDF5 writes the mirror image in the classic layout +// or, with Latest, the superblock 3 layout, and SaveHDF5Text writes +// fixed-length string datasets. +// +// NetCDF classic (CDF-1 and CDF-2) is the archival format of climate +// and ocean science: LoadNetCDF returns named dimensions, variables +// and global attributes, and SaveNetCDF writes CDF-1. Each variable +// lands the core dtype its classic type code carries: NC_BYTE as +// int8, NC_CHAR as uint8 raw bytes (CHAR carries bytes at the array +// level, never text), NC_SHORT as int16, NC_INT as int32, and +// NC_FLOAT and NC_DOUBLE as float64. A variable that lands a narrow +// dtype from a file is refused by SaveNetCDF, whose writer stores +// float64, float32 and int64 arrays only; convert with Astype first. +// Record dimensions are read and written in both directions; the +// writer stores float64, float32 and int64 arrays, and a type code +// beyond the classic six is refused by name. +// +// MapFloats, MapFloat32s and MapInts open a native-endian file as a +// read-only array without reading it, which suits data far larger than +// memory; release unmaps it and comes strictly last. SaveNativeFloats +// writes the format MapFloats reads. +// +// Each format is covered by a documented subset and the rest is +// refused by name, never half-read. Errors carry the library's +// "tensor: " prefix and the name of the call that produced them. +package io diff --git a/io/example_test.go b/io/example_test.go new file mode 100644 index 0000000..dfd47b4 --- /dev/null +++ b/io/example_test.go @@ -0,0 +1,237 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io_test + +// The godoc examples: one runnable, checked snippet per format the +// package speaks. `go test` executes them, so the documentation cannot +// rot. Each one writes into a fresh temporary directory and removes it +// again. + +import ( + "fmt" + "log" + "os" + "path/filepath" + "strings" + + tensor "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/io" +) + +// tempDir makes a fresh temporary directory for one example; the +// example removes it with a deferred os.RemoveAll. +func tempDir() string { + dir, err := os.MkdirTemp("", "tensor-io-example-") + if err != nil { + log.Fatal(err) + } + return dir +} + +// A 2-D array goes out as comma-separated text and reads back as the +// same numbers. +func ExampleSaveCSV() { + dir := tempDir() + defer os.RemoveAll(dir) + path := filepath.Join(dir, "readings.csv") + + a, err := tensor.FromFloats([]float64{18.5, 21.25, 19.75, 23, 20.5, 17.25}, 2, 3) + if err != nil { + log.Fatal(err) + } + if err := io.SaveCSV(path, a); err != nil { + log.Fatal(err) + } + back, err := io.LoadCSV(path, false) + if err != nil { + log.Fatal(err) + } + fmt.Println("shape:", back.Shape()) + fmt.Println("first:", back.FloatAt(0), "last:", back.FloatAt(back.Len()-1)) + // Output: + // shape: [2 3] + // first: 18.5 last: 17.25 +} + +// The stream form writes CSV to any writer and reads it from any +// reader, skipping a header row on the way back. +func ExampleSaveCSVWriter() { + a, err := tensor.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if err != nil { + log.Fatal(err) + } + var buf strings.Builder + if err := io.SaveCSVWriter(&buf, a); err != nil { + log.Fatal(err) + } + fmt.Print(buf.String()) + + back, err := io.LoadCSVReader(strings.NewReader("row,c1,c2\n"+buf.String()), true) + if err != nil { + log.Fatal(err) + } + fmt.Println("shape:", back.Shape(), "last:", back.FloatAt(3)) + // Output: + // 1,2 + // 3,4 + // shape: [2 2] last: 4 +} + +// A float64 image goes out as a FITS primary image with header cards +// and reads back with the cards beside the values. +func ExampleSaveFITS() { + dir := tempDir() + defer os.RemoveAll(dir) + path := filepath.Join(dir, "image.fits") + + a, err := tensor.FromFloats([]float64{1, 2.5, 3, 4, 5, 6}, 2, 3) + if err != nil { + log.Fatal(err) + } + headers := map[string]string{"OBJECT": "M31", "EXPTIME": "600"} + if err := io.SaveFITS(path, a, headers); err != nil { + log.Fatal(err) + } + img, header, err := io.LoadFITS(path) + if err != nil { + log.Fatal(err) + } + fmt.Println("shape:", img.Shape(), "dtype:", img.Dtype()) + fmt.Println("OBJECT:", header["OBJECT"], "EXPTIME:", header["EXPTIME"]) + fmt.Println("last:", img.FloatAt(img.Len()-1)) + // Output: + // shape: [2 3] dtype: float + // OBJECT: M31 EXPTIME: 600 + // last: 6 +} + +// A table extension holds a character column and a numeric column; the +// reader returns both, parallel to the file's column list. +func ExampleSaveFITSTable() { + dir := tempDir() + defer os.RemoveAll(dir) + path := filepath.Join(dir, "catalogue.fits") + + flux, err := tensor.FromFloats([]float64{1.5, 2.25, 3.75}, 3) + if err != nil { + log.Fatal(err) + } + cols := []io.FITSTableColumn{ + {Name: "source", Form: "8A", Text: []string{"alpha", "beta", "gamma"}}, + {Name: "flux", Unit: "Jy", Form: "D", Data: flux}, + } + if err := io.SaveFITSTable(path, false, cols, nil); err != nil { + log.Fatal(err) + } + table, err := io.LoadFITSTable(path) + if err != nil { + log.Fatal(err) + } + fmt.Println("kind:", table.Kind, "rows:", table.Rows) + fmt.Println("names:", table.Names, "unit:", table.Units[1]) + fmt.Println("source[1]:", table.Text[0][1], "flux[1]:", table.Columns[1].FloatAt(1)) + // Output: + // kind: BINTABLE rows: 3 + // names: [source flux] unit: Jy + // source[1]: beta flux[1]: 2.25 +} + +// Two datasets, one in a group, with attributes on the datasets and on +// the root and the group: the file reads back with the same paths, +// shapes and dtypes, and each dataset carries the attributes of the +// enclosing groups. +func ExampleSaveHDF5() { + dir := tempDir() + defer os.RemoveAll(dir) + path := filepath.Join(dir, "scan.h5") + + temp, err := tensor.FromFloat32s([]float32{1.5, 2.5, 3.5, 4.5}, 2, 2) + if err != nil { + log.Fatal(err) + } + counts, err := tensor.FromInts([]int64{100, 200, 300}, 3) + if err != nil { + log.Fatal(err) + } + sets := []io.HDF5Dataset{ + {Path: "/scan/temperature", Values: temp, Attrs: map[string]string{"units": "degC"}}, + {Path: "/scan/background", Values: counts, Attrs: map[string]string{"units": "counts"}}, + } + groupAttrs := map[string]map[string]string{ + "/": {"title": "cruise"}, + "/scan": {"instrument": "thermistor"}, + } + if err := io.SaveHDF5(path, sets, groupAttrs); err != nil { + log.Fatal(err) + } + got, err := io.LoadHDF5(path) + if err != nil { + log.Fatal(err) + } + for _, d := range got { + fmt.Printf("%s %v %s %s %s\n", d.Path, d.Shape, d.Values.Dtype(), d.Attrs["instrument"], d.Attrs["units"]) + } + // Output: + // /scan/background [3] int thermistor counts + // /scan/temperature [2 2] float32 thermistor degC +} + +// Dimensions, a variable with an attribute and a global attribute go +// through NetCDF classic and back; the writer stores float64 as +// NC_DOUBLE, which the reader lands as float64 again, and the integer +// type codes land their own dtypes the same way. +func ExampleSaveNetCDF() { + dir := tempDir() + defer os.RemoveAll(dir) + path := filepath.Join(dir, "field.nc") + + dims := []io.NetCDFDim{{Name: "lat", Length: 2}, {Name: "lon", Length: 3}} + temp, err := tensor.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + log.Fatal(err) + } + vars := []io.NetCDFVar{{ + Name: "temp", + Dims: []string{"lat", "lon"}, + Values: temp, + Attrs: map[string]string{"units": "degC"}, + }} + if err := io.SaveNetCDF(path, dims, vars, map[string]string{"title": "cruise"}); err != nil { + log.Fatal(err) + } + gotDims, gotVars, gotAttrs, err := io.LoadNetCDF(path) + if err != nil { + log.Fatal(err) + } + fmt.Println("dims:", gotDims) + fmt.Println("var:", gotVars[0].Name, gotVars[0].Dims, gotVars[0].Attrs["units"]) + fmt.Println("title:", gotAttrs["title"], "last:", gotVars[0].Values.FloatAt(5)) + // Output: + // dims: [{lat 2} {lon 3}] + // var: temp [lat lon] degC + // title: cruise last: 6 +} + +// A native-endian file of float64 values maps into a read-only array +// without being read; release unmaps it and comes strictly last. +func ExampleSaveNativeFloats() { + dir := tempDir() + defer os.RemoveAll(dir) + path := filepath.Join(dir, "values.bin") + + values := []float64{1.5, 2.5, 3.5, 4.5, 5.5} + if err := io.SaveNativeFloats(path, values); err != nil { + log.Fatal(err) + } + a, release, err := io.MapFloats(path, 0, len(values)) + if err != nil { + log.Fatal(err) + } + fmt.Println("mapped:", a.Len(), a.Dtype(), a.FloatAt(0), a.FloatAt(4)) + if err := release(); err != nil { + log.Fatal(err) + } + // Output: + // mapped: 5 float 1.5 5.5 +} diff --git a/io/extent_wrap_pins_test.go b/io/extent_wrap_pins_test.go new file mode 100644 index 0000000..2414ecf --- /dev/null +++ b/io/extent_wrap_pins_test.go @@ -0,0 +1,573 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "fmt" + "math" + "os" + "path/filepath" + "runtime" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression tests for hostile inputs across the io package: the size +// arithmetic of the HDF5 reader, the NetCDF record-size +// pre-pass, the FITS table encoders and decoders, the mmap count +// conversion, and the CSV shape contract. Every hostile file below is +// built by hand, byte by byte, because the point of each test is a +// declared size the writer of a well-formed file would never produce. + +// h5HostileFile returns an n-byte HDF5 file with the signature and a +// version 0 superblock (the classic layout: eight-byte addresses and +// lengths) whose root object header sits at offset 96. +func h5HostileFile(n int) []byte { + f := make([]byte, n) + copy(f, hdf5Magic) + f[8] = 0 // superblock version 0 + f[13] = 8 + f[14] = 8 + binary.LittleEndian.PutUint64(f[32:], math.MaxUint64) // free space undefined + binary.LittleEndian.PutUint64(f[40:], uint64(n)) // end of file + binary.LittleEndian.PutUint64(f[48:], math.MaxUint64) // driver undefined + binary.LittleEndian.PutUint64(f[64:], 96) // root object header + return f +} + +// h5Msg is one object header message: its type and its body bytes. +type h5Msg struct { + typ uint16 + body []byte +} + +// h5ObjectHeader writes a version 1 object header at off carrying msgs +// and returns the offset just past its message region. +func h5ObjectHeader(f []byte, off int, msgs ...h5Msg) int { + f[off] = 1 + binary.LittleEndian.PutUint16(f[off+2:], uint16(len(msgs))) + binary.LittleEndian.PutUint32(f[off+4:], 1) // reference count + region := off + 16 + size := 0 + for _, m := range msgs { + size = alignUp(size+8+len(m.body), 8) + } + binary.LittleEndian.PutUint32(f[off+8:], uint32(size)) + for _, m := range msgs { + binary.LittleEndian.PutUint16(f[region:], m.typ) + binary.LittleEndian.PutUint16(f[region+2:], uint16(len(m.body))) + copy(f[region+8:], m.body) + region += alignUp(8+len(m.body), 8) + } + return region +} + +// h5Dataspace renders a version 1 dataspace message with the given +// extents, each written as an eight-byte length. +func h5Dataspace(dims ...uint64) []byte { + m := make([]byte, 8+8*len(dims)) + m[0] = 1 // version 1 + if len(dims) == 0 { + return m + } + m[1] = byte(len(dims)) + for i, d := range dims { + binary.LittleEndian.PutUint64(m[8+8*i:], d) + } + return m +} + +// h5FloatType renders a version 1 floating-point datatype message of the +// given element width. +func h5FloatType(size uint32) []byte { + m := make([]byte, 20) + m[0] = 0x11 // version 1, class 1 (floating-point) + binary.LittleEndian.PutUint32(m[4:], size) + return m +} + +// h5ChunkLayoutV4 renders a version 4 chunked layout message whose chunk +// dimensions are eight bytes wide, the flag the version 4 message +// carries, so an extent beyond 2^32 can be declared at all. +func h5ChunkLayoutV4(btree uint64, dims ...uint64) []byte { + m := make([]byte, 3+8+1+8*len(dims)) + m[0] = 4 + m[1] = 2 // chunked + m[2] = byte(len(dims)) + binary.LittleEndian.PutUint64(m[3:], btree) + m[11] = 1 // eight-byte chunk dimensions + for i, d := range dims { + binary.LittleEndian.PutUint64(m[12+8*i:], d) + } + return m +} + +// h5ChunkLayoutV3 renders a version 3 chunked layout message with +// four-byte chunk dimensions. +func h5ChunkLayoutV3(btree uint64, dims ...uint32) []byte { + m := make([]byte, 3+8+4*len(dims)) + m[0] = 3 + m[1] = 2 + m[2] = byte(len(dims)) + binary.LittleEndian.PutUint64(m[3:], btree) + for i, d := range dims { + binary.LittleEndian.PutUint32(m[11+4*i:], d) + } + return m +} + +// h5ContiguousLayout renders a version 3 contiguous layout message. +func h5ContiguousLayout(addr, size uint64) []byte { + m := make([]byte, 2+8+8) + m[0] = 3 + m[1] = 1 // contiguous + binary.LittleEndian.PutUint64(m[2:], addr) + binary.LittleEndian.PutUint64(m[10:], size) + return m +} + +// h5ChunkTree writes a one-entry leaf chunk B-tree at off for a dataset +// of the given rank: one chunk of size stored bytes at chunkAt, with no +// filter mask, at the given chunk offsets. The key carries one more +// offset than the dataset has axes, the element size, which is why the +// entry is 8+(8*(rank+1))+offSize bytes long. +func h5ChunkTree(f []byte, off, rank int, size uint32, offsets []uint64, chunkAt uint64) { + copy(f[off:], hdf5Tree) + f[off+4] = 1 // chunk tree + f[off+5] = 0 // leaf level + binary.LittleEndian.PutUint16(f[off+6:], 1) + p := off + 24 + binary.LittleEndian.PutUint32(f[p:], size) + binary.LittleEndian.PutUint32(f[p+4:], 0) // filter mask + for i := range rank { + binary.LittleEndian.PutUint64(f[p+8+8*i:], offsets[i]) + } + binary.LittleEndian.PutUint64(f[p+8+8*rank:], 8) // the element-size key + binary.LittleEndian.PutUint64(f[p+8+8*(rank+1):], chunkAt) +} + +// TestLoadHDF5ChunkExtentWrap pins the chunk-dimension bound. The chunk +// shape 2^32 by 2^29 float64 elements is 2^64 bytes, which wraps to +// zero: the "chunk above the reader budget" guard and the "chunk holds a +// full chunk" guard both saw 0, so the chunk strides were walked out of +// the stored chunk (a slice out of range on a 2^61-element declaration, +// and a silent read of the file's other bytes on a milder one). +func TestLoadHDF5ChunkExtentWrap(t *testing.T) { + cases := []struct { + name string + dims []uint64 // the chunk's shape plus the element-size slot + }{ + {"2^32 by 2^29 elements", []uint64{1 << 32, 1 << 29, 8}}, + {"2^61 elements", []uint64{1 << 61, 1, 8}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + const btree, chunkAt, n = 320, 448, 512 + f := h5HostileFile(n) + end := h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(2, 1)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV4(btree, tc.dims...)}, + ) + if end > btree { + t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree) + } + h5ChunkTree(f, btree, 2, 16, []uint64{0, 0}, chunkAt) + path := writeHostile(t, "chunkextent.h5", f) + sets, err := LoadHDF5(path) + if err == nil { + t.Fatalf("LoadHDF5 accepted a chunk whose extent wraps: %d datasets", len(sets)) + } + if !strings.Contains(err.Error(), "budget") { + t.Fatalf("error = %v, want the reader-budget refusal", err) + } + }) + } +} + +// TestLoadHDF5DatasetExtentWrap pins the dataspace bound on both storage +// classes: a shape of 2^32 by 2^32 float64 elements wraps the byte count +// to zero, which passed the cap, so the output buffer was allocated +// empty while the copy needed a whole element (the chunked class +// panicked), and the contiguous class returned a dataset claiming 2^64 +// elements with an empty payload and no error. +func TestLoadHDF5DatasetExtentWrap(t *testing.T) { + t.Run("chunked", func(t *testing.T) { + const btree, chunkAt, n = 320, 448, 512 + f := h5HostileFile(n) + end := h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(1<<32, 1<<32)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(btree, 1, 1, 8)}, + ) + if end > btree { + t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree) + } + h5ChunkTree(f, btree, 2, 8, []uint64{0, 0}, chunkAt) + sets, err := LoadHDF5(writeHostile(t, "shapedwrap.h5", f)) + if err == nil { + t.Fatalf("LoadHDF5 accepted a wrapping dataspace: %d datasets", len(sets)) + } + if !strings.Contains(err.Error(), "budget") { + t.Fatalf("error = %v, want the reader-budget refusal", err) + } + }) + t.Run("contiguous", func(t *testing.T) { + const dataAt, n = 448, 512 + f := h5HostileFile(n) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(1<<32, 1<<32)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)}, + ) + sets, err := LoadHDF5(writeHostile(t, "contigwrap.h5", f)) + if err == nil { + shape := []int(nil) + if len(sets) > 0 { + shape = sets[0].Shape + } + t.Fatalf("LoadHDF5 returned %v for a wrapping dataspace, want an error", shape) + } + if !strings.Contains(err.Error(), "budget") { + t.Fatalf("error = %v, want the reader-budget refusal", err) + } + }) +} + +// TestLoadHDF5AttributeExtentWrap pins the attribute dataspace bound: a +// 3 by 2^62 element dataspace wraps its element count to a negative +// number in the multiply-first check, which let it past the bounds test +// and into an allocation with a negative capacity. The attribute must be +// refused (dropped), not accepted with a wrapped value. +func TestLoadHDF5AttributeExtentWrap(t *testing.T) { + const n = 512 + f := h5HostileFile(n) + // The attribute message: version, flags, then the three size fields + // with their fields behind them. A 2^62 extent is far past the + // message, so no legal value can follow it. + attr := make([]byte, 88) + attr[0] = 1 + binary.LittleEndian.PutUint16(attr[2:], 2) // name size + binary.LittleEndian.PutUint16(attr[4:], 20) // datatype size + binary.LittleEndian.PutUint16(attr[6:], 24) // dataspace size + copy(attr[8:], "n\x00") + copy(attr[16:], h5FloatType(8)) + copy(attr[40:], h5Dataspace(3, 1<<62)) + // A hard link to a one-element dataset, so the attribute's fate is + // observable: an accepted attribute lands in the dataset's map. + link := func(addr uint64) []byte { + return binary.LittleEndian.AppendUint64([]byte{1, 0, 1, 'd'}, addr) + } + datasetAt := 288 + end := h5ObjectHeader(f, 96, + h5Msg{hdf5MsgLink, link(uint64(datasetAt))}, + h5Msg{hdf5MsgAttribute, attr}, + ) + if end > datasetAt { + t.Fatalf("the test object header runs to %d, past the dataset at %d", end, datasetAt) + } + h5ObjectHeader(f, datasetAt, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, 8)}, + ) + sets, err := LoadHDF5(writeHostile(t, "attrextent.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5 refused the file: %v", err) + } + if len(sets) != 1 || sets[0].Path != "/d" { + t.Fatalf("datasets = %d, want the linked /d", len(sets)) + } + if v, ok := sets[0].Attrs["n"]; ok { + t.Fatalf("the attribute with a wrapping dataspace was accepted as %q", v) + } +} + +// TestLoadHDF5MessageCountBounded pins the object header preallocation: +// the slice used to be sized from the header's declared message count +// (65535 messages of 32 bytes is 2 MiB) before the region that would +// have to hold them was consulted, so 16 bytes of file ordered a +// two-megabyte allocation. +func TestLoadHDF5MessageCountBounded(t *testing.T) { + f := h5HostileFile(160) + f[96] = 1 + binary.LittleEndian.PutUint16(f[98:], 65535) // declared messages + binary.LittleEndian.PutUint32(f[100:], 1) // reference count + binary.LittleEndian.PutUint32(f[104:], 16) // the region that holds them + path := writeHostile(t, "msgcount.h5", f) + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + _, err := LoadHDF5(path) + runtime.ReadMemStats(&after) + if err == nil { + t.Fatal("LoadHDF5 accepted a header whose messages run past the region") + } + if used := after.TotalAlloc - before.TotalAlloc; used > 512<<10 { + t.Fatalf("loading a 160-byte file allocated %d bytes: the declared message count still sizes the preallocation", used) + } +} + +// TestLoadNetCDFRecordSizeWrap pins the record-size pre-pass. It used to +// multiply every record variable's non-record dimensions before any +// variable was validated, so two hostile declarations inflated the +// record size for an innocuous variable to 2^54 bytes: one element per +// record then made (numrecs-1)*recordSize exactly 2^64, which wrapped to +// zero, the span check passed and the reader walked 2^54 bytes into a +// 4 KiB file. +func TestLoadNetCDFRecordSizeWrap(t *testing.T) { + const numrecs = 1<<10 + 1 + // The arithmetic the reader performed: two slabs of 2^25 by 2^25 + // doubles are 2^53 bytes each and the innocuous variable's slab is + // four, so the record size is 2^54+4 and the last record's slab + // offset wraps to 2^64+4096, i.e. 4096. A file larger than that + // passes the span check and the following records then index far + // past the file. + recordSize := 2*int64(1<<53) + 4 + if span := (numrecs-1)*recordSize + 1; span != 1<<12+1 { + t.Fatalf("construction is wrong: the wrapped span is %d, want %d", span, 1<<12+1) + } + + var b []byte + b = append(b, 'C', 'D', 'F', 1) + b = binary.BigEndian.AppendUint32(b, numrecs) + // dim_list: the record dimension leads, then one one-element + // dimension and the two hostile ones. + b = binary.BigEndian.AppendUint32(b, ncTagDimension) + b = binary.BigEndian.AppendUint32(b, 4) + b = hostileNCName(b, "rec") + b = binary.BigEndian.AppendUint32(b, 0) + b = hostileNCName(b, "one") + b = binary.BigEndian.AppendUint32(b, 1) + b = hostileNCName(b, "big1") + b = binary.BigEndian.AppendUint32(b, 1<<25) + b = hostileNCName(b, "big2") + b = binary.BigEndian.AppendUint32(b, 1<<25) + // gatt_list: absent. + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, 0) + // var_list: the innocuous record variable, then the two that inflate + // the record size. + b = binary.BigEndian.AppendUint32(b, ncTagVariable) + b = binary.BigEndian.AppendUint32(b, 3) + b = hostileNCName(b, "a") + b = binary.BigEndian.AppendUint32(b, 2) // rank: rec, one + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, 1) + b = binary.BigEndian.AppendUint32(b, 0) // variable attributes: absent + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, ncTypeByte) + b = binary.BigEndian.AppendUint32(b, 4) // vsize + b = binary.BigEndian.AppendUint32(b, 0) // begin + for i := range 2 { + b = hostileNCName(b, fmt.Sprintf("b%d", i)) + b = binary.BigEndian.AppendUint32(b, 3) // rank: rec, big1, big2 + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, 2) + b = binary.BigEndian.AppendUint32(b, 3) + b = binary.BigEndian.AppendUint32(b, 0) // variable attributes: absent + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, ncTypeDouble) + b = binary.BigEndian.AppendUint32(b, 0) // vsize + b = binary.BigEndian.AppendUint32(b, 0) // begin + } + // A file long enough for the innocuous variable's own records and for + // the wrapped span the buggy check computes. + for len(b) < 8192 { + b = append(b, 0) + } + path := writeHostile(t, "recsize-wrap.nc", b) + _, _, _, err := LoadNetCDF(path) + if err == nil { + t.Fatal("LoadNetCDF accepted a record variable whose dimensions cannot fit the file") + } + if !strings.Contains(err.Error(), "more elements than the file holds") { + t.Fatalf("error = %v, want the per-variable element bound", err) + } +} + +// TestSaveFITSTableASCIIFloat32 pins the ASCII float encoder for a +// float32 column. The encoder read the float64 payload first and only +// then substituted the float32 one, but RawFloats is nil for every +// float32 array, so the first row of an ordinary float32 column indexed +// a nil slice. +func TestSaveFITSTableASCIIFloat32(t *testing.T) { + vals := []float32{1.5, -2.25, 3e8} + f32 := core.New(core.Float32, len(vals)) + copy(f32.RawFloat32s(), vals) + cols := []FITSTableColumn{{Name: "F", Form: "E12.4", Data: f32}} + path := filepath.Join(t.TempDir(), "f32ascii.fits") + if err := SaveFITSTable(path, true, cols, nil); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + for i, want := range vals { + if got := table.Columns[0].FloatAt(i); got != float64(want) { + t.Fatalf("F[%d] = %g, want %g", i, got, float64(want)) + } + } +} + +// TestLoadFITSTableFortranNumbers pins the number forms a Fortran writer +// puts in an ASCII table: strconv.ParseFloat accepts neither the Dw.d +// exponent ("1.5D+03") nor the exponent-less form ("1.5+03"), and one +// unparseable cell refused the whole table. +func TestLoadFITSTableFortranNumbers(t *testing.T) { + hdr := cardBlock( + card("XTENSION= 'TABLE '"), + card("BITPIX = 8"), + card("NAXIS = 2"), + card("NAXIS1 = 20"), + card("NAXIS2 = 4"), + card("TFIELDS = 1"), + card("TBCOL1 = 1"), + card("TFORM1 = 'D20.12'"), + card("END"), + ) + rows := []string{"1.5D+03", "1.5+03", "-2.5d-2", "1.5E+03"} + payload := make([]byte, 0, 20*len(rows)) + for _, r := range rows { + payload = append(payload, fmt.Sprintf("%-20s", r)...) + } + path := writeHostile(t, "fortran.fits", append(hdr, payload...)) + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable refused a Fortran table: %v", err) + } + want := []float64{1500, 1500, -0.025, 1500} + for i, w := range want { + if got := table.Columns[0].FloatAt(i); got != w { + t.Fatalf("row %d (%q) = %g, want %g", i, rows[i], got, w) + } + } +} + +// TestLoadFITSTableScaleOverflow pins the integral-scaling bound. The +// guard used 9.3e18, which is above MaxInt64 (9.223372036854775807e18), +// so a scaled value in [2^63, 9.3e18) reached an out-of-range +// float-to-integer conversion: the column kept the int dtype and the +// value came back as MinInt64 instead of the documented promotion to +// float64. +func TestLoadFITSTableScaleOverflow(t *testing.T) { + const raw = int64(1) << 62 + cols := []FITSTableColumn{{Name: "V", Form: "K", Data: mustInts(t, []int64{raw}, 1)}} + path := filepath.Join(t.TempDir(), "scaled.fits") + if err := SaveFITSTable(path, false, cols, map[string]string{"TSCAL1": "2", "TZERO1": "0"}); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + col := table.Columns[0] + want := float64(raw) * 2 // 2^63, exactly representable in float64 + if col.Dtype() != core.Float { + t.Fatalf("dtype = %s, want float64: a scaling past MaxInt64 must promote the column", col.Dtype()) + } + if got := col.FloatAt(0); got != want { + t.Fatalf("scaled value = %g, want %g", got, want) + } +} + +// TestMapCountsWithoutWrap pins the mmap count arithmetic: the byte +// length used to be int64(n)*elementSize, which wraps for a hostile +// count, so the wrapped length passed the file-size checks and the typed +// view was then built with the original count, which unsafe.Slice +// rejects with a panic instead of an error. +func TestMapCountsWithoutWrap(t *testing.T) { + path := filepath.Join(t.TempDir(), "small.bin") + if err := SaveNativeFloats(path, []float64{1, 2, 3, 4}); err != nil { + t.Fatal(err) + } + // 2^62+3 float64 values are 2^65+24 bytes, which wraps to 24: the + // file holds 32. The others already failed the non-positive-length + // check by luck and must keep failing it. + for _, n := range []int{1<<62 + 3, 1 << 61, 1 << 60, math.MaxInt} { + if _, _, err := MapFloats(path, 0, n); err == nil { + t.Errorf("MapFloats(n = %d) accepted a count the file cannot hold", n) + } + if _, _, err := MapFloat32s(path, 0, n); err == nil { + t.Errorf("MapFloat32s(n = %d) accepted a count the file cannot hold", n) + } + if _, _, err := MapInts(path, 0, n); err == nil { + t.Errorf("MapInts(n = %d) accepted a count the file cannot hold", n) + } + } +} + +// TestSaveCSVZeroColumns pins the shape contract: an array with rows and +// no columns wrote one empty record per row, and encoding/csv reads +// blank lines as no records at all, so (3,0) came back as (0,0). It is +// refused instead, because no CSV text can carry the column count. +func TestSaveCSVZeroColumns(t *testing.T) { + a, err := core.FromFloats(nil, 3, 0) + if err != nil { + t.Fatalf("FromFloats(nil, 3, 0): %v", err) + } + var sb strings.Builder + if err := SaveCSVWriter(&sb, a); err == nil { + t.Fatalf("SaveCSVWriter accepted the shape %v and wrote %q, which reads back as a different shape", + a.Shape(), sb.String()) + } + // A 2-D array with columns still round-trips. + ok := mustFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + var back strings.Builder + if err := SaveCSVWriter(&back, ok); err != nil { + t.Fatalf("SaveCSVWriter: %v", err) + } + got, err := LoadCSVReader(strings.NewReader(back.String()), false) + if err != nil { + t.Fatalf("LoadCSVReader: %v", err) + } + if got.Shape()[0] != 2 || got.Shape()[1] != 3 { + t.Fatalf("round trip gave %v, want [2 3]", got.Shape()) + } +} + +// TestIOSourceTypography pins the plain-ASCII rule for the package's own +// source: an em dash, an en dash, a Unicode minus, a middle dot or a +// multiplication sign in a comment or a message is a typographic symbol +// where the repository's style asks for ASCII. +func TestIOSourceTypography(t *testing.T) { + banned := []struct { + r rune + name string + }{ + {'\u2014', "em dash"}, + {'\u2013', "en dash"}, + {'\u2212', "Unicode minus"}, + {'\u00b7', "middle dot"}, + {'\u00d7', "multiplication sign"}, + } + entries, err := os.ReadDir(".") + if err != nil { + t.Fatalf("ReadDir: %v", err) + } + checked := 0 + for _, e := range entries { + name := e.Name() + if !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") { + continue + } + data, err := os.ReadFile(name) + if err != nil { + t.Fatalf("ReadFile(%s): %v", name, err) + } + checked++ + for _, b := range banned { + if strings.ContainsRune(string(data), b.r) { + t.Errorf("%s carries a %s (U+%04X)", name, b.name, b.r) + } + } + } + if checked == 0 { + t.Fatal("no source files were checked") + } +} diff --git a/io/fits.go b/io/fits.go new file mode 100644 index 0000000..d4dbb28 --- /dev/null +++ b/io/fits.go @@ -0,0 +1,561 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "fmt" + "maps" + "math" + "os" + "slices" + "strconv" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// FITS image I/O. The Flexible Image Transport System is astronomy's +// archival format: a self-describing header of 80-character ASCII +// cards in 2880-byte blocks, followed by big-endian binary data +// padded to the same block size. Every observation archived by a +// telescope in the last four decades reads back with the same parser, +// which is the property that makes a format worth speaking. +// +// This implementation covers the primary HDU image: BITPIX -64 +// (float64) and -32 (float32), any rank with positive axes. FITS +// orders axes Fortran-style with NAXIS1 varying fastest, the opposite +// of Go's row-major convention, so NAXISj is declared from the shape +// in reverse and the flat payload needs no permutation. Extensions +// (XTENSION), tables and the integer bit depths are refused with an +// error rather than half-read. + +// fitsCardsPerBlock is the number of 80-byte cards in one 2880-byte +// FITS block. +const fitsCardsPerBlock = 2880 / 80 + +// SaveFITS writes a float64 or float32 array as a FITS primary image +// with the given header entries. Keywords are uppercased and must be +// 1 to 8 characters from A-Z, 0-9, '-' and '_' (the format's reserved +// SIMPLE, BITPIX, NAXIS, NAXISn, EXTEND and END are refused); values +// are written as FITS strings, at most 68 characters after the +// format's quote escaping. +func SaveFITS(path string, a *core.Array, headers map[string]string) error { + var bitpix int + switch a.Dtype() { + case core.Float: + bitpix = -64 + case core.Float32: + bitpix = -32 + default: + return base.Errf("SaveFITS: supports float64 and float32 arrays, got dtype %s", a.Dtype()) + } + if a.NDim() == 0 { + return base.Errf("SaveFITS: the image needs at least one axis") + } + for i, d := range a.Shape() { + if d <= 0 { + return base.Errf("SaveFITS: axis %d has extent %d, every axis must be positive", i+1, d) + } + } + + cards := []string{ + fitsBoolCard("SIMPLE", true), + fitsIntCard("BITPIX", bitpix), + fitsIntCard("NAXIS", a.NDim()), + } + // NAXIS1 is the fastest-varying axis, the last one in Go's + // row-major order. + for j := range a.NDim() { + cards = append(cards, fitsIntCard("NAXIS"+strconv.Itoa(j+1), a.Shape()[a.NDim()-1-j])) + } + cards = append(cards, fitsBoolCard("EXTEND", true)) + userCards, err := fitsUserCards(headers) + if err != nil { + return base.Errf("SaveFITS: %w", err) + } + cards = append(cards, userCards...) + cards = append(cards, fitsEndCard()) + + elem := a.Len() + width := 8 + if a.Dtype() == core.Float32 { + width = 4 + } + // The final size is known up front: both the header and the payload + // pad to whole blocks, so one allocation serves the whole file. + out := make([]byte, 0, fitsBlockSize(len(cards))+fitsBlockSize(elem*width)) + out = fitsAppendCards(out, cards) + if a.Dtype() == core.Float { + raw := a.RawFloats() + for i := range elem { + out = binary.BigEndian.AppendUint64(out, math.Float64bits(raw[i])) + } + } else { + raw := a.RawFloat32s() + for i := range elem { + out = binary.BigEndian.AppendUint32(out, math.Float32bits(raw[i])) + } + } + // Zero bytes pad the data to the block boundary. + out = fitsAppendZeroPad(out) + return os.WriteFile(path, out, 0o644) +} + +// LoadFITS reads a FITS primary image into a float64 (BITPIX -64) or +// float32 (BITPIX -32) array, returning every non-structural header +// entry alongside it. String values are unquoted and unescaped, +// logical values come back as "T" or "F", numbers as their literal +// text; COMMENT, HISTORY and blank cards carry no value and are +// skipped. +func LoadFITS(path string) (*core.Array, map[string]string, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, nil, base.Errf("LoadFITS: %w", err) + } + return parseFITS(data) +} + +// parseFITS decodes a FITS primary image from raw bytes. +func parseFITS(data []byte) (*core.Array, map[string]string, error) { + if len(data) < 80 { + return nil, nil, base.Errf("LoadFITS: file is shorter than one header card") + } + if string(data[:9]) == "XTENSION " { + return nil, nil, base.Errf("LoadFITS: extensions are not supported, only the primary image") + } + if strings.TrimRight(string(data[:8]), " ") != "SIMPLE" { + return nil, nil, base.Errf("LoadFITS: the first card must be SIMPLE") + } + + cards, dataAt, err := scanFITSCards(data, 0) + if err != nil { + return nil, nil, base.Errf("LoadFITS: %w", err) + } + + var ( + headers = map[string]string{} + bitpix int + naxis = -1 + axisVals = map[int]int{} + dims []int + ) + for _, c := range cards { + switch { + case c.key == "SIMPLE": + if c.value != "T" { + return nil, nil, base.Errf("LoadFITS: SIMPLE = F marks a non-conformant file") + } + case c.key == "EXTEND": + // Structural; EXTEND still reports itself to the caller. + headers[c.key] = c.value + case c.key == "BITPIX": + v, cerr := strconv.Atoi(c.value) + if cerr != nil { + return nil, nil, base.Errf("LoadFITS: %s = %q is not an integer", c.key, c.value) + } + bitpix = v + case c.key == "NAXIS": + v, cerr := strconv.Atoi(c.value) + if cerr != nil { + return nil, nil, base.Errf("LoadFITS: %s = %q is not an integer", c.key, c.value) + } + naxis = v + default: + // NAXISn is an axis only when a number follows the prefix, + // the rule fitsCheckKeyword applies on write kept symmetric; + // NAXISREF and every other spelling is a user keyword. The + // value is stored under its axis number, so the card order + // cannot re-bind the axes. + if fitsAxisKeyword(c.key) { + v, cerr := strconv.Atoi(c.value) + if cerr != nil { + return nil, nil, base.Errf("LoadFITS: %s = %q is not an integer", c.key, c.value) + } + j, _ := strconv.Atoi(c.key[len("NAXIS"):]) + if j < 1 { + return nil, nil, base.Errf("LoadFITS: %s is not an axis keyword", c.key) + } + if _, dup := axisVals[j]; dup { + return nil, nil, base.Errf("LoadFITS: %s repeats", c.key) + } + axisVals[j] = v + continue + } + headers[c.key] = c.value + } + } + // A zero-axis primary HDU (the standard container for extension + // files) answers an empty image. + if naxis == 0 { + return core.New(core.Float, 0), headers, nil + } + if bitpix != -64 && bitpix != -32 { + return nil, nil, base.Errf("LoadFITS: BITPIX %d is not supported (want -64 or -32)", bitpix) + } + // NAXISn values are bound by their axis number, not by card order: + // a header writing NAXIS2 before NAXIS1, or omitting an axis, is + // malformed and must be refused rather than silently re-bound. + if naxis < 0 { + return nil, nil, base.Errf("LoadFITS: NAXIS = %d is negative or the card is missing", naxis) + } + // The card count gates the allocation: a hostile NAXIS far beyond + // the NAXISn cards the file actually carries is refused here, not + // turned into a slice of that length. + if len(axisVals) != naxis { + return nil, nil, base.Errf("LoadFITS: NAXIS = %d with %d NAXISn cards", naxis, len(axisVals)) + } + dims = make([]int, naxis) + for j := 1; j <= naxis; j++ { + d, ok := axisVals[j] + if !ok { + return nil, nil, base.Errf("LoadFITS: NAXIS%d is missing under NAXIS = %d", j, naxis) + } + dims[j-1] = d + } + width := -bitpix / 8 + if dataAt > len(data) { + return nil, nil, base.Errf("LoadFITS: data is truncated (%d bytes present)", len(data)) + } + // Every axis extent and running product is bounded by the bytes + // the file actually holds, so a hostile header cannot overflow the + // int product before the truncation check rejects it. + avail := (len(data) - dataAt) / width + shape := make([]int, naxis) + total := 1 + for i, d := range dims { + if d <= 0 { + return nil, nil, base.Errf("LoadFITS: NAXIS%d = %d, every axis must be positive", i+1, d) + } + if d > avail || total > avail/d { + return nil, nil, base.Errf("LoadFITS: data is truncated (%d bytes present, more needed)", + len(data)-dataAt) + } + shape[naxis-1-i] = d + total *= d + } + + payload := data[dataAt : dataAt+total*width] + var out *core.Array + if bitpix == -64 { + out = core.New(core.Float, shape...) + raw := out.RawFloats() + for i := range total { + raw[i] = math.Float64frombits(binary.BigEndian.Uint64(payload[i*8:])) + } + } else { + out = core.New(core.Float32, shape...) + raw := out.RawFloat32s() + for i := range total { + raw[i] = math.Float32frombits(binary.BigEndian.Uint32(payload[i*4:])) + } + } + // BSCALE/BZERO scaling: physical = raw*scale + zero (FITS 4.1). + // A silent skip would hand back storage values as physical ones, so + // the affine map is applied whenever the keywords deviate from the + // identity; float32 values are computed in float64 and rounded once. + scale, serr := fitsScaledHeader(headers, "BSCALE", 1) + if serr != nil { + return nil, nil, base.Errf("LoadFITS: %w", serr) + } + zero, zerr := fitsScaledHeader(headers, "BZERO", 0) + if zerr != nil { + return nil, nil, base.Errf("LoadFITS: %w", zerr) + } + if scale != 1 || zero != 0 { + if out.Dtype() == core.Float { + raw := out.RawFloats() + for i := range raw { + raw[i] = raw[i]*scale + zero + } + } else { + raw := out.RawFloat32s() + for i := range raw { + raw[i] = float32(float64(raw[i])*scale + zero) + } + } + } + return out, headers, nil +} + +// fitsScaledHeader parses a floating-point header entry, falling back +// to def when the keyword is absent. A present but malformed value is +// an error, never silently the default. +func fitsScaledHeader(headers map[string]string, key string, def float64) (float64, error) { + v, ok := headers[key] + if !ok || v == "" { + return def, nil + } + f, err := strconv.ParseFloat(strings.TrimSpace(v), 64) + if err != nil { + return 0, base.Errf("%s = %q is not a number", key, v) + } + return f, nil +} + +// fitsValue extracts the value field of a card: everything after the +// "= " indicator, minus any trailing comment, with FITS string +// quoting resolved. The second return says whether the value was a +// quoted string, which the long-string CONTINUE convention keys on. +func fitsValue(card string) (string, bool, error) { + field := card[10:] + if strings.HasPrefix(strings.TrimLeft(field, " "), "'") { + // A quoted string: '' inside escapes one quote. + var b strings.Builder + in := field[strings.Index(field, "'"):] + i := 1 + for i < len(in) { + if in[i] == '\'' { + if i+1 < len(in) && in[i+1] == '\'' { + b.WriteByte('\'') + i += 2 + continue + } + return strings.TrimRight(b.String(), " "), true, nil + } + b.WriteByte(in[i]) + i++ + } + return "", false, base.Errf("card %q has an unterminated string value", card[:min(20, len(card))]) + } + // Free-format value: cut at the comment slash and trim. + if slash := strings.IndexByte(field, '/'); slash >= 0 { + field = field[:slash] + } + return strings.TrimSpace(field), false, nil +} + +// fitsContinueString extracts the quoted segment of a CONTINUE card. +// The keyword occupies columns 1-8 and no "= " indicator follows, so +// the string opens at the first quote anywhere in the card. +func fitsContinueString(card string) (string, error) { + q := strings.IndexByte(card, '\'') + if q < 0 { + return "", base.Errf("CONTINUE card %q has no quoted segment", card[:min(20, len(card))]) + } + var b strings.Builder + i := q + 1 + for i < len(card) { + if card[i] == '\'' { + if i+1 < len(card) && card[i+1] == '\'' { + b.WriteByte('\'') + i += 2 + continue + } + return b.String(), nil + } + b.WriteByte(card[i]) + i++ + } + return "", base.Errf("CONTINUE card %q has an unterminated string", card[:min(20, len(card))]) +} + +// fitsCheckKeyword validates a header keyword the caller supplies. +// The format's commentary keywords COMMENT and HISTORY carry no +// value, so they are refused like the structural ones. +func fitsCheckKeyword(kw string) error { + if len(kw) < 1 || len(kw) > 8 { + return base.Errf("keyword %q must be 1 to 8 characters", kw) + } + for _, r := range kw { + switch { + case r >= 'A' && r <= 'Z', r >= '0' && r <= '9', r == '-', r == '_': + default: + return base.Errf("keyword %q may only contain A-Z, 0-9, '-' and '_'", kw) + } + } + switch kw { + case "SIMPLE", "BITPIX", "NAXIS", "EXTEND", "END", "COMMENT", "HISTORY": + return base.Errf("keyword %q is reserved by the format", kw) + case "BSCALE", "BZERO": + // SaveFITS writes physical values directly; a user scaling card + // would make conforming readers scale them a second time. + return base.Errf("keyword %q is reserved by the format (values are written unscaled)", kw) + } + if rest, ok := strings.CutPrefix(kw, "NAXIS"); ok && rest != "" { + if _, err := strconv.Atoi(rest); err == nil { + return base.Errf("keyword %q is reserved by the format", kw) + } + } + return nil +} + +// fitsAxisKeyword reports whether key is a structural NAXISn card: the +// prefix followed by a number, the same rule fitsCheckKeyword applies +// when it refuses a reserved keyword on write. Any other spelling of +// the prefix (NAXISREF and the like) is a user keyword, and counting it +// as an axis used to make the reader answer "NAXIS = 1 with 2 NAXISn +// cards" for a file that carries one. +func fitsAxisKeyword(key string) bool { + rest, ok := strings.CutPrefix(key, "NAXIS") + if !ok { + return false + } + _, err := strconv.Atoi(rest) + return err == nil +} + +// fitsLookupN reads headers[prefix+index] without building the key on +// the heap: the digits are formatted into a stack buffer, and the map +// lookup over that buffer compiles to a lookup over the bytes with no +// string conversion. +func fitsLookupN(headers map[string]string, prefix string, index int) string { + var kb [24]byte + b := append(kb[:0], prefix...) + b = strconv.AppendInt(b, int64(index), 10) + return headers[string(b)] +} + +// fitsIntCard renders an integer-valued card with the value +// right-justified in columns 11 to 30. +func fitsIntCard(keyword string, v int) string { + return fitsPadCard(fmt.Sprintf("%-8s= %20d", keyword, v)) +} + +// fitsBoolCard renders a logical-valued card. +func fitsBoolCard(keyword string, v bool) string { + t := "F" + if v { + t = "T" + } + return fitsPadCard(fmt.Sprintf("%-8s= %20s", keyword, t)) +} + +// fitsStringCard renders a string-valued card; the format pads the +// quoted value to at least eight characters. +func fitsStringCard(keyword, v string) (string, error) { + escaped := strings.ReplaceAll(v, "'", "''") + inner := escaped + if len(inner) < 8 { + inner += strings.Repeat(" ", 8-len(inner)) + } + body := fmt.Sprintf("%-8s= '%s'", keyword, inner) + if len(body) > 80 { + return "", base.Errf("value for %q does not fit a card after quote escaping (%d characters)", + keyword, len(escaped)) + } + return fitsPadCard(body), nil +} + +// fitsEndCard renders the header terminator. +func fitsEndCard() string { + return fitsPadCard("END") +} + +// fitsPadCard right-pads a card body with spaces to the full 80 bytes. +func fitsPadCard(body string) string { + return body + strings.Repeat(" ", 80-len(body)) +} + +// fitsUserCards renders the caller's header entries as cards, keywords +// uppercased and sorted so the output is deterministic. +func fitsUserCards(headers map[string]string) ([]string, error) { + cards := make([]string, 0, len(headers)) + for _, key := range slices.Sorted(maps.Keys(headers)) { + kw := strings.ToUpper(key) + if err := fitsCheckKeyword(kw); err != nil { + return nil, err + } + card, err := fitsStringCard(kw, headers[key]) + if err != nil { + return nil, err + } + cards = append(cards, card) + } + return cards, nil +} + +// fitsCard is one parsed header card: its keyword and the value field +// with quoting resolved. +type fitsCard struct { + key, value string +} + +// scanFITSCards walks the 80-byte cards of one header starting at off +// and returns every value card up to the END terminator together with +// the block-aligned header length in bytes, which is where the data +// block begins. Blank, COMMENT and HISTORY cards carry no value and +// are skipped before the value-indicator check, because a commentary +// card may legitimately carry an "= " sequence in columns 9-10; a +// header that never terminates is an error. +func scanFITSCards(data []byte, off int) ([]fitsCard, int, error) { + var cards []fitsCard + seen := 0 + for pos := off; pos+80 <= len(data); pos += 80 { + line := string(data[pos : pos+80]) + seen++ + key := strings.TrimRight(line[:8], " ") + if key == "END" { + return cards, (seen + fitsCardsPerBlock - 1) / fitsCardsPerBlock * 2880, nil + } + // Commentary cards have no value regardless of what follows + // columns 9-10; check them before the "= " indicator. + if key == "" || key == "COMMENT" || key == "HISTORY" { + continue + } + if key == "CONTINUE" { + // A CONTINUE card only makes sense after an open long + // string; reaching one here means the base card never + // ended in '&', and silently dropping it would lose the + // caller's value. + return nil, 0, base.Errf("CONTINUE card without an open long string") + } + if line[8:10] != "= " { + continue + } + value, wasString, err := fitsValue(line) + if err != nil { + return nil, 0, err + } + // Long-string convention: a string value ending in '&' is + // continued by the following CONTINUE cards, each carrying the + // next quoted segment. Dropping them would truncate the value + // at the card boundary. + for wasString && strings.HasSuffix(value, "&") && pos+160 <= len(data) { + next := string(data[pos+80 : pos+160]) + if strings.TrimSpace(next[:8]) != "CONTINUE" { + break + } + pos += 80 + seen++ + cont, cerr := fitsContinueString(next) + if cerr != nil { + return nil, 0, cerr + } + value = strings.TrimSuffix(value, "&") + strings.TrimRight(cont, " ") + } + cards = append(cards, fitsCard{key, value}) + } + return nil, 0, base.Errf("no END card terminates the header") +} + +// fitsBlockSize returns n rounded up to a whole FITS block. +func fitsBlockSize(n int) int { + return (n + 2879) / 2880 * 2880 +} + +// fitsAppendCards appends the rendered cards as one header block, +// blank-padded to the block boundary. +func fitsAppendCards(dst []byte, cards []string) []byte { + for _, card := range cards { + dst = append(dst, card...) + } + // Blank cards pad the header to the block boundary. + if rem := len(dst) % 2880; rem != 0 { + blank := fitsPadCard("") + for i := 0; i < 2880-rem; i += 80 { + dst = append(dst, blank...) + } + } + return dst +} + +// fitsAppendZeroPad appends zero bytes up to the block boundary. +func fitsAppendZeroPad(dst []byte) []byte { + if rem := len(dst) % 2880; rem != 0 { + dst = append(dst, make([]byte, 2880-rem)...) + } + return dst +} diff --git a/io/fits_hostile_pins_test.go b/io/fits_hostile_pins_test.go new file mode 100644 index 0000000..ae6d478 --- /dev/null +++ b/io/fits_hostile_pins_test.go @@ -0,0 +1,272 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "strings" + "testing" +) + +// Regression pins for hostile FITS input: negative and missing axis +// cards, TBCOL and row-byte arithmetic that wraps, repeat counts past +// the int range, unbounded HDU skips and garbage PCOUNT, all refused by +// name before any payload is read. + +// tableHostile builds a FITS table file from card bodies plus one +// 2880-byte data block of the given row content. +func tableHostile(cards []string, rows int) []byte { + var b []byte + for _, body := range cards { + b = append(b, card(body)...) + } + b = append(b, card("END")...) + if pad := len(b) % 2880; pad != 0 { + b = append(b, make([]byte, 2880-pad)...) + } + b = append(b, make([]byte, 2880)...) + _ = rows + return b +} + +// TestFITSNegativeAxisCard pins the refusal of a negative or +// missing NAXIS before the axis slice is allocated, and the image skip +// past a zero axis in the table reader. +func TestFITSNegativeAxisCard(t *testing.T) { + build := func(bodies []string) []byte { + var b []byte + for _, body := range bodies { + b = append(b, card(body)...) + } + b = append(b, card("END")...) + if pad := len(b) % 2880; pad != 0 { + b = append(b, make([]byte, 2880-pad)...) + } + return b + } + noNaxis := writeHostile(t, "no-naxis.fits", build([]string{ + "SIMPLE = T", + "BITPIX = -64", + })) + if _, _, err := LoadFITS(noNaxis); err == nil { + t.Fatal("expected an error for a missing NAXIS card") + } + negative := writeHostile(t, "neg-naxis.fits", build([]string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = -3", + })) + if _, _, err := LoadFITS(negative); err == nil { + t.Fatal("expected an error for a negative NAXIS") + } + // A zero axis followed by another divides the running product in + // the table reader's image skip; the skip must pass it without + // dividing by the zeroed product and report that no table follows. + zeroAxis := writeHostile(t, "zero-axis.fits", build([]string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 2", + "NAXIS1 = 0", + "NAXIS2 = 1", + })) + if _, err := LoadFITSTable(zeroAxis); err == nil || !strings.Contains(err.Error(), "no table extension") { + t.Fatalf("LoadFITSTable past a zero axis: err = %v, want the no-table report", err) + } + if _, _, err := LoadFITS(zeroAxis); err == nil { + t.Fatal("LoadFITS: expected an error for a zero axis under positive NAXIS") + } +} + +// TestFITSHugeNaxisRefused pins that a NAXIS beyond the NAXISn cards +// the file carries is refused before the axis slice is allocated: a +// NAXIS of MaxInt64 used to reach make([]int, naxis) and panic with +// makeslice instead of answering the card-count error. +func TestFITSHugeNaxisRefused(t *testing.T) { + build := func(bodies []string) []byte { + var b []byte + for _, body := range bodies { + b = append(b, card(body)...) + } + b = append(b, card("END")...) + if pad := len(b) % 2880; pad != 0 { + b = append(b, make([]byte, 2880-pad)...) + } + return b + } + huge := writeHostile(t, "huge-naxis.fits", build([]string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 9223372036854775807", + })) + if _, _, err := LoadFITS(huge); err == nil || !strings.Contains(err.Error(), "NAXISn cards") { + t.Fatalf("LoadFITS with NAXIS = MaxInt64: err = %v, want the card-count refusal", err) + } +} + +// TestASCIITableTBCOLWrap pins the overflow-free bound on +// TBCOL + width: a TBCOL near MaxInt64 used to wrap the sum negative, +// pass the guard and panic on the slice. +func TestASCIITableTBCOLWrap(t *testing.T) { + data := tableHostile([]string{ + "XTENSION= 'TABLE '", + "BITPIX = 8", + "NAXIS = 2", + "NAXIS1 = 20", + "NAXIS2 = 1", + "TFIELDS = 1", + "TFORM1 = 'I10 '", + "TBCOL1 = '9223372036854775807'", + }, 1) + path := writeHostile(t, "tbcol-wrap.fits", data) + if _, err := LoadFITSTable(path); err == nil { + t.Fatal("expected an error for a TBCOL whose sum with the width wraps") + } +} + +// TestBINTABLERowBytesWrap pins the per-column bound on the +// row prefix sum: two A columns of 2^62 repeat each used to wrap +// rowBytes negative, skip the truncation guard and read past the file. +func TestBINTABLERowBytesWrap(t *testing.T) { + data := tableHostile([]string{ + "XTENSION= 'BINTABLE'", + "BITPIX = 8", + "NAXIS = 2", + "NAXIS1 = 16", + "NAXIS2 = 1", + "TFIELDS = 2", + "TFORM1 = '4611686018427387904A'", + "TFORM2 = '4611686018427387904A'", + }, 1) + path := writeHostile(t, "rowbytes-wrap.fits", data) + if _, err := LoadFITSTable(path); err == nil { + t.Fatal("expected an error for column widths whose sum wraps") + } +} + +// TestTFORMRepeatOverflowRefused pins the named refusal for a +// repeat count that does not fit an int, which used to fall back to a +// silent scalar of width 1. +func TestTFORMRepeatOverflowRefused(t *testing.T) { + if _, _, err := parseTFORM("99999999999999999999E"); err == nil { + t.Fatal("parseTFORM: expected an error for an overflowing repeat") + } + if r, code, err := parseTFORM("16A"); err != nil || r != 16 || code != "A" { + t.Fatalf("parseTFORM(16A) = %d, %q, %v", r, code, err) + } +} + +// TestImageHDUSkipBounded pins the image-HDU skip arithmetic: +// axis products that would wrap the int cannot skip into attacker +// bytes; the file is refused as truncated instead. +func TestImageHDUSkipBounded(t *testing.T) { + var b []byte + for _, body := range []string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 2", + "NAXIS1 = 2147483647", + "NAXIS2 = 4", + "PCOUNT = 0", + "GCOUNT = 2", + } { + b = append(b, card(body)...) + } + b = append(b, card("END")...) + if pad := len(b) % 2880; pad != 0 { + b = append(b, make([]byte, 2880-pad)...) + } + b = append(b, make([]byte, 2880)...) + path := writeHostile(t, "skip-wrap.fits", b) + // The table reader skips the image HDU to look for a table after + // it; the wrapped product used to land the skip inside the image + // data. Now the skip is bounded and the file has no table. + if _, err := LoadFITSTable(path); err == nil { + t.Fatal("expected an error: no table HDU follows the bounded skip") + } +} + +// TestPCOUNTGarbageRefused pins that a non-integer PCOUNT is +// an error, not a silent zero. +func TestPCOUNTGarbageRefused(t *testing.T) { + var b []byte + for _, body := range []string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 2", + "NAXIS1 = 2", + "NAXIS2 = 2", + "PCOUNT = '1e3'", + } { + b = append(b, card(body)...) + } + b = append(b, card("END")...) + if pad := len(b) % 2880; pad != 0 { + b = append(b, make([]byte, 2880-pad)...) + } + b = append(b, make([]byte, 2880)...) + path := writeHostile(t, "pcount.fits", b) + if _, err := LoadFITSTable(path); err == nil { + t.Fatal("expected an error for a non-integer PCOUNT") + } +} + +// TestFITSNAXISnOrder pins that NAXISn cards are bound by +// their axis number: a header writing NAXIS2 before NAXIS1 under +// NAXIS = 2 with matching extents stays the image it declares, and a +// missing axis is an error. +func TestFITSNAXISnOrder(t *testing.T) { + build := func(bodies []string) []byte { + var b []byte + for _, body := range bodies { + b = append(b, card(body)...) + } + b = append(b, card("END")...) + if pad := len(b) % 2880; pad != 0 { + b = append(b, make([]byte, 2880-pad)...) + } + // 2x2 float64 image data. + for range 4 { + b = binary.BigEndian.AppendUint64(b, 1) + } + return b + } + // Reversed card order: the image is still 3 by 2. + path := writeHostile(t, "order.fits", build([]string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 2", + "NAXIS2 = 2", + "NAXIS1 = 3", + })) + // 3x2 = 6 doubles but only 4 present: truncated either way, so the + // claim is checked by the value the reader reports. Give it enough + // data instead: rebuild with full payload. + full := build([]string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 2", + "NAXIS2 = 2", + "NAXIS1 = 3", + }) + full = append(full, make([]byte, 2880)...) + path = writeHostile(t, "order-full.fits", full) + a, _, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS with reversed NAXISn cards: %v", err) + } + if a.Shape()[0] != 2 || a.Shape()[1] != 3 { + t.Fatalf("shape = %v, want [2 3] (Fortran order, NAXIS1 = 3)", a.Shape()) + } + // A missing axis is refused by name. + missing := build([]string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 2", + "NAXIS1 = 2", + }) + mpath := writeHostile(t, "missing.fits", missing) + if _, _, err := LoadFITS(mpath); err == nil || !strings.Contains(err.Error(), "NAXIS") { + t.Fatalf("err = %v, want the missing NAXISn card named", err) + } +} diff --git a/io/fits_table_pins_test.go b/io/fits_table_pins_test.go new file mode 100644 index 0000000..f0d49a0 --- /dev/null +++ b/io/fits_table_pins_test.go @@ -0,0 +1,206 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "fmt" + "math" + "path/filepath" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression pins for the FITS table codecs: the binary and ASCII +// column round trips, the header cards that size them, and the hostile +// counts and widths the reader must refuse by name. + +// card renders one 80-byte FITS header card from a full-line body. +func card(body string) []byte { + return []byte(body + strings.Repeat(" ", 80-len(body))) +} + +// cardBlock concatenates cards and pads to the FITS block size. +func cardBlock(cards ...[]byte) []byte { + var b []byte + for _, c := range cards { + b = append(b, c...) + } + if rem := len(b) % 2880; rem != 0 { + b = append(b, make([]byte, 2880-rem)...) + } + return b +} + +// TestBinaryStringColumnRoundTrip pins the nA binary column contract: +// the reader must accept the repeat-count width the writer emits, so +// SaveFITSTable to LoadFITSTable survives character columns. +func TestBinaryStringColumnRoundTrip(t *testing.T) { + names := []string{"M31", "NGC 1275", "SMC"} + cols := []FITSTableColumn{ + {Name: "NAME", Form: "8A", Text: names}, + {Name: "FLUX", Form: "D", Data: mustFloats(t, []float64{1, 2, 3}, 3)}, + } + path := filepath.Join(t.TempDir(), "names.fits") + if err := SaveFITSTable(path, false, cols, nil); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + for i, want := range names { + if table.Text[0][i] != want { + t.Fatalf("NAME[%d] = %q, want %q", i, table.Text[0][i], want) + } + } +} + +// TestTableTSCALTZERO pins the scaling keywords: a binary int column +// with TSCAL/TZERO must read back through the affine map, the standard +// unsigned-integer convention included. +func TestTableTSCALTZERO(t *testing.T) { + raw := []int64{1, 2, 3} + ints := core.New(core.Int, len(raw)) + copy(ints.RawInts(), raw) + cols := []FITSTableColumn{ + {Name: "U16", Form: "K", Data: ints}, + } + path := filepath.Join(t.TempDir(), "scaled.fits") + headers := map[string]string{"TSCAL1": "1", "TZERO1": "32768"} + if err := SaveFITSTable(path, false, cols, headers); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + col := table.Columns[0] + for i, v := range raw { + want := float64(v) + 32768 + got := col.FloatAt(i) + if got != want { + t.Fatalf("U16[%d] = %g, want %g", i, got, want) + } + } + // A fractional scale promotes the column to float. + if err := SaveFITSTable(path, false, cols, map[string]string{"TSCAL1": "0.5", "TZERO1": "0"}); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err = LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + if table.Columns[0].Dtype() != core.Float { + t.Fatalf("scaled int column dtype %s, want float", table.Columns[0].Dtype()) + } +} + +// TestLoadFITSTableSkips16BitImage pins the image-HDU skip arithmetic +// on the case that used to break: a positive BITPIX (16) with a +// non-square shape, where |BITPIX|/8 · Π NAXISn ≠ Σ NAXISn. +func TestLoadFITSTableSkips16BitImage(t *testing.T) { + // A 5×3 16-bit image (NAXIS1 = 5 is the fastest axis): 15 pixels, + // 30 bytes of data. + hdr := cardBlock( + card("SIMPLE = T"), + card("BITPIX = 16"), + card("NAXIS = 2"), + card("NAXIS1 = 5"), + card("NAXIS2 = 3"), + card("END"), + ) + data := make([]byte, 0, 30) + for range 15 { + data = binary.BigEndian.AppendUint16(data, 42) + } + image := append(hdr, data...) + if rem := len(image) % 2880; rem != 0 { + image = append(image, make([]byte, 2880-rem)...) + } + + cols := []FITSTableColumn{ + {Name: "V", Form: "K", Data: mustInts(t, []int64{7, 8}, 2)}, + } + tablePath := filepath.Join(t.TempDir(), "table.fits") + if err := SaveFITSTable(tablePath, false, cols, nil); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + tableBytes, err := osReadFile(tablePath) + if err != nil { + t.Fatalf("read table: %v", err) + } + // Drop the table file's empty primary HDU (first block), keep the + // extension behind the hand-built image. + path := filepath.Join(t.TempDir(), "im_then_table.fits") + if err := osWriteFile(path, image, tableBytes[2880:]); err != nil { + t.Fatalf("write: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable after a 16-bit image: %v", err) + } + if table.Rows != 2 || table.Columns[0].RawInts()[1] != 8 { + t.Fatalf("table landed wrong: rows %d, V[1] = %d", table.Rows, table.Columns[0].RawInts()[1]) + } +} + +// TestLoadFITSBSCALE pins image scaling: physical = raw·BSCALE + BZERO. +func TestLoadFITSBSCALE(t *testing.T) { + hdr := cardBlock( + card("SIMPLE = T"), + card("BITPIX = -64"), + card("NAXIS = 1"), + card("NAXIS1 = 2"), + card("BSCALE = 2.0"), + card("BZERO = 10.0"), + card("END"), + ) + var data []byte + data = binary.BigEndian.AppendUint64(data, math.Float64bits(1)) + data = binary.BigEndian.AppendUint64(data, math.Float64bits(2)) + path := filepath.Join(t.TempDir(), "scaled.fits") + var payload []byte + payload = append(payload, data...) + if rem := len(payload) % 2880; rem != 0 { + payload = append(payload, make([]byte, 2880-rem)...) + } + if err := osWriteFile(path, hdr, payload); err != nil { + t.Fatalf("write: %v", err) + } + img, _, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS: %v", err) + } + if img.FloatAt(0) != 12 || img.FloatAt(1) != 14 { + t.Fatalf("scaled pixels (%g, %g), want (12, 14)", img.FloatAt(0), img.FloatAt(1)) + } +} + +// TestFITSContinueLongString pins the long-string convention: a value +// ending in '&' continues on the CONTINUE cards instead of truncating. +func TestFITSContinueLongString(t *testing.T) { + hdr := cardBlock( + card("SIMPLE = T"), + card("BITPIX = 8"), + card("NAXIS = 0"), + card(fmt.Sprintf("%-8s= '%s&'", "LONGKEY", "the first part of a long value ")), + card(fmt.Sprintf("CONTINUE '%s'", "and the second part")), + card("END"), + ) + path := filepath.Join(t.TempDir(), "long.fits") + if err := osWriteFile(path, hdr); err != nil { + t.Fatalf("write: %v", err) + } + _, headers, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS: %v", err) + } + want := "the first part of a long value and the second part" + if headers["LONGKEY"] != want { + t.Fatalf("LONGKEY = %q, want %q", headers["LONGKEY"], want) + } +} diff --git a/io/fits_test.go b/io/fits_test.go new file mode 100644 index 0000000..985d909 --- /dev/null +++ b/io/fits_test.go @@ -0,0 +1,338 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "math" + "os" + "path/filepath" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +func fitsTempPath(t *testing.T) string { + t.Helper() + return filepath.Join(t.TempDir(), "image.fits") +} + +// padCard right-pads a card body with spaces to the full 80 bytes. +func padCard(text string) []byte { + return []byte(text + strings.Repeat(" ", 80-len(text))) +} + +// TestFITSRoundTripFloat64 moves a rank-2 float64 image with negative +// and non-round values through the file format and back, values and +// header strings included. +func TestFITSRoundTripFloat64(t *testing.T) { + a := mustFloats(t, []float64{ + 1.5, -2.25, 3.125, 4, + -5.5, 6.75, -7.875, 8, + 9.25, -10.5, 11.125, -12, + }, 3, 4) + path := fitsTempPath(t) + headers := map[string]string{ + "OBJECT": "M31", + "OBSERVER": "petr's dome", + "EXPTIME": "600", + } + if err := SaveFITS(path, a, headers); err != nil { + t.Fatalf("SaveFITS: %v", err) + } + back, hdr, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS: %v", err) + } + if back.Dtype() != core.Float || back.NDim() != 2 || back.Shape()[0] != 3 || back.Shape()[1] != 4 { + t.Fatalf("shape/dtype mismatch: %v %s", back.Shape(), back.Dtype()) + } + for i := range a.Len() { + if back.FloatAt(i) != a.FloatAt(i) { + t.Fatalf("value[%d] = %g, want %g", i, back.FloatAt(i), a.FloatAt(i)) + } + } + for key, want := range map[string]string{ + "OBJECT": "M31", + "OBSERVER": "petr's dome", + "EXPTIME": "600", + } { + if hdr[key] != want { + t.Fatalf("header %q = %q, want %q", key, hdr[key], want) + } + } +} + +// TestFITSAxisConvention pins the wire format against the raw bytes: +// NAXIS1 must carry the fastest (last Go) axis and the payload must +// be big-endian in flat row-major order. +func TestFITSAxisConvention(t *testing.T) { + a := mustFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + path := fitsTempPath(t) + if err := SaveFITS(path, a, nil); err != nil { + t.Fatalf("SaveFITS: %v", err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if len(raw)%2880 != 0 { + t.Fatalf("file length %d is not a multiple of 2880", len(raw)) + } + header := raw[:2880] + for _, want := range []string{ + "SIMPLE = T", + "BITPIX = -64", + "NAXIS = 2", + "NAXIS1 = 3", + "NAXIS2 = 2", + } { + if !strings.Contains(string(header), want) { + t.Fatalf("header misses %q", want) + } + } + endAt := strings.Index(string(header), "END") + if endAt < 0 { + t.Fatal("header has no END card") + } + var word [8]byte + for i := range 6 { + binary.BigEndian.PutUint64(word[:], math.Float64bits(float64(i+1))) + got := raw[2880+i*8 : 2880+i*8+8] + if string(got) != string(word[:]) { + t.Fatalf("payload word %d = % x, want % x", i, got, word) + } + } +} + +// TestFITSRoundTripFloat32 keeps the float32 element type through the +// round trip, which is what BITPIX −32 stores. +func TestFITSRoundTripFloat32(t *testing.T) { + a := core.New(core.Float32, 5) + for i := range 5 { + a.RawFloat32s()[i] = float32(i) * 1.25 + } + path := fitsTempPath(t) + if err := SaveFITS(path, a, nil); err != nil { + t.Fatalf("SaveFITS: %v", err) + } + back, _, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS: %v", err) + } + if back.Dtype() != core.Float32 { + t.Fatalf("dtype = %s, want core.Float32", back.Dtype()) + } + for i := range 5 { + if back.RawFloat32s()[i] != a.RawFloat32s()[i] { + t.Fatalf("value[%d] = %g, want %g", i, back.RawFloat32s()[i], a.RawFloat32s()[i]) + } + } +} + +// TestFITSSkipsValuelessCards checks COMMENT and HISTORY cards are +// tolerated and skipped rather than parsed as values. +func TestFITSSkipsValuelessCards(t *testing.T) { + a := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2) + path := fitsTempPath(t) + if err := SaveFITS(path, a, map[string]string{"OBJECT": "TEST"}); err != nil { + t.Fatalf("SaveFITS: %v", err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + // Splice two valueless cards in before the END card and re-pad the + // header to the block boundary. The END card is matched in full + // ("END" plus its padding) because EXTEND contains the same three + // letters. + endCard := strings.Index(string(raw), "END"+strings.Repeat(" ", 77)) + if endCard < 0 { + t.Fatal("no END card") + } + spliced := append([]byte{}, raw[:endCard]...) + spliced = append(spliced, padCard("COMMENT a note without a value")...) + spliced = append(spliced, padCard("HISTORY an audit trail entry")...) + spliced = append(spliced, padCard("END")...) + for len(spliced)%2880 != 0 { + spliced = append(spliced, padCard("")...) + } + spliced = append(spliced, raw[2880:]...) + if err := os.WriteFile(path, spliced, 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + back, hdr, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS: %v", err) + } + if back.Len() != 4 || back.FloatAt(3) != 4 { + t.Fatalf("payload damaged: %v %v", back.Shape(), back.RawFloats()) + } + if hdr["OBJECT"] != "TEST" { + t.Fatalf("OBJECT = %q, want %q", hdr["OBJECT"], "TEST") + } + if _, ok := hdr["COMMENT"]; ok { + t.Fatal("COMMENT card must not enter the header map") + } +} + +// TestFITSErrors covers every refusal path in the format layer. +func TestFITSErrors(t *testing.T) { + path := fitsTempPath(t) + + ints, ierr := core.FromInts([]int64{1, 2, 3, 4}, 2, 2) + if ierr != nil { + t.Fatalf("FromInts: %v", ierr) + } + if err := SaveFITS(path, ints, nil); err == nil { + t.Fatal("int64 input: want an error") + } + if err := SaveFITS(path, mustComplexes(t, []complex128{1, 2, 3, 4}, 2, 2), nil); err == nil { + t.Fatal("complex input: want an error") + } + ok := mustFloats(t, []float64{1, 2}, 2) + if err := SaveFITS(path, ok, map[string]string{"TOOLONGKEYWORD": "x"}); err == nil { + t.Fatal("long keyword: want an error") + } + if err := SaveFITS(path, ok, map[string]string{"BITPIX": "x"}); err == nil { + t.Fatal("reserved keyword: want an error") + } + if err := SaveFITS(path, ok, map[string]string{"naxis2": "x"}); err == nil { + t.Fatal("NAXISn keyword: want an error") + } + if err := SaveFITS(path, ok, map[string]string{"BAD KEY": "x"}); err == nil { + t.Fatal("space in keyword: want an error") + } + if err := SaveFITS(path, ok, map[string]string{"NOTE": strings.Repeat("x", 69)}); err == nil { + t.Fatal("overlong value: want an error") + } + + // A valid file to damage in every way the parser must catch. + if err := SaveFITS(path, ok, nil); err != nil { + t.Fatalf("SaveFITS: %v", err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + with := func(mutate func([]byte) []byte) { + t.Helper() + if err := os.WriteFile(path, mutate(append([]byte{}, raw...)), 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if _, _, err := LoadFITS(path); err == nil { + t.Fatal("damaged file: want an error") + } + } + with(func(b []byte) []byte { return b[:100] }) // no END card + with(func(b []byte) []byte { return b[:2880] }) // data truncated + with(func(b []byte) []byte { b[30] = 'F'; return b }) // SIMPLE = F + with(func(b []byte) []byte { copy(b[11:20], "XTENSION"); return b }) // wrong first card + + // BITPIX 8 (unsigned bytes) is outside the supported image types. + unsupported := append(append(append(append([]byte{}, + padCard("SIMPLE = T")...), + padCard("BITPIX = 8")...), + padCard("NAXIS = 1")...), + padCard("NAXIS1 = 4")...) + unsupported = append(unsupported, padCard("END")...) + unsupported = append(unsupported, make([]byte, 2880-5*80)...) + unsupported = append(unsupported, make([]byte, 2880)...) + if err := os.WriteFile(path, unsupported, 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if _, _, err := LoadFITS(path); err == nil { + t.Fatal("BITPIX 8: want an error") + } + + // A hostile header whose axis product overflows int must be + // refused as truncated data, never panic the process. + hostile := append(append(append(append([]byte{}, + padCard("SIMPLE = T")...), + padCard("BITPIX = -64")...), + padCard("NAXIS = 2")...), + padCard("NAXIS1 = 1099511627776")...) + hostile = append(hostile, padCard("NAXIS2 = 1099511627776")...) + hostile = append(hostile, padCard("END")...) + hostile = append(hostile, make([]byte, 2880-6*80)...) + hostile = append(hostile, make([]byte, 2880)...) + if err := os.WriteFile(path, hostile, 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if _, _, err := LoadFITS(path); err == nil { + t.Fatal("overflowing axis product: want an error") + } + + // An XTENSION-first file is an extension, not a primary image. + ext := append([]byte{}, padCard("XTENSION= 'IMAGE '")...) + ext = append(ext, raw[80:]...) + if err := os.WriteFile(path, ext, 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if _, _, err := LoadFITS(path); err == nil { + t.Fatal("extension header: want an error") + } + + // The intact file still loads after all the damage around it. + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if _, _, err := LoadFITS(path); err != nil { + t.Fatalf("intact file: %v", err) + } +} + +// TestCommentaryCardsAreNotValueCards pins the commentary rule: a +// COMMENT or HISTORY card may legitimately carry an "= " sequence in +// columns 9-10, and such a card must be skipped as commentary, not +// parsed as a keyword with a value. +func TestCommentaryCardsAreNotValueCards(t *testing.T) { + cards := []string{ + fitsBoolCard("SIMPLE", true), + fitsIntCard("BITPIX", -64), + fitsIntCard("NAXIS", 1), + fitsIntCard("NAXIS1", 2), + fitsPadCard("COMMENT = this looks like a value card"), + fitsPadCard("HISTORY = so does this one"), + fitsStringCardRaw("OBSERVER", "tester"), + fitsEndCard(), + } + data := fitsAppendCards(nil, cards) + payload := []byte{0x3f, 0xf0, 0, 0, 0, 0, 0, 0, 0x40, 0, 0, 0, 0, 0, 0, 0} // 1.0, 2.0 + data = append(data, payload...) + data = fitsAppendZeroPad(data) + + path := filepath.Join(t.TempDir(), "commentary.fits") + if err := osWriteFile(path, data); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } + img, headers, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS: %v", err) + } + if img.Len() != 2 || img.FloatAt(0) != 1 || img.FloatAt(1) != 2 { + t.Fatalf("image = %s, want [1, 2]", img) + } + if _, ok := headers["COMMENT"]; ok { + t.Error("COMMENT was parsed as a value keyword") + } + if _, ok := headers["HISTORY"]; ok { + t.Error("HISTORY was parsed as a value keyword") + } + if headers["OBSERVER"] != "tester" { + t.Errorf("OBSERVER = %q, want tester", headers["OBSERVER"]) + } +} + +// TestFitsCheckKeywordRefusesCommentary pins that COMMENT and HISTORY +// are refused as user keywords: they carry no value in the format. +func TestFitsCheckKeywordRefusesCommentary(t *testing.T) { + dir := t.TempDir() + img := mustFloats(t, []float64{1, 2}, 2) + for _, kw := range []string{"COMMENT", "HISTORY"} { + if err := SaveFITS(filepath.Join(dir, strings.ToLower(kw)+".fits"), img, map[string]string{kw: "x"}); err == nil { + t.Errorf("SaveFITS accepted the reserved keyword %q", kw) + } + } +} diff --git a/io/fitstable.go b/io/fitstable.go new file mode 100644 index 0000000..99243f3 --- /dev/null +++ b/io/fitstable.go @@ -0,0 +1,1082 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "bytes" + "encoding/binary" + "fmt" + "math" + "os" + "strconv" + "strings" + "unsafe" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// FITS table extensions. Catalogues, source lists and observation +// logs live in the format's table HDUs: a binary table (XTENSION +// 'BINTABLE') packs each column big-endian according to its TFORM +// descriptor, an ASCII table (XTENSION 'TABLE') lays fixed-width text +// columns into rows of NAXIS1 characters. Both share the header +// vocabulary: TTYPEn names a column, TUNITn gives its unit, TFIELDS +// counts them. + +// FITSTableColumn describes one column of a FITS table. Form carries +// the TFORM-style descriptor: "D" (float64), "E" (float32), "K" +// (int64), "L" (logical), "B" (unsigned byte), "I" (16-bit) or "J" +// (32-bit) for a binary table, and "nA" (a character string of n +// characters) for either kind; "Inw", "Fw.d", "Ew.d" or "Dw.d" lay out +// an ASCII table column, where w is the column width in characters. A +// numeric form reads its values from Data (Float answers D, Float32 +// answers E, Int answers K, L, B, I and J); a character form reads +// them from Text. +type FITSTableColumn struct { + Name string + Unit string + Form string + Data *core.Array + Text []string +} + +// FITSTable is a parsed table extension. Names, Units, Columns and +// Text run parallel to the file's column list: a numeric column holds +// its values in Columns and nil in Text, a character column holds its +// strings in Text and nil in Columns. +type FITSTable struct { + Kind string // "BINTABLE" or "TABLE" + Names []string + Units []string + Columns []*core.Array + Text [][]string + Rows int + Headers map[string]string +} + +// SaveFITSTable writes a primary HDU followed by one table extension: +// a binary table by default, an ASCII table when ascii is set. The +// headers land in the extension's header beside the column cards. A +// column with a numeric form but nil Data, a character form without +// Text, mismatched row counts, an unknown form, an empty column list, +// or a value that does not fit its ASCII form's width is an error. +func SaveFITSTable(path string, ascii bool, cols []FITSTableColumn, headers map[string]string) error { + const name = "SaveFITSTable" + if len(cols) == 0 { + return base.Errf("%s: at least one column is required", name) + } + rows := -1 + forms := make([]string, len(cols)) + for i := range cols { + c := &cols[i] + if err := fitsCheckString(c.Name, "column name"); err != nil { + return base.Errf("%s: %w", name, err) + } + form := strings.ToUpper(strings.TrimSpace(c.Form)) + forms[i] = form + if form == "" { + return base.Errf("%s: column %d (%s) has an empty form", name, i, c.Name) + } + if strings.HasSuffix(form, "A") { + if c.Text == nil { + return base.Errf("%s: character column %s needs Text", name, c.Name) + } + if c.Data != nil { + return base.Errf("%s: character column %s must not carry Data", name, c.Name) + } + if rows >= 0 && len(c.Text) != rows { + return base.Errf("%s: column %s has %d rows, want %d", name, c.Name, len(c.Text), rows) + } + rows = len(c.Text) + continue + } + if c.Data == nil { + return base.Errf("%s: numeric column %s needs Data", name, c.Name) + } + if c.Data.NDim() != 1 { + return base.Errf("%s: column %s must be a vector, got shape %s", name, c.Name, base.ShapeText(c.Data.Shape())) + } + if rows >= 0 && c.Data.Len() != rows { + return base.Errf("%s: column %s has %d rows, want %d", name, c.Name, c.Data.Len(), rows) + } + rows = c.Data.Len() + } + if rows == 0 { + return base.Errf("%s: tables need at least one row", name) + } + + // Row image: for each column its encoded byte width (binary) or + // character width (ASCII), and an encoder closure. An encoder + // reports an error when the rendered value does not fit its field. + var rowBytes int + type encoder func(dst []byte, row int) error + var encoders []encoder + var asciiEncoders []encoder + asciiWidths := make([]int, len(cols)) + for i := range cols { + c := &cols[i] + form := forms[i] + if ascii { + width, prec, code, perr := parseASCIISaveForm(form) + if perr != nil { + return base.Errf("%s: column %s: %w", name, c.Name, perr) + } + asciiWidths[i] = width + rowBytes += width + switch code { + case 'A': + text := c.Text + if text == nil { + return base.Errf("%s: character column %s needs Text", name, c.Name) + } + asciiEncoders = append(asciiEncoders, func(dst []byte, row int) error { + if len(text[row]) > width { + return base.Errf("text %q is %d characters, form %s allows %d", text[row], len(text[row]), form, width) + } + copy(dst, text[row]) + for p := len(text[row]); p < width; p++ { + dst[p] = ' ' + } + return nil + }) + case 'I': + col := c.Data + if col == nil || col.Dtype() != core.Int { + return base.Errf("%s: form %s needs an int column", name, form) + } + ints := col.RawInts() + asciiEncoders = append(asciiEncoders, func(dst []byte, row int) error { + word := fmt.Sprintf("%d", ints[row]) + if len(word) > width { + return base.Errf("value %d does not fit the width %d of form %s", ints[row], width, form) + } + copy(dst, word) + for p := len(word); p < width; p++ { + dst[p] = ' ' + } + return nil + }) + case 'E', 'F', 'D': + col := c.Data + if col == nil || (col.Dtype() != core.Float && col.Dtype() != core.Float32) { + return base.Errf("%s: form %s needs a float column", name, form) + } + scientific := form[0] == 'E' || form[0] == 'D' + // Keep every value inside the declared width: shrink + // the printed precision to what the column holds + // rather than letting the field overflow it. + overhead := 2 // sign room and decimal point + if scientific { + overhead = 8 // sign, one digit, point and the e+/-xx tail + } + if width <= overhead { + return base.Errf("%s: form %s is too narrow for any value", name, form) + } + effPrec := min(prec, width-overhead) + // The payload the column's dtype actually carries: a + // float32 array keeps its values in the float32 slice + // alone, so reading the float64 one first indexes a nil + // slice. + isFloat32 := col.Dtype() == core.Float32 + floats := col.RawFloats() + floats32 := col.RawFloat32s() + asciiEncoders = append(asciiEncoders, func(dst []byte, row int) error { + var v float64 + if isFloat32 { + v = float64(floats32[row]) + } else { + v = floats[row] + } + var word string + if scientific { + word = fmt.Sprintf("%*.*e", width, effPrec, v) + } else { + word = fmt.Sprintf("%*.*f", width, effPrec, v) + } + if len(word) > width { + return base.Errf("value %g does not fit the width %d of form %s", v, width, form) + } + copy(dst, word) + for p := len(word); p < width; p++ { + dst[p] = ' ' + } + return nil + }) + } + continue + } + switch { + case strings.HasSuffix(form, "A"): + width, aerr := strconv.Atoi(strings.TrimSuffix(form, "A")) + if aerr != nil || width < 1 { + return base.Errf("%s: character form %q needs a positive repeat", name, c.Form) + } + text := c.Text + if text == nil { + return base.Errf("%s: character column %s needs Text", name, c.Name) + } + for r := range rows { + if len(text[r]) > width { + return base.Errf("%s: row %d of %s is %d characters, form %q allows %d", + name, r, c.Name, len(text[r]), form, width) + } + } + off := rowBytes + encoders = append(encoders, func(dst []byte, row int) error { + copy(dst[off:], text[row]) + for p := len(text[row]); p < width; p++ { + dst[off+p] = ' ' + } + return nil + }) + rowBytes += width + case form == "D": + if c.Data.Dtype() != core.Float { + return base.Errf("%s: form D needs a float64 column, %s is %s", name, c.Name, c.Data.Dtype()) + } + off := rowBytes + floats := c.Data.RawFloats() + encoders = append(encoders, func(dst []byte, row int) error { + binary.BigEndian.PutUint64(dst[off:], math.Float64bits(floats[row])) + return nil + }) + rowBytes += 8 + case form == "E": + if c.Data.Dtype() != core.Float32 { + return base.Errf("%s: form E needs a float32 column, %s is %s", name, c.Name, c.Data.Dtype()) + } + off := rowBytes + floats32 := c.Data.RawFloat32s() + encoders = append(encoders, func(dst []byte, row int) error { + binary.BigEndian.PutUint32(dst[off:], math.Float32bits(floats32[row])) + return nil + }) + rowBytes += 4 + case form == "K" || form == "L": + if c.Data.Dtype() != core.Int { + return base.Errf("%s: form %s needs an int column, %s is %s", name, form, c.Name, c.Data.Dtype()) + } + off := rowBytes + ints := c.Data.RawInts() + if form == "K" { + encoders = append(encoders, func(dst []byte, row int) error { + binary.BigEndian.PutUint64(dst[off:], uint64(ints[row])) + return nil + }) + rowBytes += 8 + } else { + encoders = append(encoders, func(dst []byte, row int) error { + dst[off] = 'F' + if ints[row] != 0 { + dst[off] = 'T' + } + return nil + }) + rowBytes += 1 + } + default: + return base.Errf("%s: unsupported form %q, want D, E, K, L or nA", name, c.Form) + } + } + if ascii { + // ASCII rows separate neighbouring columns by one space. + rowBytes += len(cols) - 1 + } + + // Extension header. + ext := "BINTABLE" + if ascii { + ext = "TABLE" + } + cards := []string{ + fitsStringCardRaw("XTENSION", ext), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 2), + fitsIntCard("NAXIS1", rowBytes), + fitsIntCard("NAXIS2", rows), + fitsIntCard("PCOUNT", 0), + fitsIntCard("GCOUNT", 1), + fitsIntCard("TFIELDS", len(cols)), + } + colStart := 1 + for i := range cols { + c := &cols[i] + n := strconv.Itoa(i + 1) + nameCard, err := fitsStringCard("TTYPE"+n, c.Name) + if err != nil { + return base.Errf("%s: %w", name, err) + } + cards = append(cards, nameCard) + if ascii { + cards = append(cards, fitsIntCard("TBCOL"+n, colStart)) + colStart += asciiWidths[i] + 1 + } + formCard, err := fitsStringCard("TFORM"+n, forms[i]) + if err != nil { + return base.Errf("%s: %w", name, err) + } + cards = append(cards, formCard) + if c.Unit != "" { + unitCard, uerr := fitsStringCard("TUNIT"+n, c.Unit) + if uerr != nil { + return base.Errf("%s: %w", name, uerr) + } + cards = append(cards, unitCard) + } + } + userCards, err := fitsUserCards(headers) + if err != nil { + return base.Errf("%s: %w", name, err) + } + cards = append(cards, userCards...) + cards = append(cards, fitsEndCard()) + + // Primary HDU comes first: a zero-axis image header. + primary := []string{ + fitsBoolCard("SIMPLE", true), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 0), + fitsBoolCard("EXTEND", true), + fitsEndCard(), + } + // Both headers pad to whole blocks and the body is rows*rowBytes, so + // the file size is known before a byte is written. + out := make([]byte, 0, + fitsBlockSize(len(primary))+fitsBlockSize(len(cards))+rows*rowBytes+2880) + out = fitsAppendCards(out, primary) + out = fitsAppendCards(out, cards) + + body := make([]byte, rows*rowBytes) + if ascii { + for r := range rows { + p := 0 + for i := range cols { + if i > 0 { + body[r*rowBytes+p] = ' ' + p++ + } + if err := asciiEncoders[i](body[r*rowBytes+p:], r); err != nil { + return base.Errf("%s: column %s row %d: %w", name, cols[i].Name, r, err) + } + p += asciiWidths[i] + } + } + } else { + for r := range rows { + for _, enc := range encoders { + if err := enc(body[r*rowBytes:], r); err != nil { + return base.Errf("%s: row %d: %w", name, r, err) + } + } + } + } + out = append(out, body...) + out = fitsAppendZeroPad(out) + return os.WriteFile(path, out, 0o644) +} + +// parseASCIISaveForm splits an ASCII-table column form into its +// character width, fractional digits and type code: "10A" gives +// (10, 0, 'A'), "I7" gives (7, 0, 'I'), "F14.6" gives (14, 6, 'F'), +// "E12.4" and "D20.12" give their widths and digits with the +// scientific code. +func parseASCIISaveForm(form string) (int, int, byte, error) { + if form == "" { + return 0, 0, 0, base.Errf("empty form") + } + // "nA" carries its repeat before the code; every other form puts + // the width after it, optionally with a .precision tail. + if form[len(form)-1] == 'A' { + head := form[:len(form)-1] + if head == "" { + return 1, 0, 'A', nil + } + width, cerr := strconv.Atoi(head) + if cerr != nil || width < 1 { + return 0, 0, 0, base.Errf("form %q needs a positive width", form) + } + return width, 0, 'A', nil + } + code := form[0] + switch code { + case 'I', 'F', 'E', 'D': + default: + return 0, 0, 0, base.Errf("form %q has no known type code", form) + } + rest := form[1:] + before, after, ok := strings.Cut(rest, ".") + num := rest + if ok { + num = before + } + width, cerr := strconv.Atoi(num) + if cerr != nil || width < 1 { + return 0, 0, 0, base.Errf("form %q needs a positive width", form) + } + prec := 0 + if ok { + p, perr := strconv.Atoi(after) + if perr != nil { + return 0, 0, 0, base.Errf("form %q has a malformed precision", form) + } + prec = p + } + return width, prec, code, nil +} + +// fitsStringCardRaw renders a string card without the reserved-keyword +// check (XTENSION is reserved but the table writes it itself). +func fitsStringCardRaw(keyword, v string) string { + card, _ := fitsStringCard(keyword, v) + return card +} + +// LoadFITSTable reads the first table extension (XTENSION 'BINTABLE' +// or 'TABLE') of a FITS file, skipping the primary HDU and any images +// before it. The result carries every column: numeric ones as arrays, +// character ones as string slices. +func LoadFITSTable(path string) (*FITSTable, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, base.Errf("LoadFITSTable: %w", err) + } + off := 0 + for { + table, next, terr := parseFITSTableHDU(data, off) + if terr != nil { + return nil, base.Errf("LoadFITSTable: %w", terr) + } + if table != nil { + return table, nil + } + if next <= off { + return nil, base.Errf("LoadFITSTable: the file stalled at byte %d", off) + } + off = next + if off >= len(data) { + return nil, base.Errf("LoadFITSTable: no table extension found") + } + } +} + +// parseFITSTableHDU parses one HDU starting at off. A nil table with +// a positive next means "an image HDU, carry on"; a table comes back +// fully populated. +func parseFITSTableHDU(data []byte, off int) (*FITSTable, int, error) { + if off+80 > len(data) { + return nil, 0, base.Errf("file ends inside a header at byte %d", off) + } + keyword := strings.TrimRight(string(data[off:off+8]), " ") + isTable := keyword == "XTENSION" + cards, headerLen, err := scanFITSCards(data, off) + if err != nil { + return nil, 0, err + } + dataAt := off + headerLen + // The keyword map answers lookups without the per-lookup scan the + // linear walk cost. A repeated keyword keeps its first value, which + // is what the scan it replaces returned, so a hostile header that + // repeats a structural card is read exactly as before. + first := make(map[string]string, len(cards)) + for _, c := range cards { + if _, ok := first[c.key]; !ok { + first[c.key] = c.value + } + } + get := func(key string) string { return first[key] } + atoi := func(key string) (int, error) { + v := get(key) + n, cerr := strconv.Atoi(v) + if cerr != nil { + return 0, base.Errf("%s = %q is not an integer", key, v) + } + return n, nil + } + // atoiN is the numbered-keyword form: "NAXIS1" is prefix NAXIS and + // index 1, formatted into a stack buffer, with the same error text + // the concatenated key produced. + atoiN := func(prefix string, n int) (int, error) { + v := fitsLookupN(first, prefix, n) + parsed, cerr := strconv.Atoi(v) + if cerr != nil { + return 0, base.Errf("%s%d = %q is not an integer", prefix, n, v) + } + return parsed, nil + } + bitpix, berr := atoi("BITPIX") + if berr != nil { + return nil, 0, berr + } + naxis, aerr := atoi("NAXIS") + if aerr != nil { + return nil, 0, aerr + } + var dims []int + for j := 1; j <= naxis; j++ { + d, derr := atoiN("NAXIS", j) + if derr != nil { + return nil, 0, derr + } + dims = append(dims, d) + } + pcount := 0 + if v := get("PCOUNT"); v != "" { + p, perr := atoi("PCOUNT") + if perr != nil { + return nil, 0, perr + } + pcount = p + } + gcount := 1 + if v := get("GCOUNT"); v != "" { + g, gerr := atoi("GCOUNT") + if gerr != nil { + return nil, 0, gerr + } + gcount = g + } + if !isTable { + // An image HDU: skip its data block. The size is |BITPIX|/8 + // bytes per element times GCOUNT*(PCOUNT + every NAXISn); positive + // BITPIX (the integer depths) counts the same as negative, and + // the axis extents multiply rather than add. Every factor is + // bounded against the bytes the file actually holds before it + // multiplies, the same discipline parseFITS applies, so a + // hostile product cannot wrap into a skip to a chosen offset. + width := 1 + if bitpix < 0 { + width = -bitpix / 8 + } else { + width = bitpix / 8 + } + if width <= 0 { + return nil, 0, base.Errf("BITPIX = %d gives a non-positive element width", bitpix) + } + avail := max(int64(len(data)-dataAt), 0) + total := 0 + if naxis > 0 { + prod := int64(1) + for _, d := range dims { + if d < 0 { + return nil, 0, base.Errf("image HDU declares the negative axis %d", d) + } + // A zero axis empties the data block whatever follows: + // the bound divides by the running product, so the + // division runs only while the product is live. + if prod > 0 && int64(d) > avail/int64(width)/prod { + return nil, 0, base.Errf("image HDU declares more data than the %d bytes the file holds", len(data)-dataAt) + } + prod *= int64(d) + } + if int64(pcount) < 0 || int64(pcount) > avail { + return nil, 0, base.Errf("image HDU declares a PCOUNT of %d against %d bytes", pcount, avail) + } + space := int64(pcount) + prod + switch { + case space == 0: + // A zero-sized group repeats into nothing whatever + // GCOUNT claims. + total = 0 + case int64(gcount) > 1 && int64(gcount) > avail/space: + return nil, 0, base.Errf("image HDU declares more data than the %d bytes the file holds", len(data)-dataAt) + default: + total = int(int64(gcount) * space) + } + } + if total < 0 { + return nil, 0, base.Errf("image HDU declares a negative data size") + } + return nil, dataAt + fitsBlockSize(total*width), nil + } + ext := strings.Trim(get("XTENSION"), " ") + if ext != "BINTABLE" && ext != "TABLE" { + // An unknown extension: refuse rather than guess its size. + return nil, 0, base.Errf("unsupported extension %q", ext) + } + if naxis != 2 || len(dims) != 2 { + return nil, 0, base.Errf("%s extension must have NAXIS = 2, got %d", ext, naxis) + } + tfields, ferr := atoi("TFIELDS") + if ferr != nil { + return nil, 0, ferr + } + // A negative count would size the per-column slices below with a + // negative length, which panics; it is a header that lies, so it is + // refused by name. + if tfields < 0 { + return nil, 0, base.Errf("TFIELDS = %d is negative", tfields) + } + rows := dims[1] + if rows < 1 { + return nil, 0, base.Errf("table declares NAXIS2 = %d, want at least 1 row", rows) + } + // The count is bounded by what could back it, the way the NetCDF + // reader bounds its own: every column of either table kind occupies + // at least one byte of every row, so a TFIELDS past the data the + // file holds is a claim nothing can follow, and refusing it here + // keeps the per-column preallocations from being sized by the claim. + if left := int64(len(data) - dataAt); int64(tfields) > left/int64(rows) { + return nil, 0, base.Errf("TFIELDS = %d names more columns than the %d bytes of table data can hold", tfields, left) + } + table := &FITSTable{ + Kind: ext, + Rows: rows, + Headers: map[string]string{}, + } + type parsed struct { + name, unit, form string + tbcol int + } + var cols []parsed + for i := 1; i <= tfields; i++ { + col := parsed{ + name: fitsLookupN(first, "TTYPE", i), + unit: fitsLookupN(first, "TUNIT", i), + form: strings.ToUpper(strings.TrimSpace(fitsLookupN(first, "TFORM", i))), + tbcol: 0, + } + if v := fitsLookupN(first, "TBCOL", i); v != "" { + tb, tberr := strconv.Atoi(strings.TrimSpace(v)) + if tberr != nil { + return nil, 0, base.Errf("column %d: TBCOL %q is not an integer", i, v) + } + col.tbcol = tb + } + if col.form == "" { + return nil, 0, base.Errf("column %d has no TFORM", i) + } + cols = append(cols, col) + table.Names = append(table.Names, col.name) + table.Units = append(table.Units, col.unit) + } + headers := table.Headers + for _, c := range cards { + switch c.key { + case "XTENSION", "BITPIX", "NAXIS", "PCOUNT", "GCOUNT", "TFIELDS": + default: + if !fitsAxisKeyword(c.key) && + !strings.HasPrefix(c.key, "TTYPE") && + !strings.HasPrefix(c.key, "TFORM") && + !strings.HasPrefix(c.key, "TUNIT") && + !strings.HasPrefix(c.key, "TBCOL") { + headers[c.key] = c.value + } + } + } + if ext == "BINTABLE" { + // The data block starts on a 2880-byte block boundary, so a file + // that ends inside the header has no data at all. The row + // arithmetic below indexes from dataAt, which must be checked + // against the file before it is used. + if dataAt > len(data) { + return nil, 0, base.Errf("table data starts at byte %d of a %d-byte file", dataAt, len(data)) + } + widths := make([]int, tfields) + forms := make([]string, tfields) + // Column offsets inside a row: prefix sums of the widths, + // computed once instead of once per row. + offs := make([]int, tfields) + rowBytes := 0 + // The widest row any present data can hold: every width is + // bounded against it as it accumulates, so the prefix sum can + // never wrap past the guard below into a small positive value. + maxRow := (int64(len(data)) - int64(dataAt)) / int64(rows) + for i, col := range cols { + r, form, perr := parseTFORM(col.form) + if perr != nil { + return nil, 0, base.Errf("column %d: %w", i+1, perr) + } + if form == "A" { + // A character column's repeat is the string width: + // "16A" is one 16-character field, the layout every + // catalogue uses. The numeric repeat rule below would + // reject the library's own output. + if r < 1 { + return nil, 0, base.Errf("column %d: width %d is not positive", i+1, r) + } + forms[i] = "A" + widths[i] = r + } else { + if r != 1 { + return nil, 0, base.Errf("column %d: repeat %d is not supported, want a scalar", i+1, r) + } + widths[i] = tfSize(form) + // A zero width is an unknown code, or a variable-length + // descriptor ("P", "Q") the reader cannot size. Refusing + // it here keeps the row arithmetic below meaningful: a + // zero-width column would defeat the truncation guard and + // push the reads past the end of the file. + if widths[i] == 0 { + return nil, 0, base.Errf("column %d: unsupported TFORM %q", i+1, form) + } + forms[i] = form + } + if int64(widths[i]) > maxRow-int64(rowBytes) { + return nil, 0, base.Errf("column %d of width %d leaves the %d bytes of row the data holds", i+1, widths[i], maxRow) + } + offs[i] = rowBytes + rowBytes += widths[i] + } + // Bound before multiplying: a hostile NAXIS2 must not overflow + // the product (mirroring parseFITS's guard), it just means the + // declared data cannot fit the file. + if rowBytes > 0 && rows > (len(data)-dataAt)/rowBytes { + return nil, 0, base.Errf("table data is truncated") + } + for i := range cols { + form := forms[i] + if form == "A" { + table.Text = append(table.Text, fitsTextColumn(data, rows, dataAt+offs[i], rowBytes, widths[i])) + table.Columns = append(table.Columns, nil) + continue + } + var dt core.Dtype = core.Float + switch form { + case "B", "I", "J", "K", "L": + dt = core.Int + case "E": + dt = core.Float32 + } + arr := core.New(dt, rows) + // The form decides the loop once per column, and the cell goes + // straight from the file's bytes into the array's payload: no + // cell is boxed into an any and the row body holds no type + // switch. colBase is the column's first byte, so a row's cell + // starts rowBytes further on. + colBase := dataAt + offs[i] + switch form { + case "D": + dst := arr.RawFloats() + for row := range rows { + dst[row] = math.Float64frombits(binary.BigEndian.Uint64(data[colBase+rowBytes*row:])) + } + case "E": + dst := arr.RawFloat32s() + for row := range rows { + dst[row] = math.Float32frombits(binary.BigEndian.Uint32(data[colBase+rowBytes*row:])) + } + case "K": + dst := arr.RawInts() + for row := range rows { + dst[row] = int64(binary.BigEndian.Uint64(data[colBase+rowBytes*row:])) + } + case "J": + dst := arr.RawInts() + for row := range rows { + dst[row] = int64(int32(binary.BigEndian.Uint32(data[colBase+rowBytes*row:]))) + } + case "I": + dst := arr.RawInts() + for row := range rows { + dst[row] = int64(int16(binary.BigEndian.Uint16(data[colBase+rowBytes*row:]))) + } + case "B": + dst := arr.RawInts() + for row := range rows { + dst[row] = int64(data[colBase+rowBytes*row]) + } + case "L": + // A false logical leaves the zero the payload was + // allocated with. + dst := arr.RawInts() + for row := range rows { + if data[colBase+rowBytes*row] == 'T' { + dst[row] = 1 + } + } + } + table.Columns = append(table.Columns, scaleColumn(table.Headers, i+1, arr)) + table.Text = append(table.Text, nil) + } + return table, dataAt + fitsBlockSize(rows*rowBytes), nil + } + // ASCII table: NAXIS1-character text rows. + widths := make([]int, tfields) + starts := make([]int, tfields) + rowBytes := dims[0] + for i, col := range cols { + w, perr := parseASCIITFORM(col.form) + if perr != nil { + return nil, 0, base.Errf("column %d: %w", i+1, perr) + } + widths[i] = w + starts[i] = col.tbcol - 1 + // The bound is written without the sum, which a hostile TBCOL + // would wrap past: starts + width must leave the row. + if starts[i] < 0 || w > rowBytes || starts[i] > rowBytes-w { + return nil, 0, base.Errf("column %d: TBCOL %d with width %d leaves the row", i+1, col.tbcol, w) + } + } + // Bound before multiplying, as in the binary branch above. + if rowBytes > 0 && rows > (len(data)-dataAt)/rowBytes { + return nil, 0, base.Errf("table data is truncated") + } + for i, col := range cols { + if strings.HasSuffix(col.form, "A") { + table.Text = append(table.Text, fitsTextColumn(data, rows, dataAt+starts[i], rowBytes, widths[i])) + table.Columns = append(table.Columns, nil) + continue + } + isInt := strings.HasPrefix(col.form, "I") + var dt core.Dtype = core.Float + if isInt { + dt = core.Int + } + arr := core.New(dt, rows) + // The cell is parsed where it lies: the trimmed field is a view of + // the row's bytes, not a copy, and the strconv call sees exactly + // the text the string form saw. + colBase := dataAt + starts[i] + for row := range rows { + p := colBase + rowBytes*row + field := bytes.TrimSpace(data[p : p+widths[i]]) + if len(field) == 0 { + continue + } + if isInt { + v, cerr := strconv.ParseInt(fitsFieldString(field), 10, 64) + if cerr != nil { + return nil, 0, base.Errf("column %d row %d: %q is not an integer", i+1, row, string(field)) + } + arr.RawInts()[row] = v + } else { + v, ok := fitsASCIIFloat(fitsFieldString(field)) + if !ok { + return nil, 0, base.Errf("column %d row %d: %q is not a number", i+1, row, string(field)) + } + arr.RawFloats()[row] = v + } + } + table.Columns = append(table.Columns, scaleColumn(table.Headers, i+1, arr)) + table.Text = append(table.Text, nil) + } + return table, dataAt + fitsBlockSize(rows*rowBytes), nil +} + +// parseTFORM splits a binary TFORM into its repeat count and type +// code: "3E" gives (3, "E"), "D" gives (1, "D"), "16A" gives (16, "A"). +func parseTFORM(form string) (int, string, error) { + if form == "" { + return 0, "", base.Errf("empty TFORM") + } + repeat, code := 1, form + for i, r := range form { + if r < '0' || r > '9' { + code = form[i:] + if i > 0 { + v, aerr := strconv.Atoi(form[:i]) + if aerr != nil { + // An overflowing repeat count must not silently + // fall back to 1: a hostile "99999999999999999999E" + // would be read as a scalar instead of refused. + return 0, "", base.Errf("TFORM %q: the repeat count does not fit an integer", form) + } + repeat = v + } + break + } + } + return repeat, code, nil +} + +// fitsTextColumn builds a character column's strings from the +// fixed-width cells that start at colBase and repeat every rowStride +// bytes, trailing spaces trimmed. The trimmed cell bytes are copied +// into one slab and every string of the column aliases its own range +// of that slab, so the column costs one allocation whatever its row +// count. The caller's guards have already bounded colBase + rowStride* +// (rows-1) + width against the file's bytes. +// +// SAFETY: the slab is fully written before the first string is formed +// from it, the aliased ranges are disjoint, and nothing writes to the +// slab afterwards, so each string keeps exactly the bytes the copy +// left there for as long as it lives. +func fitsTextColumn(data []byte, rows, colBase, rowStride, width int) []string { + text := make([]string, rows) + total := 0 + for row := range rows { + base := colBase + rowStride*row + p := base + width + for p > base && data[p-1] == ' ' { + p-- + } + total += p - base + } + slab := make([]byte, total) + at := 0 + for row := range rows { + base := colBase + rowStride*row + p := base + width + for p > base && data[p-1] == ' ' { + p-- + } + n := p - base + copy(slab[at:at+n], data[base:p]) + if n == 0 { + text[row] = "" + } else { + text[row] = unsafe.String(unsafe.SliceData(slab[at:at+n]), n) + } + at += n + } + return text +} + +// scaleColumn applies the TSCALn/TZEROn affine map (physical = raw* +// scale + zero) to a decoded numeric column. Integer scaling that +// stays integral keeps the int dtype (the unsigned-integer convention +// TZERO = 32768/2147483648 lands here); any fractional or overflowing +// scaling promotes the column to float64, mirroring cfitsio. A +// malformed value is treated as unscaled: the raw storage values are +// the best available answer and an error here would lose the whole +// table over one keyword. +func scaleColumn(headers map[string]string, n int, arr *core.Array) *core.Array { + ss := fitsLookupN(headers, "TSCAL", n) + zs := fitsLookupN(headers, "TZERO", n) + if ss == "" && zs == "" { + return arr + } + scale, zero := 1.0, 0.0 + if ss != "" { + if v, err := strconv.ParseFloat(strings.TrimSpace(ss), 64); err == nil { + scale = v + } + } + if zs != "" { + if v, err := strconv.ParseFloat(strings.TrimSpace(zs), 64); err == nil { + zero = v + } + } + if scale == 1 && zero == 0 { + return arr + } + switch arr.Dtype() { + case core.Float: + raw := arr.RawFloats() + for i := range raw { + raw[i] = raw[i]*scale + zero + } + return arr + case core.Float32: + // Scaled values leave the float32 guarantee; keep the full + // float64 result instead of rounding twice. + raw32 := arr.RawFloat32s() + out := core.New(core.Float, arr.Shape()...) + raw := out.RawFloats() + for i := range raw32 { + raw[i] = float64(raw32[i])*scale + zero + } + return out + default: + rawInts := arr.RawInts() + if scale == math.Trunc(scale) && zero == math.Trunc(zero) && + math.Abs(scale) < 1e15 && math.Abs(zero) < 9e15 { + out := core.New(core.Int, arr.Shape()...) + dst := out.RawInts() + integral := true + for i, v := range rawInts { + sv := float64(v)*scale + zero + // The exact float64 bounds of the int64 range: MaxInt64 + // is 9223372036854775807, which no float64 holds, so the + // first value past it is 2^63, while MinInt64 is exact. + // A conversion out of range is implementation-defined + // (amd64 answers MinInt64), which is the silent corruption + // this path exists to avoid. + if sv >= 9223372036854775808.0 || sv < -9223372036854775808.0 { + integral = false + break + } + dst[i] = int64(sv) + } + if integral { + return out + } + } + out := core.New(core.Float, arr.Shape()...) + raw := out.RawFloats() + for i, v := range rawInts { + raw[i] = float64(v)*scale + zero + } + return out + } +} + +// tfSize returns the byte width of one element of a binary column. +func tfSize(code string) int { + switch code { + case "L", "B", "A": + return 1 + case "I": + return 2 + case "J", "E": + return 4 + case "K", "D": + return 8 + } + return 0 +} + +// parseASCIITFORM returns the character width of an ASCII-table +// column: "10A" gives 10, "I7" gives 7, "F14.6" gives 14, "E12.4" +// gives 12, "D20.12" gives 20. +func parseASCIITFORM(form string) (int, error) { + w, _, _, err := parseASCIISaveForm(strings.ToUpper(strings.TrimSpace(form))) + return w, err +} + +// fitsASCIIFloat parses one ASCII-table cell. Beyond the forms +// strconv.ParseFloat accepts, the cells of a Fortran-written table +// carry the Dw.d exponent ("1.5D+03", the D column type this reader +// accepts) and the exponent-less form ("1.5+03"), where the sign after +// the mantissa introduces the exponent; both mean what the same text +// with an E means, and refusing them refused the whole table. +func fitsASCIIFloat(field string) (float64, bool) { + if v, err := strconv.ParseFloat(field, 64); err == nil { + return v, true + } + norm := field + switch i := strings.IndexAny(norm, "dD"); { + case i >= 0: + // A Fortran D exponent is the E exponent under another letter. + norm = norm[:i] + "E" + norm[i+1:] + case strings.IndexAny(norm, "eE") < 0: + // No exponent at all: the exponent-less Fortran form puts the + // exponent's sign directly behind the mantissa, "1.5+03". + if j := strings.IndexAny(norm[1:], "+-"); j >= 0 { + norm = norm[:j+1] + "E" + norm[j+1:] + } + } + if norm == field { + return 0, false + } + v, err := strconv.ParseFloat(norm, 64) + if err != nil { + return 0, false + } + return v, true +} + +// fitsFieldString views one trimmed ASCII-table field as a string +// without copying it, the form strconv needs and the byte slice +// already holds. +// +// SAFETY: the view aliases the file's bytes for the length of one +// strconv call, which reads the string and retains nothing of it. The +// parse errors are formatted from a copy, so no view of the file +// outlives the slice it points into. +func fitsFieldString(b []byte) string { + return unsafe.String(unsafe.SliceData(b), len(b)) +} + +// fitsCheckString validates a human-supplied string field. +func fitsCheckString(v, what string) error { + if len(v) == 0 { + return base.Errf("%s must not be empty", what) + } + if len(v) > 68 { + return base.Errf("%s is %d characters, at most 68 fit a card", what, len(v)) + } + return nil +} diff --git a/io/fitstable_keyword_pin_test.go b/io/fitstable_keyword_pin_test.go new file mode 100644 index 0000000..c191f53 --- /dev/null +++ b/io/fitstable_keyword_pin_test.go @@ -0,0 +1,80 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "os" + "path/filepath" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The binary-table reader resolves header keywords through a map whose +// documented rule is first-occurrence-wins, the behaviour the previous +// linear scan answered. A hostile header may repeat a structural card +// with a different value: this pin fixes that the reader follows the +// first value of each repeated card. +func TestLoadFITSTableRepeatedKeywordsFirstWins(t *testing.T) { + body := make([]byte, 16) + for r := range 2 { + binary.BigEndian.PutUint64(body[r*8:], uint64(1000+r)) + } + cards := []string{ + fitsStringCardRaw("XTENSION", "BINTABLE"), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 2), + fitsIntCard("NAXIS1", 8), + fitsIntCard("NAXIS2", 2), + fitsIntCard("NAXIS2", 99), + fitsIntCard("PCOUNT", 0), + fitsIntCard("GCOUNT", 1), + fitsIntCard("TFIELDS", 1), + fitsIntCard("TFIELDS", 5), + fitsStringCardRaw("TTYPE1", "COL1"), + fitsStringCardRaw("TFORM1", "K"), + fitsStringCardRaw("TFORM1", "D"), + fitsEndCard(), + } + out := fitsAppendCards(nil, []string{ + fitsBoolCard("SIMPLE", true), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 0), + fitsBoolCard("EXTEND", true), + fitsEndCard(), + }) + out = fitsAppendCards(out, cards) + out = append(out, body...) + out = fitsAppendZeroPad(out) + path := filepath.Join(t.TempDir(), "repeated.fits") + if err := os.WriteFile(path, out, 0o644); err != nil { + t.Fatal(err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + if table.Rows != 2 { + t.Fatalf("rows = %d, want 2 from the first NAXIS2", table.Rows) + } + if len(table.Columns) != 1 { + t.Fatalf("columns = %d, want 1 from the first TFIELDS", len(table.Columns)) + } + col := table.Columns[0] + if col == nil { + t.Fatal("the first column came back nil; the repeated TFORM1 must still decode the first form") + } + if len(table.Text) > 0 && table.Text[0] != nil { + t.Fatalf("column text = %v, want nil for a K column decoded from the first TFORM1", table.Text[0]) + } + if col.Dtype() != core.Int { + t.Fatalf("column dtype = %s, want Int from the first TFORM1 (K)", col.Dtype()) + } + for i, want := range []int64{1000, 1001} { + if got := col.RawInts()[i]; got != want { + t.Fatalf("value %d = %d, want %d", i, got, want) + } + } +} diff --git a/io/fitstable_test.go b/io/fitstable_test.go new file mode 100644 index 0000000..37ee0f2 --- /dev/null +++ b/io/fitstable_test.go @@ -0,0 +1,441 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "bytes" + "encoding/binary" + "math" + "os" + "path/filepath" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// mustFloats32 builds a float32 array. +func mustFloats32(t *testing.T, vals []float32, shape ...int) *core.Array { + t.Helper() + if len(shape) == 0 { + shape = []int{len(vals)} + } + a, err := core.FromFloat32s(vals, shape...) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + return a +} + +// mustInts builds an int array. +func mustInts(t *testing.T, vals []int64, shape ...int) *core.Array { + t.Helper() + if len(shape) == 0 { + shape = []int{len(vals)} + } + a, err := core.FromInts(vals, shape...) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + return a +} + +// osReadFile and osWriteFile wrap the os calls so the test table +// splicing stays terse. +func osReadFile(path string) ([]byte, error) { return os.ReadFile(path) } + +func osWriteFile(path string, parts ...[]byte) error { + var out []byte + for _, p := range parts { + out = append(out, p...) + } + return os.WriteFile(path, out, 0o644) +} + +// starCatalogue returns a deterministic three-column table: integer +// identifiers, float64 magnitudes and float32 temperatures. +func starCatalogue(t *testing.T) ([]FITSTableColumn, int) { + t.Helper() + const n = 5 + ids := make([]float64, n) + for i := range n { + ids[i] = float64(1000 + i) + } + ints := make([]int64, n) + for i := range n { + ints[i] = int64(1000 + i) + } + mags := make([]float64, n) + for i := range n { + mags[i] = math.Sin(float64(3*i+1)) * 5 + } + temps := make([]float32, n) + for i := range n { + temps[i] = float32(3000 + 700*i) + } + idArr := core.New(core.Int, n) + copy(idArr.RawInts(), ints) + cols := []FITSTableColumn{ + {Name: "ID", Unit: "", Form: "K", Data: idArr}, + {Name: "MAG", Unit: "mag", Form: "D", Data: mustFloats(t, mags, n)}, + {Name: "TEMP", Unit: "K", Form: "E", Data: mustFloats32(t, temps, n)}, + } + return cols, n +} + +// TestSaveLoadFITSTableBinary round-trips a binary table: names, +// units, and every value read back unchanged. +func TestSaveLoadFITSTableBinary(t *testing.T) { + cols, n := starCatalogue(t) + path := filepath.Join(t.TempDir(), "catalogue.fits") + if err := SaveFITSTable(path, false, cols, map[string]string{"ORIGIN": "tensor test"}); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + if table.Kind != "BINTABLE" { + t.Fatalf("kind %q, want BINTABLE", table.Kind) + } + if table.Rows != n { + t.Fatalf("rows = %d, want %d", table.Rows, n) + } + if table.Headers["ORIGIN"] != "tensor test" { + t.Fatalf("header ORIGIN = %q", table.Headers["ORIGIN"]) + } + wantNames := []string{"ID", "MAG", "TEMP"} + for i, want := range wantNames { + if table.Names[i] != want { + t.Fatalf("column %d named %q, want %q", i, table.Names[i], want) + } + } + if table.Units[1] != "mag" { + t.Fatalf("MAG unit = %q, want mag", table.Units[1]) + } + for i := range n { + if table.Columns[0].RawInts()[i] != int64(1000+i) { + t.Fatalf("ID[%d] = %d", i, table.Columns[0].RawInts()[i]) + } + if math.Abs(table.Columns[1].FloatAt(i)-cols[1].Data.FloatAt(i)) > 1e-12 { + t.Fatalf("MAG[%d] = %.14g", i, table.Columns[1].FloatAt(i)) + } + if math.Abs(float64(table.Columns[2].RawFloat32s()[i]-cols[2].Data.RawFloat32s()[i])) > 1e-4 { + t.Fatalf("TEMP[%d] = %.6g", i, table.Columns[2].RawFloat32s()[i]) + } + } + // The primary image still reads through the image loader. + img, headers, err := LoadFITS(path) + if err != nil { + t.Fatalf("LoadFITS on a table file: %v", err) + } + if img.Len() != 0 { + t.Fatalf("primary image has %d elements, want an empty zero-axis HDU", img.Len()) + } + if headers["EXTEND"] != "T" { + t.Fatalf("EXTEND = %q, want T", headers["EXTEND"]) + } +} + +// TestSaveLoadFITSTableBinaryStringColumnNotFirst round-trips a +// binary table whose character column sits between two numeric ones. +// A character encoder that writes at the start of the row instead of +// its own column offset corrupts every column beside it, and the +// damage is silent: the file still loads. +func TestSaveLoadFITSTableBinaryStringColumnNotFirst(t *testing.T) { + const n = 3 + mags := []float64{1.25, -0.5, 3.75} + names := []string{"alf", "bet", "gam"} + ids := []int64{7, 8, 9} + cols := []FITSTableColumn{ + {Name: "MAG", Unit: "mag", Form: "D", Data: mustFloats(t, mags, n)}, + {Name: "STAR", Form: "8A", Text: names}, + {Name: "ID", Form: "K", Data: mustInts(t, ids, n)}, + } + path := filepath.Join(t.TempDir(), "stars.fits") + if err := SaveFITSTable(path, false, cols, nil); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + for i := range n { + if got := table.Columns[0].FloatAt(i); got != mags[i] { + t.Fatalf("MAG[%d] = %v, want %v", i, got, mags[i]) + } + if got := table.Text[1][i]; got != names[i] { + t.Fatalf("STAR[%d] = %q, want %q", i, got, names[i]) + } + if got := table.Columns[2].RawInts()[i]; got != ids[i] { + t.Fatalf("ID[%d] = %d, want %d", i, got, ids[i]) + } + } + // The row layout itself: 8 bytes of float64, 8 of text, 8 of int. + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read back: %v", err) + } + start := bytes.Index(raw, []byte("XTENSION")) + if start < 0 { + t.Fatal("no XTENSION card in the file") + } + // The data starts on the 2880-byte boundary after the extension + // header, whose card count the loader has already validated. + data := raw[(start/2880+1)*2880:] + if got := math.Float64frombits(binary.BigEndian.Uint64(data[0:])); got != mags[0] { + t.Fatalf("first row's first field = %v, want %v", got, mags[0]) + } + if got := string(bytes.TrimRight(data[8:16], " ")); got != names[0] { + t.Fatalf("first row's text field = %q, want %q", got, names[0]) + } + if got := int64(binary.BigEndian.Uint64(data[16:24])); got != ids[0] { + t.Fatalf("first row's int field = %d, want %d", got, ids[0]) + } +} + +// TestSaveLoadFITSTableASCII round-trips an ASCII table with integer, +// float and string columns. +func TestSaveLoadFITSTableASCII(t *testing.T) { + const n = 4 + names := []string{"ALF Cen", "Betel", "Rigel", "Deneb"} + mags := make([]float64, n) + for i := range n { + mags[i] = -1.5 + 1.3*float64(i) + } + cols := []FITSTableColumn{ + {Name: "STAR", Form: "10A", Text: names}, + {Name: "MAG", Unit: "mag", Form: "D20.14", Data: mustFloats(t, mags, n)}, + } + path := filepath.Join(t.TempDir(), "stars.fits") + if err := SaveFITSTable(path, true, cols, nil); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + if table.Kind != "TABLE" { + t.Fatalf("kind %q, want TABLE", table.Kind) + } + for i := range n { + if table.Text[0][i] != names[i] { + t.Fatalf("STAR[%d] = %q, want %q", i, table.Text[0][i], names[i]) + } + if math.Abs(table.Columns[1].FloatAt(i)-mags[i]) > 1e-9 { + t.Fatalf("MAG[%d] = %.14g, want %.14g", i, table.Columns[1].FloatAt(i), mags[i]) + } + } +} + +// TestLoadFITSTableSkipsImage writes an image first and a table +// second by hand-concatenating the two files' HDUs: the loader must +// skip the image and land on the table. +func TestLoadFITSTableSkipsImage(t *testing.T) { + dir := t.TempDir() + img := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2) + imagePath := filepath.Join(dir, "image.fits") + if err := SaveFITS(imagePath, img, nil); err != nil { + t.Fatalf("SaveFITS: %v", err) + } + cols := []FITSTableColumn{ + {Name: "X", Form: "D", Data: mustFloats(t, []float64{1.5, 2.5}, 2)}, + } + tablePath := filepath.Join(dir, "table.fits") + if err := SaveFITSTable(tablePath, false, cols, nil); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + combined := filepath.Join(dir, "combined.fits") + imageData, err := osReadFile(imagePath) + if err != nil { + t.Fatal(err) + } + tableData, err := osReadFile(tablePath) + if err != nil { + t.Fatal(err) + } + // The image file's primary already declares EXTEND; append the + // table extension with its own primary stripped (the extension + // starts at its XTENSION card, 2880 bytes in). + if err := osWriteFile(combined, imageData, tableData[2880:]); err != nil { + t.Fatal(err) + } + table, err := LoadFITSTable(combined) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + if table.Columns[0].FloatAt(1) != 2.5 { + t.Fatalf("X[1] = %.4g, want 2.5", table.Columns[0].FloatAt(1)) + } +} + +// TestSaveFITSTableErrors pins the validation contract. +func TestSaveFITSTableErrors(t *testing.T) { + if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, nil, nil); err == nil { + t.Fatal("expected an error for an empty column list") + } + noData := []FITSTableColumn{{Name: "X", Form: "D"}} + if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, noData, nil); err == nil { + t.Fatal("expected an error for a numeric column without data") + } + noText := []FITSTableColumn{{Name: "S", Form: "8A"}} + if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, noText, nil); err == nil { + t.Fatal("expected an error for a character column without text") + } + badForm := []FITSTableColumn{{Name: "X", Form: "Q", Data: mustFloats(t, []float64{1}, 1)}} + if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, badForm, nil); err == nil { + t.Fatal("expected an error for an unknown form") + } + ragged := []FITSTableColumn{ + {Name: "X", Form: "D", Data: mustFloats(t, []float64{1, 2}, 2)}, + {Name: "Y", Form: "D", Data: mustFloats(t, []float64{1}, 1)}, + } + if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, ragged, nil); err == nil { + t.Fatal("expected an error for mismatched row counts") + } + longName := []FITSTableColumn{{Name: "S", Form: "4A", Text: []string{"too long"}}} + if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, longName, nil); err == nil { + t.Fatal("expected an error for text wider than the form") + } +} + +// TestLoadFITSTableHostileRows pins the NAXIS2 guards: a hostile +// header declaring a colossal row count must report truncation (not +// overflow or an out-of-memory allocation), and a negative one must +// be rejected instead of answering an empty table. +func TestLoadFITSTableHostileRows(t *testing.T) { + cols, n := starCatalogue(t) + path := filepath.Join(t.TempDir(), "catalogue.fits") + if err := SaveFITSTable(path, false, cols, nil); err != nil { + t.Fatalf("SaveFITSTable: %v", err) + } + raw, err := osReadFile(path) + if err != nil { + t.Fatal(err) + } + hostile := filepath.Join(t.TempDir(), "hostile.fits") + huge := osWriteFile(hostile, bytes.Replace(raw, []byte(fitsIntCard("NAXIS2", n)), []byte(fitsIntCard("NAXIS2", math.MaxInt64)), 1)) + if huge != nil { + t.Fatalf("os.WriteFile: %v", huge) + } + if _, err := LoadFITSTable(hostile); err == nil { + t.Fatal("expected an error for a colossal NAXIS2") + } + negative := filepath.Join(t.TempDir(), "negative.fits") + if err := osWriteFile(negative, bytes.Replace(raw, []byte(fitsIntCard("NAXIS2", n)), []byte(fitsIntCard("NAXIS2", -5)), 1)); err != nil { + t.Fatalf("os.WriteFile: %v", err) + } + if _, err := LoadFITSTable(negative); err == nil { + t.Fatal("expected an error for a negative NAXIS2") + } +} + +// TestSaveFITSTableASCIIWidthErrors pins the ASCII field contract: a +// value whose rendering exceeds the form's declared width is an +// error, never silently truncated digits or an overflowing field. +func TestSaveFITSTableASCIIWidthErrors(t *testing.T) { + dir := t.TempDir() + wideInt := []FITSTableColumn{{Name: "N", Form: "I5", Data: mustInts(t, []int64{12345678}, 1)}} + if err := SaveFITSTable(filepath.Join(dir, "i.fits"), true, wideInt, nil); err == nil { + t.Error("expected an error for an integer that does not fit I5") + } + wideFloat := []FITSTableColumn{{Name: "F", Form: "F10.4", Data: mustFloats(t, []float64{1e300}, 1)}} + if err := SaveFITSTable(filepath.Join(dir, "f.fits"), true, wideFloat, nil); err == nil { + t.Error("expected an error for a float that does not fit F10.4") + } + wideText := []FITSTableColumn{{Name: "S", Form: "3A", Text: []string{"abcd"}}} + if err := SaveFITSTable(filepath.Join(dir, "a.fits"), true, wideText, nil); err == nil { + t.Error("expected an error for text wider than the form") + } + // Fitting values keep working. + okCols := []FITSTableColumn{ + {Name: "N", Form: "I5", Data: mustInts(t, []int64{42}, 1)}, + {Name: "F", Form: "F10.4", Data: mustFloats(t, []float64{3.14159}, 1)}, + {Name: "E", Form: "E13.4", Data: mustFloats(t, []float64{-1.5e300}, 1)}, + } + path := filepath.Join(dir, "ok.fits") + if err := SaveFITSTable(path, true, okCols, nil); err != nil { + t.Fatalf("SaveFITSTable with fitting values: %v", err) + } + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + if table.Columns[0].RawInts()[0] != 42 { + t.Errorf("N[0] = %d, want 42", table.Columns[0].RawInts()[0]) + } + if math.Abs(table.Columns[1].FloatAt(0)-3.14159) > 1e-4 { + t.Errorf("F[0] = %.6g, want 3.14159", table.Columns[1].FloatAt(0)) + } + if got, want := table.Columns[2].FloatAt(0), -1.5e300; math.Abs(got-want) > 1e-6*math.Abs(want) { + t.Errorf("E[0] = %.6g, want %.6g", got, want) + } +} + +// TestLoadFITSTableSignedBinaryForms pins the four numeric decodes the +// package's own writer cannot emit: J is a signed 32-bit big-endian +// integer, I a signed 16-bit one, B an unsigned byte and L a logical +// whose false value is the zero the payload was allocated with. The +// table is laid out here, so every value that distinguishes the forms +// is present: a negative integer in each of J and I, a byte above the +// signed range and a false logical beside a true one. +func TestLoadFITSTableSignedBinaryForms(t *testing.T) { + hdr := cardBlock( + card("XTENSION= 'BINTABLE'"), + card("BITPIX = 8"), + card("NAXIS = 2"), + card("NAXIS1 = 8"), + card("NAXIS2 = 3"), + card("TFIELDS = 4"), + card("TTYPE1 = 'JCOL '"), + card("TFORM1 = 'J '"), + card("TTYPE2 = 'ICOL '"), + card("TFORM2 = 'I '"), + card("TTYPE3 = 'BCOL '"), + card("TFORM3 = 'B '"), + card("TTYPE4 = 'LCOL '"), + card("TFORM4 = 'L '"), + card("END"), + ) + rows := []struct { + j int32 + i int16 + b byte + l byte + }{ + {-123456, -7, 200, 'T'}, + {123456, 30000, 255, 'F'}, + {-1, -32768, 128, 'T'}, + } + var body []byte + for _, r := range rows { + body = binary.BigEndian.AppendUint32(body, uint32(r.j)) + body = binary.BigEndian.AppendUint16(body, uint16(r.i)) + body = append(body, r.b, r.l) + } + path := writeHostile(t, "forms.fits", append(hdr, body...)) + table, err := LoadFITSTable(path) + if err != nil { + t.Fatalf("LoadFITSTable: %v", err) + } + if len(table.Columns) != 4 { + t.Fatalf("the table carries %d columns, want 4", len(table.Columns)) + } + for row, want := range rows { + if got := table.Columns[0].RawInts()[row]; got != int64(want.j) { + t.Fatalf("row %d: J = %d, want %d (a signed 32-bit big-endian integer)", row, got, want.j) + } + if got := table.Columns[1].RawInts()[row]; got != int64(want.i) { + t.Fatalf("row %d: I = %d, want %d (a signed 16-bit big-endian integer)", row, got, want.i) + } + if got := table.Columns[2].RawInts()[row]; got != int64(want.b) { + t.Fatalf("row %d: B = %d, want %d (an unsigned byte)", row, got, want.b) + } + wantL := int64(0) + if want.l == 'T' { + wantL = 1 + } + if got := table.Columns[3].RawInts()[row]; got != wantL { + t.Fatalf("row %d: L = %d, want %d (%q in the file)", row, got, wantL, want.l) + } + } +} diff --git a/io/fuzz_test.go b/io/fuzz_test.go new file mode 100644 index 0000000..33cda36 --- /dev/null +++ b/io/fuzz_test.go @@ -0,0 +1,474 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "math" + "os" + "path/filepath" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Coverage-guided fuzzing over the four binary readers. Every input is +// written to one shared temp file per worker process and handed to the +// reader; a panic, a hang or an allocation blow-up fails the input, and +// so does a returned result whose shape disagrees with its payload, +// because a silent mismatch is the worst outcome a parser can produce. +// Seeds are the real fixtures plus files built by the own writers, so +// the mutator starts from genuinely valid bytes rather than from magic +// constants. Run one target at a time, for example: +// +// go test -run '^$' -fuzz FuzzLoadHDF5 -fuzztime 2m ./io/ +// +// Inputs that expose a defect land in testdata/fuzz// and +// stay there as regression seeds. + +// fuzzWrite seeds the corpus with one file built by a writer. +func fuzzWrite(f *testing.F, name string, build func(path string) error) { + path := filepath.Join(f.TempDir(), name) + if err := build(path); err != nil { + f.Fatalf("build seed %s: %v", name, err) + } + data, err := os.ReadFile(path) + if err != nil { + f.Fatalf("read seed %s: %v", name, err) + } + f.Add(data) +} + +// fuzzFloats builds a float64 vector of n recognisably distinct values. +func fuzzFloats(f *testing.F, n int) *core.Array { + f.Helper() + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(i+1) + 0.5 + } + a, err := core.FromFloats(vals, n) + if err != nil { + f.Fatalf("FromFloats(%d): %v", n, err) + } + return a +} + +// fuzzInts builds an int64 vector of n recognisably distinct values. +func fuzzInts(f *testing.F, n int) *core.Array { + f.Helper() + vals := make([]int64, n) + for i := range vals { + vals[i] = int64(i) - 2 + } + a, err := core.FromInts(vals, n) + if err != nil { + f.Fatalf("FromInts(%d): %v", n, err) + } + return a +} + +// fuzzShapeProduct multiplies the shape with the overflow guard the +// parsers themselves are held to: an attacker-controlled shape must +// never wrap around into a plausible small length. +func fuzzShapeProduct(t *testing.T, shape []int) int { + n := 1 + for _, d := range shape { + if d < 0 || d > (1<<31-1)/n { + t.Fatalf("hostile shape %v accepted", shape) + } + n *= d + } + return n +} + +func FuzzLoadHDF5(f *testing.F) { + fuzzWrite(f, "fixture.h5", func(p string) error { + data, err := os.ReadFile("testdata/h5/fixture.h5") + if err != nil { + return err + } + return os.WriteFile(p, data, 0o644) + }) + fuzzWrite(f, "fletcher.h5", func(p string) error { + data, err := os.ReadFile("testdata/h5/fletcher.h5") + if err != nil { + return err + } + return os.WriteFile(p, data, 0o644) + }) + fuzzWrite(f, "latest.h5", func(p string) error { + data, err := os.ReadFile("testdata/h5/latest.h5") + if err != nil { + return err + } + return os.WriteFile(p, data, 0o644) + }) + path := filepath.Join(f.TempDir(), "fuzz.h5") + f.Fuzz(func(t *testing.T, data []byte) { + if err := os.WriteFile(path, data, 0o644); err != nil { + t.Fatal(err) + } + sets, err := LoadHDF5(path) + if err != nil { + return + } + for _, s := range sets { + if n := fuzzShapeProduct(t, s.Shape); n != s.Values.Len() { + t.Fatalf("dataset %s: shape product %d != payload %d", s.Path, n, s.Values.Len()) + } + } + }) +} + +func FuzzLoadNetCDF(f *testing.F) { + fuzzWrite(f, "classic.nc", func(p string) error { + dims := []NetCDFDim{{Name: "lat", Length: 3}, {Name: "lon", Length: 4}} + vars := []NetCDFVar{ + {Name: "temp", Dims: []string{"lat", "lon"}, Values: fuzzFloats(f, 12)}, + {Name: "mask", Dims: []string{"lon"}, Values: fuzzInts(f, 4)}, + } + return SaveNetCDF(p, dims, vars, map[string]string{"title": "seed"}) + }) + fuzzWrite(f, "record.nc", func(p string) error { + dims := []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 2}} + vars := []NetCDFVar{ + {Name: "series", Dims: []string{"t", "x"}, Values: fuzzFloats(f, 6)}, + } + return SaveNetCDF(p, dims, vars, nil) + }) + path := filepath.Join(f.TempDir(), "fuzz.nc") + f.Fuzz(func(t *testing.T, data []byte) { + if err := os.WriteFile(path, data, 0o644); err != nil { + t.Fatal(err) + } + dims, vars, _, err := LoadNetCDF(path) + if err != nil { + return + } + lengths := make(map[string]int, len(dims)) + for _, d := range dims { + lengths[d.Name] = d.Length + } + for _, v := range vars { + // A rank-0 (scalar) variable comes back as a one-element + // vector with no dimension names, the documented shape. + if len(v.Values.Shape()) != len(v.Dims) && !(len(v.Dims) == 0 && v.Values.Len() == 1) { + t.Fatalf("variable %s: rank %d != %d named dimensions", v.Name, len(v.Values.Shape()), len(v.Dims)) + } + record := false + n := 1 + for _, name := range v.Dims { + length, ok := lengths[name] + if !ok { + t.Fatalf("variable %s: unknown dimension %s", v.Name, name) + } + if length == 0 { + record = true + continue + } + n *= length + } + if !record && n != v.Values.Len() { + t.Fatalf("variable %s: dimension product %d != payload %d", v.Name, n, v.Values.Len()) + } + } + }) +} + +func FuzzLoadFITS(f *testing.F) { + fuzzWrite(f, "image.fits", func(p string) error { + a, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}, 4, 4) + if err != nil { + return err + } + return SaveFITS(p, a, map[string]string{"OBJECT": "seed", "TELESCOP": "TENSOR"}) + }) + path := filepath.Join(f.TempDir(), "fuzz.fits") + f.Fuzz(func(t *testing.T, data []byte) { + if err := os.WriteFile(path, data, 0o644); err != nil { + t.Fatal(err) + } + a, _, err := LoadFITS(path) + if err != nil { + return + } + if n := fuzzShapeProduct(t, a.Shape()); n != a.Len() { + t.Fatalf("image: shape product %d != payload %d", n, a.Len()) + } + }) +} + +func FuzzLoadFITSTable(f *testing.F) { + fuzzWrite(f, "binary.fits", func(p string) error { + cols := []FITSTableColumn{ + {Name: "flux", Form: "D", Data: fuzzFloats(f, 5)}, + {Name: "id", Form: "K", Data: fuzzInts(f, 5)}, + {Name: "name", Form: "8A", Text: []string{"alpha", "beta", "gamma", "delta", "epsilon"}}, + } + return SaveFITSTable(p, false, cols, nil) + }) + fuzzWrite(f, "ascii.fits", func(p string) error { + cols := []FITSTableColumn{ + {Name: "flux", Form: "E12.5", Data: fuzzFloats(f, 3)}, + {Name: "name", Form: "6A", Text: []string{"one", "two", "three"}}, + } + return SaveFITSTable(p, true, cols, nil) + }) + path := filepath.Join(f.TempDir(), "fuzztable.fits") + f.Fuzz(func(t *testing.T, data []byte) { + if err := os.WriteFile(path, data, 0o644); err != nil { + t.Fatal(err) + } + table, err := LoadFITSTable(path) + if err != nil { + return + } + if table.Rows < 0 { + t.Fatalf("negative row count %d", table.Rows) + } + n := len(table.Names) + if len(table.Columns) != n || len(table.Text) != n { + t.Fatalf("column lists disagree: %d names, %d value columns, %d text columns", n, len(table.Columns), len(table.Text)) + } + for i := range table.Names { + switch { + case table.Columns[i] != nil && table.Text[i] != nil: + t.Fatalf("column %d carries both values and text", i) + case table.Columns[i] != nil: + if table.Columns[i].Len() != table.Rows { + t.Fatalf("column %d: length %d != %d rows", i, table.Columns[i].Len(), table.Rows) + } + case table.Text[i] != nil: + if len(table.Text[i]) != table.Rows { + t.Fatalf("text column %d: length %d != %d rows", i, len(table.Text[i]), table.Rows) + } + default: + t.Fatalf("column %d carries neither values nor text", i) + } + } + }) +} + +// FuzzHDF5WriteRead fuzzes the writer's round-trip contract: whatever +// parameters and values the mutator picks, SaveHDF5 must produce a +// file the reader decodes back to the very same shapes, bit patterns +// and attributes. A write that fails, a read that fails or a value +// that moves is a bug, not a skipped input: unlike the readers there +// is no untrusted bytes here, the writer owns every byte it emits. +func FuzzHDF5WriteRead(f *testing.F) { + for _, seed := range [][]byte{ + {0, 0, 7, 3, 0, 0, 0, 0}, + {0, 0, 7, 3, 0, 0, 1, 0}, + {1, 1, 19, 5, 1, 1, 0, 1}, + {2, 0, 33, 1, 0, 1, 1, 1}, + {0, 1, 5, 7, 1, 0, 0, 2}, + {2, 1, 11, 4, 1, 1, 1, 1}, + {3, 0, 5, 4, 0, 0, 0, 0}, + {4, 1, 9, 3, 0, 1, 0, 1}, + {5, 0, 6, 2, 0, 0, 1, 0}, + {6, 1, 12, 5, 0, 0, 0, 1}, + {7, 0, 8, 3, 1, 1, 0, 0}, + {8, 1, 15, 4, 0, 0, 0, 1}, + {9, 0, 10, 2, 1, 0, 0, 0}, + } { + f.Add(seed) + } + f.Fuzz(func(t *testing.T, data []byte) { + if len(data) < 8 { + t.Skip() + } + dtype := int(data[0] % 10) + rank := 1 + int(data[1]%2) + d0 := 1 + int(data[2])%31 + d1 := 1 + int(data[3])%9 + gzip := 0 + if data[4]%2 == 1 { + gzip = 6 + } + shuffle := data[5]%2 == 1 + // A filtered dataset in a latest-version file is refused by + // design; the fuzz contract covers the accepted combinations. + latest := data[6]%2 == 1 && gzip == 0 && !shuffle + chunk := 0 + if data[7]%2 == 1 { + chunk = 32 + } + n := d0 + if rank == 2 { + n *= d1 + } + shape := []int{d0} + if rank == 2 { + shape = append(shape, d1) + } + var values *core.Array + var err error + switch dtype { + case 0: + v := make([]float64, n) + for i := range v { + v[i] = float64(data[(8+i)%len(data)]) - 128 + float64(i%5)*0.25 + } + values, err = core.FromFloats(v, shape...) + case 1: + v := make([]float32, n) + for i := range v { + v[i] = float32(data[(8+i)%len(data)]) - 128 + float32(i%5)*0.25 + } + values, err = core.FromFloat32s(v, shape...) + case 2: + v := make([]int64, n) + for i := range v { + v[i] = int64(data[(8+i)%len(data)])*1000 - 128000 + int64(i) + } + values, err = core.FromInts(v, shape...) + case 3: + v := make([]bool, n) + for i := range v { + v[i] = data[(8+i)%len(data)]%2 == 1 + } + values, err = core.FromBools(v, shape...) + case 4: + v := make([]int8, n) + for i := range v { + v[i] = int8(int(data[(8+i)%len(data)]) - 128) + } + values, err = core.FromInt8s(v, shape...) + case 5: + v := make([]uint8, n) + for i := range v { + v[i] = data[(8+i)%len(data)] + } + values, err = core.FromUint8s(v, shape...) + case 6: + v := make([]int16, n) + for i := range v { + v[i] = int16(int(data[(8+i)%len(data)])*257 - 32768) + } + values, err = core.FromInt16s(v, shape...) + case 7: + v := make([]uint16, n) + for i := range v { + v[i] = uint16(int(data[(8+i)%len(data)]) * 257) + } + values, err = core.FromUint16s(v, shape...) + case 8: + v := make([]int32, n) + for i := range v { + v[i] = int32(int64(data[(8+i)%len(data)])*1000000 - 128000000 + int64(i)) + } + values, err = core.FromInt32s(v, shape...) + default: + v := make([]uint32, n) + for i := range v { + v[i] = uint32(uint64(data[(8+i)%len(data)])*1000000 + uint64(i)) + } + values, err = core.FromUint32s(v, shape...) + } + if err != nil { + t.Fatalf("build input: %v", err) + } + sets := []HDF5Dataset{{Path: "/d", Shape: shape, Values: values, Attrs: map[string]string{"k": "1"}}} + attrs := map[string]map[string]string{"/": {"title": "fuzz"}} + path := filepath.Join(t.TempDir(), "fuzz.h5") + if err := SaveHDF5(path, sets, attrs, HDF5WriteOptions{Gzip: gzip, Shuffle: shuffle, Latest: latest, ChunkBytes: chunk}); err != nil { + t.Fatalf("SaveHDF5 (%v): %v", shape, err) + } + back, err := LoadHDF5(path) + if err != nil { + t.Fatalf("LoadHDF5 round trip (%v): %v", shape, err) + } + if len(back) != 1 { + t.Fatalf("round trip returned %d datasets, want 1", len(back)) + } + d := back[0] + if d.Path != "/d" { + t.Fatalf("path = %q, want /d", d.Path) + } + if !slices.Equal(d.Shape, shape) { + t.Fatalf("shape = %v, want %v", d.Shape, shape) + } + if d.Values.Dtype() != values.Dtype() { + t.Fatalf("dtype = %s, want %s", d.Values.Dtype(), values.Dtype()) + } + switch dtype { + case 0: + got, want := d.Values.RawFloats()[:n], values.RawFloats()[:n] + for i := range want { + if math.Float64bits(got[i]) != math.Float64bits(want[i]) { + t.Fatalf("[%d] = %v, want %v (bit-exact)", i, got[i], want[i]) + } + } + case 1: + got, want := d.Values.RawFloat32s()[:n], values.RawFloat32s()[:n] + for i := range want { + if math.Float32bits(got[i]) != math.Float32bits(want[i]) { + t.Fatalf("[%d] = %v, want %v (bit-exact)", i, got[i], want[i]) + } + } + case 2: + got, want := d.Values.RawInts()[:n], values.RawInts()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %d, want %d", i, got[i], want[i]) + } + } + case 3: + got, want := d.Values.RawBools()[:n], values.RawBools()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %v, want %v", i, got[i], want[i]) + } + } + case 4: + got, want := d.Values.RawInt8s()[:n], values.RawInt8s()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %d, want %d", i, got[i], want[i]) + } + } + case 5: + got, want := d.Values.RawUint8s()[:n], values.RawUint8s()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %d, want %d", i, got[i], want[i]) + } + } + case 6: + got, want := d.Values.RawInt16s()[:n], values.RawInt16s()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %d, want %d", i, got[i], want[i]) + } + } + case 7: + got, want := d.Values.RawUint16s()[:n], values.RawUint16s()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %d, want %d", i, got[i], want[i]) + } + } + case 8: + got, want := d.Values.RawInt32s()[:n], values.RawInt32s()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %d, want %d", i, got[i], want[i]) + } + } + default: + got, want := d.Values.RawUint32s()[:n], values.RawUint32s()[:n] + for i := range want { + if got[i] != want[i] { + t.Fatalf("[%d] = %d, want %d", i, got[i], want[i]) + } + } + } + if d.Attrs["k"] != "1" { + t.Fatalf("dataset attr k = %q, want 1", d.Attrs["k"]) + } + if d.Attrs["title"] != "fuzz" { + t.Fatalf("merged root attr title = %q, want fuzz", d.Attrs["title"]) + } + }) +} diff --git a/io/hdf5.go b/io/hdf5.go new file mode 100644 index 0000000..8683953 --- /dev/null +++ b/io/hdf5.go @@ -0,0 +1,2303 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "bytes" + "compress/zlib" + "encoding/binary" + "io" + "maps" + "math" + "math/bits" + "os" + "slices" + "strconv" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// HDF5 (HDF5 1.8/1.10 file format), read-only, for the numeric arrays +// scientific files carry: every dataset of a file comes back as an +// Array with its path, shape and attributes. +// +// Supported: superblocks 0 to 3 (the classic layout and the "latest" +// library version, whose checksums are the lookup3 sum and are +// verified), object headers versions 1 and 2, groups stored as symbol +// tables or as link messages, datasets stored contiguously, compactly +// or in chunks through a version 1 B-tree, the deflate, shuffle and +// fletcher32 filters, fixed-point and floating-point datatypes of the +// usual widths, the boolean enumeration convention, and attributes in +// the object header. Fixed-point datasets land the core dtype their +// stored width and signedness declare; floating-point datasets land +// float64 at width 8 and float32 at width 4. +// +// Refused with an error naming what is missing, never guessed at: the +// superblock extension, the fractal-heap link storage of dense groups, +// the version 2 B-tree a latest-version file uses to index chunks, +// variable-length, compound, array and every non-boolean enumeration +// datatype, bit fields, string datasets, and the szip, nbit and +// scale-offset filters. + +// Superblock and object header signatures. +var ( + hdf5Magic = []byte{0x89, 'H', 'D', 'F', '\r', '\n', 0x1a, '\n'} + hdf5Tree = []byte{'T', 'R', 'E', 'E'} + hdf5SymbolNode = []byte{'S', 'N', 'O', 'D'} + hdf5LocalHeap = []byte{'H', 'E', 'A', 'P'} + hdf5ObjHdr2 = []byte{'O', 'H', 'D', 'R'} + hdf5Chunk2 = []byte{'O', 'C', 'H', 'K'} +) + +// hdf5MaxDatasetBytes bounds the raw byte extent of one dataset and of +// one chunk. The reader materialises whole datasets, so a declared size +// beyond this budget is a hostile header, not data: without the cap a +// 16-byte file could order a terabyte-scale allocation through the +// dataspace, the fill-value path or a single chunk declaration. Honest +// datasets above the budget belong behind mmap, not behind an eager +// read. +const hdf5MaxDatasetBytes = 2 << 30 // 2 GiB + +// hdf5ByteExtent multiplies declared extents by an element width into a +// byte count, bounding every factor against budget *before* it is +// multiplied. Multiplying first and comparing the product afterwards is +// the defect the review found: a dataspace of 2^32 by 2^32 float64 +// elements is 2^64 bytes, which wraps to zero, and a cap that runs after +// the multiplication therefore accepts a declaration the file cannot +// back, while a chunk of 2^32 by 2^29 elements looks empty the same way. +// Every partial product here stays inside budget, so no caller can be +// handed a wrapped count either. A zero extent is legal and yields zero +// bytes (an empty dataset); a negative one is a hostile header, as is a +// non-positive element width. +func hdf5ByteExtent(dims []int, width int, budget uint64) (uint64, error) { + total := uint64(1) + empty := false + for _, d := range dims { + if d < 0 { + return 0, base.Errf("a declared extent of %d", d) + } + if d == 0 { + empty = true + continue + } + if uint64(d) > budget/total { + return 0, base.Errf("the declared extents multiply past the %d byte budget", budget) + } + total *= uint64(d) + } + if empty { + return 0, nil + } + if width <= 0 { + return 0, base.Errf("a declared element width of %d bytes", width) + } + if total > budget/uint64(width) { + return 0, base.Errf("the declared extents multiply past the %d byte budget", budget) + } + return total * uint64(width), nil +} + +// HDF5 dataset object header message types. +const ( + hdf5MsgDataspace = 1 + hdf5MsgLinkInfo = 2 + hdf5MsgDatatype = 3 + hdf5MsgFillValueOld = 4 + hdf5MsgFillValue = 5 + hdf5MsgLink = 6 + hdf5MsgDataLayout = 8 + hdf5MsgGroupInfo = 10 + hdf5MsgFilterPipeline = 11 + hdf5MsgAttribute = 12 + hdf5MsgContinuation = 16 + hdf5MsgSymbolTable = 17 +) + +// HDF5 filter identifiers. +const ( + hdf5FilterDeflate = 1 + hdf5FilterShuffle = 2 + hdf5FilterFletcher32 = 3 +) + +// HDF5Dataset is one dataset of a file: its path from the root, its +// shape (row-major, as the file stores it), its values and its +// attributes. The attribute map also carries the attributes of the +// groups the dataset sits in, the nearest one winning, because that is +// where files usually put the units and titles that apply to a whole +// group. +type HDF5Dataset struct { + Path string + Shape []int + Values *core.Array + Attrs map[string]string +} + +// LoadHDF5 reads every dataset of an HDF5 file, in path order, as +// numeric arrays. The values keep the file's own dtype where the core +// has one: fixed-point data lands by stored width and signedness +// (int8, uint8, int16, uint16, int32, uint32, and int64 as Int), the +// boolean enumeration convention lands Bool, and floating-point data +// lands float64 at width 8 and float32 at width 4. Unsigned 64-bit +// data is refused, because no core dtype holds every value of it. +// +// An object hard-linked under several names is read once: its datasets +// appear under the first path the traversal reaches, the links of a +// group being visited in their file order, and never under the later +// ones. A link that closes a cycle along the path to it is an error, +// because no linear listing of the file exists. +func LoadHDF5(path string) ([]HDF5Dataset, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, base.Errf("LoadHDF5: %w", err) + } + f, err := newHDF5File(data) + if err != nil { + return nil, err + } + out := []HDF5Dataset{} + if err := f.walk(f.rootAddress, newHDF5WalkState(), &out); err != nil { + return nil, err + } + slices.SortFunc(out, func(x, y HDF5Dataset) int { return strings.Compare(x.Path, y.Path) }) + return out, nil +} + +// hdf5File is a parsed superblock plus the raw bytes of the file every +// offset is resolved against. +type hdf5File struct { + data []byte + superblock byte + offSize int + lenSize int + rootAddress uint64 + groupLeafK int + groupInnerK int +} + +func newHDF5File(data []byte) (*hdf5File, error) { + const name = "LoadHDF5" + if len(data) < 96 || !bytes.Equal(data[:8], hdf5Magic) { + return nil, base.Errf("%s: not an HDF5 file, the signature is % x", name, data[:min(8, len(data))]) + } + version := data[8] + f := &hdf5File{data: data, superblock: version} + switch version { + case 0, 1: + // The classic layout: the sizes and the root group entry follow + // the fixed header. + case 2, 3: + // The modern layout: the sizes follow the version byte, the + // four addresses follow without the group K values, and the + // root group is named by its object header address. The whole + // superblock carries a lookup3 checksum. + f.offSize = int(data[9]) + f.lenSize = int(data[10]) + switch f.offSize { + case 4, 8: + default: + return nil, base.Errf("%s: %d-byte addresses are not supported", name, f.offSize) + } + if f.lenSize != f.offSize && f.lenSize != 8 && f.lenSize != 4 { + return nil, base.Errf("%s: %d-byte lengths are not supported", name, f.lenSize) + } + stored := 12 + 4*f.offSize + 4 + if len(data) < stored { + return nil, base.Errf("%s: superblock version %d is truncated", name, version) + } + // Every address the format carries is absolute, so a non-zero + // base would need base-relative resolution, which is not + // implemented; the extension holds the file space strategy and + // driver settings, which the reader does not interpret. Both + // are refused by name rather than read wrong. + if baseAddr := f.u64At(12, f.offSize); baseAddr != 0 { + return nil, base.Errf("%s: a non-zero base address of %d is not supported", name, baseAddr) + } + if ext := f.u64At(12+f.offSize, f.offSize); ext != math.MaxUint64 { + return nil, base.Errf("%s: the superblock extension at %d is not supported; rewrite the file with the default library version", name, ext) + } + if got, want := binary.LittleEndian.Uint32(data[12+4*f.offSize:]), hdf5Lookup3(data[:12+4*f.offSize]); got != want { + return nil, base.Errf("%s: the superblock fails its checksum (%#08x, want %#08x)", name, got, want) + } + f.rootAddress = f.u64At(12+3*f.offSize, f.offSize) + if f.rootAddress == math.MaxUint64 { + return nil, base.Errf("%s: the root group has no object header", name) + } + return f, nil + default: + return nil, base.Errf("%s: superblock version %d is not supported; rewrite the file with the default version", name, version) + } + f.offSize = int(data[13]) + f.lenSize = int(data[14]) + switch f.offSize { + case 4, 8: + default: + return nil, base.Errf("%s: %d-byte addresses are not supported", name, f.offSize) + } + if f.lenSize != f.offSize && f.lenSize != 8 && f.lenSize != 4 { + return nil, base.Errf("%s: %d-byte lengths are not supported", name, f.lenSize) + } + f.groupLeafK = int(binary.LittleEndian.Uint16(data[16:])) + f.groupInnerK = int(binary.LittleEndian.Uint16(data[18:])) + pos := 24 + if version == 1 { + // Version 1 carries the indexed-storage K before the addresses. + pos = 28 + } + // base address (every file this reader accepts carries 0 there: + // a non-zero base would require base-relative offset resolution, + // which is not implemented), free space information, end of file, + // driver information: four addresses, then the root group's symbol + // table entry, whose first field is the link name offset and is + // sized by the length size, not the address size. + pos += 4 * f.offSize + f.rootAddress = f.u64At(pos+f.lenSize, f.offSize) + if f.rootAddress == math.MaxUint64 { + return nil, base.Errf("%s: the root group has no object header", name) + } + return f, nil +} + +// u64At reads an address or length of size bytes at an absolute file +// offset; every offset the reader resolves is absolute. +func (f *hdf5File) u64At(off int, size int) uint64 { + if off < 0 || off+size > len(f.data) { + return math.MaxUint64 + } + switch size { + case 4: + return uint64(binary.LittleEndian.Uint32(f.data[off:])) + case 8: + return binary.LittleEndian.Uint64(f.data[off:]) + } + return math.MaxUint64 +} + +// at reads an address field of the file's address size at off. +func (f *hdf5File) at(off int) uint64 { return f.u64At(off, f.offSize) } + +// len reads a length field of the file's length size at off. +func (f *hdf5File) length(off int) uint64 { return f.u64At(off, f.lenSize) } + +// bytes returns the file bytes of the given region, or nil when the +// region lies outside the file. +func (f *hdf5File) bytes(addr uint64, n uint64) []byte { + if addr == math.MaxUint64 || n > uint64(len(f.data)) { + return nil + } + start := addr + if start > uint64(len(f.data)) || n > uint64(len(f.data))-start { + return nil + } + return f.data[start : start+n] +} + +// hdf5Message is one header message with its raw payload. +type hdf5Message struct { + typ uint16 + data []byte +} + +// messages reads an object header's messages, following continuation +// blocks. Versions 1 and 2 are supported: version 1 is what the +// library writes by default, version 2 what the "latest" library +// version writes. +func (f *hdf5File) messages(addr uint64) ([]hdf5Message, error) { + const name = "LoadHDF5" + head := f.bytes(addr, 4) + if head == nil { + return nil, base.Errf("%s: object header at %d lies outside the file", name, addr) + } + if bytes.Equal(head, hdf5ObjHdr2) { + return f.messagesV2(addr) + } + start := f.bytes(addr, 16) + if start == nil { + return nil, base.Errf("%s: object header at %d lies outside the file", name, addr) + } + if start[0] != 1 { + return nil, base.Errf("%s: object header version %d is not supported; rewrite the file with the default library version", name, start[0]) + } + // version(1), reserved(1), messages(2), references(4), data size(4), + // padding(4), then the message data itself. + nmsg := int(binary.LittleEndian.Uint16(start[2:])) + dataSize := int(binary.LittleEndian.Uint32(start[8:])) + pos := addr + 16 + region := f.bytes(pos, uint64(dataSize)) + if region == nil { + return nil, base.Errf("%s: object header at %d is truncated", name, addr) + } + // The declared count sizes the preallocation, so it is capped by the + // region that would have to hold the messages: a header claiming + // 65535 messages costs 16 bytes in the file, and every message is at + // least eight bytes, so the region bounds the count that can follow. + out := make([]hdf5Message, 0, min(nmsg, len(region)/8+1)) + p := 0 + for i := range nmsg { + if p+8 > len(region) { + return nil, base.Errf("%s: object header at %d ends inside a message", name, addr) + } + typ := binary.LittleEndian.Uint16(region[p:]) + size := int(binary.LittleEndian.Uint16(region[p+2:])) + if p+8+size > len(region) { + return nil, base.Errf("%s: object header at %d ends inside message %d", name, addr, i) + } + body := region[p+8 : p+8+size] + if typ == hdf5MsgContinuation { + // offset, length: another block of messages. The declared + // count includes the messages the chain carries, so the + // link is the last message physically present in this + // block; every further link is followed inside messagesIn, + // whose visited set keeps a hostile self-referencing block + // from looping. The link body must hold an address and a + // length outright: reading the fields from whatever follows + // a short message would guess at bytes the message never + // carried. + if size < f.offSize+f.lenSize { + return nil, base.Errf("%s: object header at %d carries a continuation message shorter than an offset and a length", name, addr) + } + blockAddr := f.at(int(pos) + p + 8) + next := f.bytes(blockAddr, f.length(int(pos)+p+8+f.offSize)) + if next == nil { + return nil, base.Errf("%s: object header continuation at %d lies outside the file", name, addr) + } + cont, err := f.messagesIn(next, blockAddr, map[uint64]bool{addr: true}, 0) + if err != nil { + return nil, err + } + out = append(out, cont...) + break + } + out = append(out, hdf5Message{typ: typ, data: body}) + p += 8 + size + p = alignUp(p, 8) + } + return out, nil +} + +// messagesV2 reads a version 2 object header: the OHDR signature, a +// version and a flag byte, the flag-selected prefix fields, the size +// of the first chunk, then the messages and a lookup3 checksum over +// everything from the signature to the end of the chunk. The format +// inserts no alignment into the stream, so the walk reads the fields +// packed as they lie. +func (f *hdf5File) messagesV2(addr uint64) ([]hdf5Message, error) { + const name = "LoadHDF5" + head := f.bytes(addr, 6) + if head == nil { + return nil, base.Errf("%s: object header at %d lies outside the file", name, addr) + } + if head[4] != 2 { + return nil, base.Errf("%s: object header version %d is not supported; rewrite the file with the default library version", name, head[4]) + } + flags := head[5] + if flags&^byte(0x3f) != 0 { + return nil, base.Errf("%s: the object header at %d carries unknown status flags %#02x", name, addr, flags) + } + p := addr + 6 + if flags&0x20 != 0 { + p += 16 // access, modification, change and birth times + } + if flags&0x10 != 0 { + p += 4 // the max compact and min dense attribute counts + } + width := 1 << (flags & 0x03) + // The checksum of four bytes follows the chunk directly, so the + // size field, the chunk and the checksum must all lie inside the + // file; the bound is settled before the size is read, so a hostile + // header cannot point the read past the end. + if p+uint64(width)+4 > uint64(len(f.data)) { + return nil, base.Errf("%s: object header at %d is truncated", name, addr) + } + chunk0 := leN(f.data[p:], width) + if chunk0 > uint64(len(f.data))-p-uint64(width)-4 { + return nil, base.Errf("%s: object header at %d is truncated", name, addr) + } + start := p + uint64(width) + end := start + chunk0 + if got, want := binary.LittleEndian.Uint32(f.data[end:]), hdf5Lookup3(f.data[addr:end]); got != want { + return nil, base.Errf("%s: the object header at %d fails its checksum (%#08x, want %#08x)", name, addr, got, want) + } + visited := map[uint64]bool{addr: true} + return f.messagesV2Stream(f.data[start:end], addr, flags&0x04 != 0, visited, 0) +} + +// messagesV2Stream walks one version 2 message region: each message is +// one type byte, a two-byte size and a flag byte, widened by a +// two-byte creation order when the header tracks creation order. A +// continuation message hands the walk to a further OCHK block, whose +// checksum covers signature and messages alike; visited refuses a +// block reached twice and depth an unbounded chain. A writer may leave +// a gap of up to three bytes before the region's checksum; anything +// wider would be a message the region cannot hold. +func (f *hdf5File) messagesV2Stream(region []byte, addr uint64, ordered bool, visited map[uint64]bool, depth int) ([]hdf5Message, error) { + const name = "LoadHDF5" + if depth > hdf5MaxHeaderBlocks { + return nil, base.Errf("%s: the object header at %d chains through more than %d continuation blocks", name, addr, hdf5MaxHeaderBlocks) + } + out := []hdf5Message{} + p := 0 + for p+4 <= len(region) { + typ := region[p] + size := int(binary.LittleEndian.Uint16(region[p+1:])) + p += 4 + if ordered { + if p+2 > len(region) { + return nil, base.Errf("%s: object header at %d ends inside a message", name, addr) + } + p += 2 + } + if p+size > len(region) { + return nil, base.Errf("%s: object header at %d ends inside a message", name, addr) + } + if typ == hdf5MsgContinuation { + // The link body is an offset of the address size followed by + // a length of the length size; a block that ends inside the + // link has neither, so the read is refused instead of taken + // from bytes past the block. + if p+f.offSize+f.lenSize > len(region) { + return nil, base.Errf("%s: object header at %d carries a continuation message shorter than an offset and a length", name, addr) + } + blockAddr := leN(region[p:], f.offSize) + block := f.bytes(blockAddr, leN(region[p+f.offSize:], f.lenSize)) + if block == nil || len(block) < 8 || !bytes.Equal(block[:4], hdf5Chunk2) { + return nil, base.Errf("%s: no version 2 continuation block at %d", name, blockAddr) + } + if got, want := binary.LittleEndian.Uint32(block[len(block)-4:]), hdf5Lookup3(block[:len(block)-4]); got != want { + return nil, base.Errf("%s: the continuation block at %d fails its checksum (%#08x, want %#08x)", name, blockAddr, got, want) + } + if visited[blockAddr] { + return nil, base.Errf("%s: object header continuation block at %d is reachable twice", name, blockAddr) + } + visited[blockAddr] = true + cont, err := f.messagesV2Stream(block[4:len(block)-4], blockAddr, ordered, visited, depth+1) + if err != nil { + return nil, err + } + out = append(out, cont...) + } else if typ != 0 { + // A null message is the residue of a deletion and names no + // content; the live messages carry on behind it. + out = append(out, hdf5Message{typ: uint16(typ), data: region[p : p+size]}) + } + p += size + } + if gap := len(region) - p; gap >= 4 { + return nil, base.Errf("%s: object header at %d ends inside a message", name, addr) + } + return out, nil +} + +// messagesIn reads the messages of one continuation block, whose size +// is the block itself. A continuation message inside the block chains +// to a further block; the visited set holds every block address already +// read, so a cycle is an error instead of an infinite walk, and depth +// bounds the chain the same way the version 2 walk bounds its own. +func (f *hdf5File) messagesIn(block []byte, addr uint64, visited map[uint64]bool, depth int) ([]hdf5Message, error) { + const name = "LoadHDF5" + if visited[addr] { + return nil, base.Errf("%s: object header continuation block at %d is reachable twice", name, addr) + } + visited[addr] = true + if depth > hdf5MaxHeaderBlocks { + return nil, base.Errf("%s: the object header at %d chains through more than %d continuation blocks", name, addr, hdf5MaxHeaderBlocks) + } + out := []hdf5Message{} + p := 0 + for p+8 <= len(block) { + typ := binary.LittleEndian.Uint16(block[p:]) + size := int(binary.LittleEndian.Uint16(block[p+2:])) + if p+8+size > len(block) { + return nil, base.Errf("%s: object header continuation block ends inside a message", name) + } + if typ == hdf5MsgContinuation { + // The link body is an offset of the address size followed + // by a length of the length size; a block that ends inside + // the link has neither, so the read is refused instead of + // taken from bytes past the block. + if p+8+f.offSize+f.lenSize > len(block) { + return nil, base.Errf("%s: object header continuation block at %d carries a continuation message shorter than an offset and a length", name, addr) + } + blockAddr := leN(block[p+8:], f.offSize) + next := f.bytes(blockAddr, leN(block[p+8+f.offSize:], f.lenSize)) + if next == nil { + return nil, base.Errf("%s: object header continuation block at %d lies outside the file", name, blockAddr) + } + cont, err := f.messagesIn(next, blockAddr, visited, depth+1) + if err != nil { + return nil, err + } + out = append(out, cont...) + p = alignUp(p+8+size, 8) + continue + } + out = append(out, hdf5Message{typ: typ, data: block[p+8 : p+8+size]}) + p = alignUp(p+8+size, 8) + } + return out, nil +} + +// leN reads a little-endian integer of n bytes from the head of b. +func leN(b []byte, n int) uint64 { + var v uint64 + for i := 0; i < n && i < len(b); i++ { + v |= uint64(b[i]) << (8 * i) + } + return v +} + +func alignUp(n, to int) int { return (n + to - 1) / to * to } + +// hdf5MaxGroupDepth bounds how deeply groups may nest. It is a stack +// guard for the walk, not a memory guard: the traversal carries one path +// buffer and one attribute map for the whole file, so the memory a deep +// file costs grows with its depth rather than with its square. Real +// files nest a handful of levels deep; the cap only refuses the extreme. +const hdf5MaxGroupDepth = 512 + +// hdf5MaxHeaderBlocks bounds how many continuation blocks one object +// header may chain through. It is a stack guard for the version 2 +// message walk, whose recursion descends one frame per block: real +// headers split over a handful of blocks, the cap only refuses the +// chain a hostile file builds to exhaust the stack. +const hdf5MaxHeaderBlocks = 512 + +// hdf5DatasetEnvelope is the fixed per-dataset charge the walk deducts +// from the read budget beyond the value bytes themselves: the result's +// struct, shape slice, path string and attribute map cost roughly the +// same for every dataset, however small its values are, and a file of +// millions of empty datasets must pay for the envelopes it orders. +// The envelope is a conservative upper estimate of that structure, not +// a measurement of it, and it is charged together with the dataset's +// path length, which the result copies out of the walk's buffer. +const hdf5DatasetEnvelope = 512 + +// hdf5WalkState is the traversal state shared by the whole walk: the +// object headers currently on the path, so a hard-link cycle is refused +// instead of recursed, the object headers already walked from any path, +// so a diamond of hard links is read once instead of once per path, the +// inherited attribute map, which every level adds to and then undoes, +// and the path buffer, which every level extends and then truncates. +// All are shared rather than copied per level: a copy per level makes +// the live memory grow with the square of the nesting, which is the +// hostile-file blow-up the review found. +type hdf5WalkState struct { + onPath map[uint64]bool + seen map[uint64]bool + attrs map[string]string + path []byte + depth int + // budget is the byte extent the walk may still hand out, the + // whole-file aggregate behind the per-dataset cap: a file whose + // datasets each fit the cap can still declare thousands of them, + // and the sum must refuse, not the machine. Every dataset pays its + // value bytes plus hdf5DatasetEnvelope and its path. + budget int64 +} + +func newHDF5WalkState() *hdf5WalkState { + return &hdf5WalkState{onPath: map[uint64]bool{}, seen: map[uint64]bool{}, attrs: map[string]string{}, path: []byte{'/'}, budget: hdf5MaxDatasetBytes} +} + +// hdf5AttrUndo records one attribute this level overwrote, so leaving +// the level restores the outer value instead of dropping it. +type hdf5AttrUndo struct { + key string + prev string + had bool +} + +// walk visits the group whose path is already in st.path and everything +// under it, appending the datasets. The attributes of a group are +// inherited by everything below it, the nearest group winning, so a +// dataset carries the units and titles its file hangs on the enclosing +// groups. +// +// An object reachable through several hard links (a diamond) is walked +// once, under the first path the traversal reaches it from: st.seen +// records every object header already walked and a second arrival +// skips. The first path wins deterministically because the links of a +// group are visited in the order their file stores them. An object that +// closes a cycle along the current path is different: it is refused +// with an error, because no order of visits can read a cycle at all. +func (f *hdf5File) walk(addr uint64, st *hdf5WalkState, out *[]HDF5Dataset) error { + const name = "LoadHDF5" + if st.onPath[addr] { + return base.Errf("%s: the group at object header %d is linked into itself along %q; a hard-link cycle cannot be walked", name, addr, string(st.path)) + } + if st.seen[addr] { + // Already walked from another path: its datasets are in out + // under that, earlier, path and the file's content is read. + return nil + } + if st.depth >= hdf5MaxGroupDepth { + return base.Errf("%s: the groups nest deeper than %d levels", name, hdf5MaxGroupDepth) + } + st.onPath[addr] = true + st.seen[addr] = true + st.depth++ + defer func() { + delete(st.onPath, addr) + st.depth-- + }() + msgs, err := f.messages(addr) + if err != nil { + return err + } + undo := make([]hdf5AttrUndo, 0, 4) + for k, v := range f.attributesOf(msgs) { + prev, had := st.attrs[k] + undo = append(undo, hdf5AttrUndo{key: k, prev: prev, had: had}) + st.attrs[k] = v + } + defer func() { + for _, u := range slices.Backward(undo) { + if u.had { + st.attrs[u.key] = u.prev + continue + } + delete(st.attrs, u.key) + } + }() + links, err := f.links(msgs) + if err != nil { + return err + } + var datasetMsgs []hdf5Message + for _, m := range msgs { + switch m.typ { + case hdf5MsgDataspace, hdf5MsgDatatype, hdf5MsgDataLayout, hdf5MsgFilterPipeline, hdf5MsgFillValue: + datasetMsgs = append(datasetMsgs, m) + } + } + // A group has no dataspace; a dataset does. + hasSpace := false + for _, m := range datasetMsgs { + if m.typ == hdf5MsgDataspace { + hasSpace = true + } + } + if hasSpace { + ds, err := f.dataset(string(st.path), datasetMsgs, st) + if err != nil { + return err + } + // The dataset owns its attributes: a clone, because the walk's + // map keeps changing for the levels below this one. + ds.Attrs = maps.Clone(st.attrs) + *out = append(*out, ds) + return nil + } + for _, l := range links { + mark := len(st.path) + if st.path[mark-1] != '/' { + st.path = append(st.path, '/') + } + st.path = append(st.path, l.name...) + err := f.walk(l.address, st, out) + st.path = st.path[:mark] + if err != nil { + return err + } + } + return nil +} + +// hdf5Link is one named hard link of a group. +type hdf5Link struct { + name string + address uint64 +} + +// links collects a group's hard links from its link messages and, when +// present, its symbol table. +func (f *hdf5File) links(msgs []hdf5Message) ([]hdf5Link, error) { + out := []hdf5Link{} + for _, m := range msgs { + switch m.typ { + case hdf5MsgLink: + l, err := f.decodeLink(m.data) + if err != nil { + return nil, err + } + if l.address != math.MaxUint64 { + out = append(out, l) + } + case hdf5MsgSymbolTable: + st, err := f.symbolTableLinks(m.data) + if err != nil { + return nil, err + } + out = append(out, st...) + case hdf5MsgLinkInfo: + // Dense storage would need the fractal heap; a group whose + // links all live in the header carries the info message + // without using the dense structures. + } + } + if len(out) == 0 { + for _, m := range msgs { + if m.typ == hdf5MsgLinkInfo { + const name = "LoadHDF5" + return nil, base.Errf("%s: this group stores its links in a dense (fractal heap) structure, which is not supported", name) + } + } + } + return out, nil +} + +// decodeLink parses one link message, ignoring soft and external links +// (they name no object in this file). +func (f *hdf5File) decodeLink(m []byte) (hdf5Link, error) { + const name = "LoadHDF5" + if len(m) < 2 { + return hdf5Link{}, base.Errf("%s: link message is truncated", name) + } + if m[0] != 1 { + return hdf5Link{}, base.Errf("%s: link message version %d is not supported", name, m[0]) + } + flags := m[1] + p := 2 + linkType := uint8(0) + if flags&0x08 != 0 { + if p >= len(m) { + return hdf5Link{}, base.Errf("%s: link message is truncated", name) + } + linkType = m[p] + p++ + } + if flags&0x04 != 0 { + if p+8 > len(m) { + return hdf5Link{}, base.Errf("%s: link message is truncated", name) + } + p += 8 // creation order + } + nameLenSize := 1 << (flags & 0x03) // 1, 2, 4 or 8 bytes + nameLen, ok := readUintLE(m[p:], nameLenSize) + if !ok { + return hdf5Link{}, base.Errf("%s: link message is truncated", name) + } + p += nameLenSize + if uint64(p)+nameLen > uint64(len(m)) { + return hdf5Link{}, base.Errf("%s: link message is truncated", name) + } + linkName := string(m[p : p+int(nameLen)]) + p += int(nameLen) + if linkType != 0 { + // A soft or external link: skip it, it names no object here. + return hdf5Link{name: linkName, address: math.MaxUint64}, nil + } + addr := f.at2(m[p:]) + return hdf5Link{name: linkName, address: addr}, nil +} + +// at2 reads an address of the file's address size from a message +// payload, or returns the undefined address. +func (f *hdf5File) at2(b []byte) uint64 { + if len(b) < f.offSize { + return math.MaxUint64 + } + switch f.offSize { + case 4: + return uint64(binary.LittleEndian.Uint32(b)) + case 8: + return binary.LittleEndian.Uint64(b) + } + return math.MaxUint64 +} + +func readUintLE(b []byte, size int) (uint64, bool) { + if len(b) < size { + return 0, false + } + switch size { + case 1: + return uint64(b[0]), true + case 2: + return uint64(binary.LittleEndian.Uint16(b)), true + case 4: + return uint64(binary.LittleEndian.Uint32(b)), true + case 8: + return binary.LittleEndian.Uint64(b), true + } + return 0, false +} + +// symbolTableLinks walks the group's version 1 B-tree of symbol table +// nodes and reads the names out of the local heap. +func (f *hdf5File) symbolTableLinks(m []byte) ([]hdf5Link, error) { + const name = "LoadHDF5" + if len(m) < 2*f.offSize { + return nil, base.Errf("%s: symbol table message is truncated", name) + } + treeAddr := f.at2(m) + heapAddr := f.at2(m[f.offSize:]) + heap, err := f.localHeap(heapAddr) + if err != nil { + return nil, err + } + out := []hdf5Link{} + set := map[uint64]bool{} + if err := f.treeLinks(treeAddr, heap, set, &out, 0); err != nil { + return nil, err + } + return out, nil +} + +// treeLinks descends a group B-tree, reading a symbol table node at +// every leaf. +func (f *hdf5File) treeLinks(addr uint64, heap []byte, seen map[uint64]bool, out *[]hdf5Link, depth int) error { + const name = "LoadHDF5" + if depth > 32 { + return base.Errf("%s: the group B-tree is more than 32 levels deep", name) + } + if addr == math.MaxUint64 || seen[addr] { + return nil + } + seen[addr] = true + // The version 1 B-tree header is the signature, the type, the level + // and the entry count (eight bytes) plus a left and a right sibling + // address, so it is 8+2*offSize wide, not the 24 an 8-byte-address + // file happens to make it. Keying the entries off a pinned 24 read + // a 4/4 file's first key as a sibling address and lost the node. + node := f.bytes(addr, uint64(8+2*f.offSize)) + if node == nil || !bytes.Equal(node[:4], hdf5Tree) { + return base.Errf("%s: no B-tree node at %d", name, addr) + } + nodeType := node[4] + level := node[5] + entries := int(binary.LittleEndian.Uint16(node[6:])) + if nodeType != 0 { + return base.Errf("%s: B-tree node type %d is not a group", name, nodeType) + } + // Each entry is a key of one address size followed by a child + // address; the node ends with one trailing key. + entrySize := f.lenSize + f.offSize + base := int(addr) + 8 + 2*f.offSize + for i := range entries { + child := f.at(base + i*entrySize + f.lenSize) + if level == 0 { + if err := f.symbolNodeLinks(child, heap, seen, out); err != nil { + return err + } + continue + } + if err := f.treeLinks(child, heap, seen, out, depth+1); err != nil { + return err + } + } + return nil +} + +// symbolNodeLinks reads one symbol table node's entries. +func (f *hdf5File) symbolNodeLinks(addr uint64, heap []byte, seen map[uint64]bool, out *[]hdf5Link) error { + const name = "LoadHDF5" + if addr == math.MaxUint64 || seen[addr] { + return nil + } + seen[addr] = true + node := f.bytes(addr, 8) + if node == nil || !bytes.Equal(node[:4], hdf5SymbolNode) { + return base.Errf("%s: no symbol table node at %d", name, addr) + } + count := int(binary.LittleEndian.Uint16(node[6:])) + // An entry is the heap offset (of the length size), the object + // address (of the address size), the cache type, a reserved word + // and the 16-byte scratch pad: lenSize+offSize+24, not the 40 an + // 8/8 file happens to make it. A pinned 40 walked a 4/4 node's + // entries into the middle of the first one and read its second + // link's name from the wrong heap offset. + entrySize := f.lenSize + f.offSize + 24 + start := int(addr) + 8 + for i := range count { + off := start + i*entrySize + nameOff := f.length(off) + objAddr := f.at(off + f.lenSize) + linkName, err := heapString(heap, nameOff) + if err != nil { + return err + } + *out = append(*out, hdf5Link{name: linkName, address: objAddr}) + } + return nil +} + +// localHeap reads a local heap's data segment, the pool of +// null-terminated names the symbol table entries point into. +func (f *hdf5File) localHeap(addr uint64) ([]byte, error) { + const name = "LoadHDF5" + // The header is the signature and version, the data segment size + // and the free-list head (both of the length size), then the data + // segment address of the address size. Sizing the header with the + // address size instead (8+3*offSize) sliced the address out of a + // header that a 4/8 file stores in 20 bytes at [24:], which + // panicked; every field here is sized by its own kind. + head := f.bytes(addr, uint64(8+2*f.lenSize+f.offSize)) + if head == nil || !bytes.Equal(head[:4], hdf5LocalHeap) { + return nil, base.Errf("%s: no local heap at %d", name, addr) + } + dataAddr := f.at2(head[8+2*f.lenSize:]) + size := f.length(int(addr) + 8) // data segment size, then the free list head + seg := f.bytes(dataAddr, size) + if seg == nil { + return nil, base.Errf("%s: the local heap at %d lies outside the file", name, addr) + } + return seg, nil +} + +// heapString reads the null-terminated string at an offset in the heap. +func heapString(heap []byte, off uint64) (string, error) { + if off >= uint64(len(heap)) { + return "", base.Errf("LoadHDF5: a symbol table entry names a name at heap offset %d, past the %d-byte heap", + off, len(heap)) + } + end := bytes.IndexByte(heap[off:], 0) + if end < 0 { + return "", base.Errf("LoadHDF5: the name at heap offset %d is not terminated", off) + } + return string(heap[off : off+uint64(end)]), nil +} + +// attributesOf collects the attributes carried in an object header. +func (f *hdf5File) attributesOf(msgs []hdf5Message) map[string]string { + out := map[string]string{} + for _, m := range msgs { + if m.typ != hdf5MsgAttribute { + continue + } + linkName, value, ok := f.decodeAttribute(m.data) + if ok { + out[linkName] = value + } + } + return out +} + +// hdf5Type is the parsed datatype of a dataset or attribute: enough of +// it to read the values. +type hdf5Type struct { + class byte + size int + signed bool + width int // fixed-point: bytes + isFloat bool + isBool bool // a boolean enumeration: lands core.Bool + bitOrder byte +} + +// hdf5ClassName names a datatype class for a refusal message: the +// number alone tells whoever wrote the file little about what the +// reader missed. +func hdf5ClassName(class byte) string { + switch class { + case 2: + return "date and time" + case 4: + return "bit field" + case 5: + return "opaque" + case 6: + return "compound" + case 7: + return "reference" + case 8: + return "enumeration" + case 9: + return "variable-length" + } + return "unknown" +} + +// decodeType parses a datatype message (version 1, which is what the +// default library version writes). Strings are refused for datasets and +// accepted for attributes, where a string is exactly what is wanted. +// An enumeration is accepted only in the boolean convention HDF5 +// writers carry booleans in, and lands core.Bool; every other +// enumeration and every bit field is refused by name. +func decodeType(m []byte, allowStrings bool) (hdf5Type, error) { + const name = "LoadHDF5" + if len(m) < 8 { + return hdf5Type{}, base.Errf("%s: datatype message is truncated", name) + } + // The first byte carries the version in its high nibble and the + // datatype class in its low one. + classAndVersion := m[0] + version := classAndVersion >> 4 + class := classAndVersion & 0x0f + if version != 1 { + return hdf5Type{}, base.Errf("%s: datatype message version %d is not supported", name, version) + } + size := int(binary.LittleEndian.Uint32(m[4:])) + t := hdf5Type{class: class, size: size, bitOrder: m[1] >> 0} + // The class bit field's bit 0 is the byte order: 0 little-endian, + // anything else big-endian. Every element below is decoded as + // little-endian, so a big-endian datatype, on a dataset or on an + // attribute alike, would come back as byte-swapped noise with no + // error; it is refused by name instead. + if (class == 0 || class == 1) && m[1]&0x01 != 0 { + return hdf5Type{}, base.Errf("%s: big-endian datatypes are not supported; rewrite the file in the little-endian order", name) + } + switch class { + case 0: // fixed-point + if size != 1 && size != 2 && size != 4 && size != 8 { + return hdf5Type{}, base.Errf("%s: %d-byte fixed-point values are not supported", name, size) + } + t.signed = m[1]&0x08 != 0 + t.width = size + case 1: // floating-point + if len(m) < 20 { + return hdf5Type{}, base.Errf("%s: floating-point datatype message is truncated", name) + } + switch size { + case 4, 8: + default: + return hdf5Type{}, base.Errf("%s: %d-byte floating-point values are not supported", name, size) + } + t.isFloat = true + t.width = size + case 3: // string + if !allowStrings { + return hdf5Type{}, base.Errf("%s: string datasets are not supported (only numeric arrays are read)", name) + } + t.width = size + if t.width <= 0 { + return hdf5Type{}, base.Errf("%s: a string attribute declares %d bytes", name, size) + } + case 8: // enumeration: only the boolean convention is read + if err := hdf5EnumBool(m); err != nil { + return hdf5Type{}, err + } + t.isBool = true + t.width = size + case 9: // variable-length: the text lives in a global heap collection + if !allowStrings { + return hdf5Type{}, base.Errf("%s: variable-length datasets are not supported (only numeric arrays are read)", name) + } + t.class = 9 + t.width = size // the descriptor: length, heap address, object index + default: + return hdf5Type{}, base.Errf("%s: datatype class %d (%s) is not supported", name, class, hdf5ClassName(class)) + } + return t, nil +} + +// hdf5EnumBool verifies that an enumeration datatype message carries +// the boolean convention HDF5 writers store booleans in: a one-byte +// unsigned base type whose member values are a subset of {0, 1}. The +// member names are irrelevant to the values, which is what the landing +// keys on, and any other enumeration is refused by name. +// +// The message layout comes from the HDF5 file format specification's +// datatype message, enumeration class: the class bit field carries the +// member count in its low sixteen bits, and the properties hold the +// base type as a complete datatype message, then the member names +// (each a NUL-terminated string stored in a multiple of eight bytes, +// padded from its own field's start) and the packed member values. The +// base type here is a complete version 1 fixed-point message: its own +// eight-byte header plus the bit offset and bit precision the +// specification's fixed-point property table defines, twelve bytes in +// total, the same shape the reader already requires of the twenty-byte +// version 1 floating-point message. +func hdf5EnumBool(m []byte) error { + const name = "LoadHDF5" + members := int(binary.LittleEndian.Uint16(m[1:])) + if m[3] != 0 { + return base.Errf("%s: enumeration datatype class 8 carries unknown bit field bits %#02x", name, m[3]) + } + if members < 1 { + return base.Errf("%s: enumeration datatype class 8 declares %d members", name, members) + } + if size := binary.LittleEndian.Uint32(m[4:]); size != 1 { + return base.Errf("%s: enumeration datatype class 8 with %d-byte values is not supported; only the one-byte unsigned boolean convention is read", name, size) + } + // The base type: a complete version 1 fixed-point message. + if len(m) < 8+12 { + return base.Errf("%s: enumeration datatype message is truncated", name) + } + b := m[8:] + if bclass := b[0] & 0x0f; b[0]>>4 != 1 || bclass != 0 { + return base.Errf("%s: enumeration datatype class 8 with a base type of class %d version %d is not supported; only the one-byte unsigned boolean convention is read", + name, bclass, b[0]>>4) + } + if b[1]&0x01 != 0 { + return base.Errf("%s: big-endian datatypes are not supported; rewrite the file in the little-endian order", name) + } + if b[1]&0x08 != 0 { + return base.Errf("%s: enumeration datatype class 8 with a signed base type is not supported; only the one-byte unsigned boolean convention is read", name) + } + if bsize := binary.LittleEndian.Uint32(b[4:]); bsize != 1 { + return base.Errf("%s: enumeration datatype class 8 with a base type of %d bytes is not supported; only the one-byte unsigned boolean convention is read", name, bsize) + } + // The names: each field is a NUL-terminated string padded, from its + // own start, to a multiple of eight bytes. The walk only needs the + // fields' extent; the values follow them. + p := 8 + 12 + for i := range members { + end := -1 + for j := p; j < len(m); j++ { + if m[j] == 0 { + end = j + break + } + } + if end < 0 { + return base.Errf("%s: enumeration datatype message ends inside member name %d", name, i+1) + } + p += alignUp(end-p+1, 8) + if p > len(m) { + return base.Errf("%s: enumeration datatype message ends inside the member names", name) + } + } + // The values: packed, one base-width byte each. + if members > len(m)-p { + return base.Errf("%s: enumeration datatype message holds %d member values, fewer than its %d members", name, len(m)-p, members) + } + for i := range members { + if v := m[p+i]; v > 1 { + return base.Errf("%s: enumeration datatype class 8 declares member value %d, outside the boolean convention {0, 1}", name, v) + } + } + return nil +} + +// decodeDataspace parses a dataspace message into the dataset's shape. +// Version 1 is what the default library version writes, version 2 what +// the "latest" one writes: it drops the reserved bytes to one and +// always stores the dimensions in eight bytes. +func decodeDataspace(m []byte, lenSize int) ([]int, error) { + const name = "LoadHDF5" + if len(m) < 8 { + return nil, base.Errf("%s: dataspace message is truncated", name) + } + version := m[0] + rank := int(m[1]) + flags := m[2] + if version != 1 && version != 2 { + return nil, base.Errf("%s: dataspace message version %d is not supported", name, version) + } + if rank == 0 { + return []int{1}, nil // a scalar: one element, no dimensions + } + pos, dimSize := 8, lenSize + if version == 2 { + pos, dimSize = 4, 8 + } + shape := make([]int, rank) + for i := range rank { + if pos+dimSize > len(m) { + return nil, base.Errf("%s: dataspace message is truncated", name) + } + shape[i] = int(uint64At(m[pos:], dimSize)) + pos += dimSize + } + if flags&0x02 != 0 { + return nil, base.Errf("%s: datasets with a dimension permutation are not supported", name) + } + if flags&0x01 != 0 { + // Maximum dimensions are not the shape; nothing to do with them. + pos += rank * dimSize + } + return shape, nil +} + +func uint64At(b []byte, size int) uint64 { + v, _ := readUintLE(b, size) + return v +} + +// hdf5Layout is a dataset's storage description. For chunked storage +// the message records one more dimension than the dataset has: the +// dimensions are the chunk's shape followed by the element size in +// bytes, and the chunk B-tree keys carry the same extra offset. +type hdf5Layout struct { + class byte + addr uint64 + size uint64 + compact []byte + dims []int +} + +func decodeLayout(m []byte, offSize int) (hdf5Layout, error) { + const name = "LoadHDF5" + if len(m) < 2 { + return hdf5Layout{}, base.Errf("%s: data layout message is truncated", name) + } + version := m[0] + class := m[1] + switch version { + case 3, 4: + default: + return hdf5Layout{}, base.Errf("%s: data layout message version %d is not supported", name, version) + } + l := hdf5Layout{class: class} + switch class { + case 0: // compact: the data lives in the message + if len(m) < 4 { + return hdf5Layout{}, base.Errf("%s: compact layout message is truncated", name) + } + size := int(binary.LittleEndian.Uint16(m[2:])) + if 4+size > len(m) { + return hdf5Layout{}, base.Errf("%s: compact data runs past the message", name) + } + l.compact = m[4 : 4+size] + case 1: // contiguous + if len(m) < 2+offSize+8 { + return hdf5Layout{}, base.Errf("%s: contiguous layout message is truncated", name) + } + l.addr = uint64At(m[2:], offSize) + l.size = uint64At(m[2+offSize:], 8) + case 2: // chunked + if len(m) < 3+offSize { + return hdf5Layout{}, base.Errf("%s: chunked layout message is truncated", name) + } + rank := int(m[2]) + p := 3 + l.addr = uint64At(m[p:], offSize) + p += offSize + // Version 4 stores one flag byte before the chunk dimensions + // and widens them to 8 bytes when the flag asks for it. + dimSize := 4 + if version == 4 { + if p >= len(m) { + return hdf5Layout{}, base.Errf("%s: chunked layout message is truncated", name) + } + if m[p]&0x01 != 0 { + dimSize = 8 + } + p++ + } + l.dims = make([]int, rank) + for i := range rank { + if p+dimSize > len(m) { + return hdf5Layout{}, base.Errf("%s: chunk dimensions run past the message", name) + } + l.dims[i] = int(uint64At(m[p:], dimSize)) + p += dimSize + } + default: + return hdf5Layout{}, base.Errf("%s: storage class %d is not supported", name, class) + } + return l, nil +} + +// hdf5Filter is one entry of a filter pipeline. +type hdf5Filter struct { + id uint16 + values []uint32 +} + +func decodeFilters(m []byte) ([]hdf5Filter, error) { + const name = "LoadHDF5" + if len(m) < 2 { + return nil, base.Errf("%s: filter pipeline message is truncated", name) + } + version := m[0] + if version != 1 { + // Only version 1's layout is known; parsing another version on + // this layout would fail far away with an unrelated truncation + // error instead of naming the defect. + return nil, base.Errf("%s: unsupported filter pipeline message version %d", name, version) + } + count := int(m[1]) + p := 2 + p += 6 // reserved + out := make([]hdf5Filter, 0, count) + for range count { + if p+8 > len(m) { + return nil, base.Errf("%s: filter pipeline message is truncated", name) + } + filter := hdf5Filter{id: binary.LittleEndian.Uint16(m[p:])} + nameLen := int(binary.LittleEndian.Uint16(m[p+2:])) + nValues := int(binary.LittleEndian.Uint16(m[p+6:])) + p += 8 + if p+nameLen > len(m) { + return nil, base.Errf("%s: filter name runs past the message", name) + } + p += nameLen + p = alignUp(p, 8) + for range nValues { + if p+4 > len(m) { + return nil, base.Errf("%s: filter values run past the message", name) + } + filter.values = append(filter.values, binary.LittleEndian.Uint32(m[p:])) + p += 4 + } + p = alignUp(p, 8) + switch filter.id { + case hdf5FilterDeflate, hdf5FilterShuffle, hdf5FilterFletcher32: + default: + return nil, base.Errf("%s: filter %d is not supported", name, filter.id) + } + out = append(out, filter) + } + return out, nil +} + +// dataset reads one dataset's values. +func (f *hdf5File) dataset(path string, msgs []hdf5Message, st *hdf5WalkState) (HDF5Dataset, error) { + const name = "LoadHDF5" + var ( + shape []int + dtype hdf5Type + layout hdf5Layout + filters []hdf5Filter + haveType bool + haveSpace bool + haveLay bool + ) + for _, m := range msgs { + switch m.typ { + case hdf5MsgDataspace: + s, err := decodeDataspace(m.data, f.lenSize) + if err != nil { + return HDF5Dataset{}, base.Errf("%s: dataset %q: %w", name, path, err) + } + shape, haveSpace = s, true + case hdf5MsgDatatype: + t, err := decodeType(m.data, false) + if err != nil { + return HDF5Dataset{}, base.Errf("%s: dataset %q: %w", name, path, err) + } + dtype, haveType = t, true + case hdf5MsgDataLayout: + l, err := decodeLayout(m.data, f.offSize) + if err != nil { + return HDF5Dataset{}, base.Errf("%s: dataset %q: %w", name, path, err) + } + layout, haveLay = l, true + case hdf5MsgFilterPipeline: + fs, err := decodeFilters(m.data) + if err != nil { + return HDF5Dataset{}, base.Errf("%s: dataset %q: %w", name, path, err) + } + filters = fs + } + } + if !haveSpace || !haveType || !haveLay { + return HDF5Dataset{}, base.Errf("%s: dataset %q is missing its %s", name, path, + missingOf(haveSpace, haveType, haveLay)) + } + if layout.class == 2 { + // A latest-version file indexes its chunks with the version 2 + // B-tree, which this reader does not walk; a classic file's + // chunks hang off the version 1 tree below. + if f.superblock >= 2 { + return HDF5Dataset{}, base.Errf("%s: dataset %q is chunked through a version 2 B-tree, which is not supported; rewrite the file with the default library version", name, path) + } + // The chunked message's last dimension is the element size, not + // a chunk dimension: the rank must be one less, and the value + // must agree with the datatype, which is a strong check that + // both were read correctly. + if len(layout.dims) != len(shape)+1 { + return HDF5Dataset{}, base.Errf("%s: dataset %q declares %d chunk dimensions for rank %d", + name, path, len(layout.dims), len(shape)) + } + if got := layout.dims[len(layout.dims)-1]; got != dtype.width { + return HDF5Dataset{}, base.Errf("%s: dataset %q declares %d-byte chunk elements against a %d-byte datatype", + name, path, got, dtype.width) + } + layout.dims = layout.dims[:len(shape)] + } + // Fixed-point data lands by its stored width and signedness; an + // unsigned 64-bit value has no exact core dtype and is refused + // before any storage is read, with the same message the value + // decode reported it with. + if !dtype.isFloat && !dtype.signed && dtype.width == 8 { + return HDF5Dataset{}, base.Errf("%s: dataset %q: %w", name, path, + base.Errf("unsigned 64-bit integers have no exact core dtype")) + } + // The size arithmetic bounds every extent against the budget before + // it is multiplied: a hostile dataspace must fail the cap, not wrap + // the product into a small want that passes it. The validated byte + // count is what every storage class below works from. + want, err := hdf5ByteExtent(shape, dtype.width, hdf5MaxDatasetBytes) + if err != nil { + return HDF5Dataset{}, base.Errf("%s: dataset %q: %w; larger datasets need mmap", name, path, err) + } + n := int(want) + // The per-dataset cap bounds one dataset; the walk's aggregate + // bounds their sum, so a file declaring thousands of cap-fitting + // datasets refuses here instead of allocating the sum. The sum is + // charged past the value bytes: every accepted dataset also costs + // its result structure and its copied path, which the envelope and + // the path length stand for, so empty datasets pay too. + charge := int64(n) + int64(len(st.path)) + hdf5DatasetEnvelope + if charge > st.budget { + return HDF5Dataset{}, base.Errf("%s: dataset %q declares %d bytes, the file's datasets are past the %d byte read budget; larger files need mmap", + name, path, n, hdf5MaxDatasetBytes) + } + st.budget -= charge + if layout.class == 2 { + values, cerr := f.chunkedArray(path, shape, dtype, layout, filters, n) + if cerr != nil { + return HDF5Dataset{}, cerr + } + return HDF5Dataset{Path: path, Shape: shape, Values: values}, nil + } + raw, rerr := f.rawData(path, shape, dtype, layout, n) + if rerr != nil { + return HDF5Dataset{}, rerr + } + values, aerr := arrayFromRaw(raw, dtype, shape) + if aerr != nil { + return HDF5Dataset{}, base.Errf("%s: dataset %q: %w", name, path, aerr) + } + return HDF5Dataset{Path: path, Shape: shape, Values: values}, nil +} + +func missingOf(haveSpace, haveType, haveLay bool) string { + switch { + case !haveSpace: + return "dataspace" + case !haveType: + return "datatype" + case !haveLay: + return "storage layout" + } + return "messages" +} + +// rawData gathers a dataset's bytes for the storage classes that hold +// them contiguously: compact storage is the message itself and +// contiguous storage is a span of the file. Chunked storage decodes +// into the values array directly, in chunkedArray. n is the dataset's +// byte size, validated against the read budget by the caller. +func (f *hdf5File) rawData(path string, shape []int, dtype hdf5Type, layout hdf5Layout, n int) ([]byte, error) { + const name = "LoadHDF5" + switch layout.class { + case 0: + if len(layout.compact) < n { + return nil, base.Errf("%s: dataset %q holds %d compact bytes, %d are needed", + name, path, len(layout.compact), n) + } + return layout.compact[:n], nil + case 1: + if layout.addr == math.MaxUint64 { + // The undefined address means storage was never allocated: + // for an empty dataset that is the whole answer (the fill + // value is the data, as the chunked branch already models), + // and for a non-empty one the bytes simply do not exist. + if n == 0 { + return []byte{}, nil + } + return nil, base.Errf("%s: dataset %q declares storage that was never allocated", name, path) + } + raw := f.bytes(layout.addr, uint64(n)) + if raw == nil { + return nil, base.Errf("%s: dataset %q lies outside the file", name, path) + } + return raw, nil + } + return nil, base.Errf("%s: dataset %q has an unsupported storage class", name, path) +} + +// hdf5ChunkStage carries what one dataset's chunk decode reuses across +// its chunks: the inflate output, the unshuffle staging, the inflate +// reader over the file's bytes, its limited view, and the index +// vectors the placement walk uses. A decode is serial, so the stage +// belongs to one call and nothing here is shared between concurrent +// loads. +type hdf5ChunkStage struct { + inflate []byte + unshuf []byte + zr io.ReadCloser + zrReset zlib.Resetter + br bytes.Reader + lim io.LimitedReader + offsets []uint64 + vec []int + seen map[uint64]bool +} + +// chunkedArray reassembles a chunked dataset through its version 1 +// B-tree, running each chunk back through the filter pipeline and +// decoding each cell into the values array as it lands. bytes is the +// dataset's size, already validated against the read budget by the +// caller, so this function never re-derives it in int. The staging +// buffers, the inflate reader and the index vectors are one stage, +// reused across every chunk of the dataset. +func (f *hdf5File) chunkedArray(path string, shape []int, dtype hdf5Type, layout hdf5Layout, filters []hdf5Filter, bytes int) (*core.Array, error) { + const name = "LoadHDF5" + rank := len(shape) + if len(layout.dims) != rank { + return nil, base.Errf("%s: dataset %q declares %d chunk dimensions for rank %d", name, path, len(layout.dims), rank) + } + for _, c := range layout.dims { + // A chunk holds at least one element of the datatype; a zero or + // negative declared dimension is a hostile header, and both the + // strides below and the placement's overlap box degenerate on it. + if c <= 0 { + return nil, base.Errf("%s: dataset %q declares a %d-element chunk dimension", name, path, c) + } + } + // One chunk may not exceed the reader budget on its own, because the + // filters inflate into a buffer of exactly this size. The product is + // bounded extent by extent, so a hostile chunk dimension fails the + // cap instead of wrapping it. + chunkBytes, err := hdf5ByteExtent(layout.dims, dtype.width, hdf5MaxDatasetBytes) + if err != nil { + return nil, base.Errf("%s: dataset %q declares a chunk: %w", name, path, err) + } + cb := int(chunkBytes) + // The destination array and the per-cell decode carry the value + // mapping arrayFromRaw applies to contiguous bytes; the cells land + // in the payload as the walk places them, so no assembled copy of + // the dataset's bytes exists. Fixed-point cells land by stored width + // and signedness, exactly as the contiguous path lands them, and a + // boolean cell outside the enumeration's {0, 1} values is refused + // rather than coerced. + nElems := 0 + if dtype.width > 0 { + nElems = bytes / dtype.width + } + var put func(data []byte, ci, oi int) error + var arr *core.Array + switch { + case dtype.isFloat && dtype.width == 8: + vals := make([]float64, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = math.Float64frombits(binary.LittleEndian.Uint64(data[ci*8:])) + return nil + } + arr, err = core.FloatsFromArray(vals, shape...) + case dtype.isFloat && dtype.width == 4: + vals := make([]float32, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = math.Float32frombits(binary.LittleEndian.Uint32(data[ci*4:])) + return nil + } + // The aliasing constructor: the decode fills the array's own + // payload through vals, which a copying constructor would have + // left behind as a separate slice. + arr, err = core.FromFloat32Slice(vals, shape...) + case dtype.isBool: + vals := make([]bool, nElems) + put = func(data []byte, ci, oi int) error { + b := data[ci] + if b > 1 { + return base.Errf("a boolean value of %d is outside the members 0 and 1", b) + } + vals[oi] = b != 0 + return nil + } + arr, err = core.BoolsFromArray(vals, shape...) + case !dtype.signed && dtype.width == 8: + // Unreachable through LoadHDF5, which refuses the datatype + // first; the decode carries the same refusal for direct callers. + return nil, base.Errf("%s: dataset %q: %w", name, path, + base.Errf("unsigned 64-bit integers have no exact core dtype")) + case dtype.width == 1 && dtype.signed: + vals := make([]int8, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = int8(data[ci]) + return nil + } + arr, err = core.Int8sFromArray(vals, shape...) + case dtype.width == 1: + vals := make([]uint8, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = data[ci] + return nil + } + arr, err = core.Uint8sFromArray(vals, shape...) + case dtype.width == 2 && dtype.signed: + vals := make([]int16, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = int16(binary.LittleEndian.Uint16(data[ci*2:])) + return nil + } + arr, err = core.Int16sFromArray(vals, shape...) + case dtype.width == 2: + vals := make([]uint16, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = binary.LittleEndian.Uint16(data[ci*2:]) + return nil + } + arr, err = core.Uint16sFromArray(vals, shape...) + case dtype.width == 4 && dtype.signed: + vals := make([]int32, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = int32(binary.LittleEndian.Uint32(data[ci*4:])) + return nil + } + arr, err = core.Int32sFromArray(vals, shape...) + case dtype.width == 4: + vals := make([]uint32, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = binary.LittleEndian.Uint32(data[ci*4:]) + return nil + } + arr, err = core.Uint32sFromArray(vals, shape...) + default: + // Signed 64-bit, the only fixed-point width left. + vals := make([]int64, nElems) + put = func(data []byte, ci, oi int) error { + vals[oi] = int64(binary.LittleEndian.Uint64(data[ci*8:])) + return nil + } + arr, err = core.IntsFromArray(vals, shape...) + } + if err != nil { + return nil, base.Errf("%s: dataset %q: %w", name, path, err) + } + if layout.addr == math.MaxUint64 { + return arr, nil // no chunks written: the fill value is the data + } + stage := &hdf5ChunkStage{ + seen: map[uint64]bool{}, + offsets: make([]uint64, rank), + vec: make([]int, 5*rank), + } + seenNodes := map[uint64]bool{} + return arr, f.chunkTreeWalk(layout.addr, layout, rank, seenNodes, func(addr uint64, offsets []uint64, mask uint32, size int) error { + if stage.seen[addr] { + return base.Errf("%s: dataset %q has two chunks at the same address", name, path) + } + stage.seen[addr] = true + stored := f.bytes(addr, uint64(size)) + if stored == nil { + return base.Errf("%s: dataset %q has a chunk outside the file", name, path) + } + data, ferr := runFilters(stored, filters, mask, dtype, cb, stage) + if ferr != nil { + return base.Errf("%s: dataset %q: %w", name, path, ferr) + } + return placeCells(shape, layout.dims, offsets, data, dtype.width, cb, nElems, stage.vec, put) + }, 0, stage.offsets) +} + +// placeCells walks the overlap of one chunk and the dataset, calling +// put for every cell with the chunk-local cell index and the output +// element index; the decode writes the value as the cell lands, and a +// cell put refuses stops the walk with that error. +// chunkBytes is the byte size of one complete chunk, computed by the +// caller with every dimension bounded against the reader budget, so +// the strides cannot wrap and a chunk is known to hold whole elements. +// The loop walks the overlap box in output coordinates, so the work is +// proportional to the placed elements: a chunk hanging past the +// shape's edge is padding the file may hold but the array has no room +// for, and one starting past an axis's end places nothing at all. +// Walking the whole declared chunk instead would let a hostile chunk +// dimension pin the reader for hours without touching a byte. vec +// carries the box bounds, the row-major strides of chunk and output, +// and the counter: five vectors of the rank, reused across chunks. +func placeCells(shape, chunkDim []int, offsets []uint64, data []byte, width, chunkBytes, maxOut int, vec []int, put func(data []byte, ci, oi int) error) error { + const name = "LoadHDF5" + if len(data) == 0 { + return nil + } + rank := len(shape) + if len(offsets) != rank { + return base.Errf("%s: a chunk carries %d offsets for rank %d", name, len(offsets), rank) + } + // The chunk must arrive complete: a truncated deflate stream would + // otherwise read past its data below. + if len(data) < chunkBytes { + return base.Errf("%s: a chunk holds %d bytes, a full chunk of %d is needed", name, len(data), chunkBytes) + } + lo := vec[:rank] + hi := vec[rank : 2*rank] + chunkStride := vec[2*rank : 3*rank] + outStride := vec[3*rank : 4*rank] + count := vec[4*rank : 5*rank] + for d := range rank { + if offsets[d] >= uint64(shape[d]) { + return nil + } + lo[d] = int(offsets[d]) + hi[d] = min(lo[d]+chunkDim[d], shape[d]) + if lo[d] == hi[d] { + return nil + } + } + // Row-major strides of the chunk and of the output. Both products + // are bounded: the chunk's by chunkBytes and the array's by the + // dataset's own validated byte count. + cs, os := 1, 1 + for i := rank - 1; i >= 0; i-- { + chunkStride[i], outStride[i] = cs, os + cs *= chunkDim[i] + os *= shape[i] + } + // The element walk indexes both sides with a computed offset. The + // chunk's data may be a view of the file buffer, whose capacity + // runs past the chunk, so a stride that reached outside the chunk + // would read the file's other bytes silently instead of failing: + // the two limits are therefore checked explicitly, once, in element + // units. + maxChunk := len(data) / width + copy(count, lo) + for { + // The element's output coordinate and chunk-local coordinate. + oi, ci := 0, 0 + for d := range rank { + oi += count[d] * outStride[d] + ci += (count[d] - lo[d]) * chunkStride[d] + } + if oi < 0 || oi >= maxOut || ci < 0 || ci >= maxChunk { + return base.Errf("%s: a chunk element lands at element %d of the %d-element array and %d of the %d-element chunk", + name, oi, maxOut, ci, maxChunk) + } + if err := put(data, ci, oi); err != nil { + return err + } + // Advance the odometer over the overlap box. + d := rank - 1 + for d >= 0 { + count[d]++ + if count[d] < hi[d] { + break + } + count[d] = lo[d] + d-- + } + if d < 0 { + return nil + } + } +} + +// chunkTree walks the chunk B-tree, calling place for every chunk, +// with a fresh per-entry offset vector. The decode reuses one vector +// across the whole walk through chunkTreeWalk instead. +func (f *hdf5File) chunkTree(addr uint64, layout hdf5Layout, rank int, nodes map[uint64]bool, place func(uint64, []uint64, uint32, int) error, depth int) error { + return f.chunkTreeWalk(addr, layout, rank, nodes, place, depth, make([]uint64, rank)) +} + +// chunkTreeWalk is the walk itself, with the per-entry offset vector +// supplied by the caller, so a decode allocates one for the dataset +// rather than one per chunk entry; a place callback consumes the +// offsets before it returns. nodes carries the addresses of the +// interior nodes already visited: a legal B-tree never revisits one, +// and without the set a hostile file whose node lists itself among its +// children would multiply the walk into entries^depth visits before +// the depth guard ever fires. +func (f *hdf5File) chunkTreeWalk(addr uint64, layout hdf5Layout, rank int, nodes map[uint64]bool, place func(uint64, []uint64, uint32, int) error, depth int, offsets []uint64) error { + const name = "LoadHDF5" + if depth > 32 { + return base.Errf("%s: the chunk B-tree is more than 32 levels deep", name) + } + if nodes[addr] { + return base.Errf("%s: the chunk B-tree revisits node %d", name, addr) + } + nodes[addr] = true + // The version 1 B-tree header is 8+2*offSize wide, the signature, + // type, level and entry count plus the two sibling addresses. + node := f.bytes(addr, uint64(8+2*f.offSize)) + if node == nil || !bytes.Equal(node[:4], hdf5Tree) { + return base.Errf("%s: no chunk B-tree node at %d", name, addr) + } + if node[4] != 1 { + return base.Errf("%s: B-tree node type %d is not a chunk tree", name, node[4]) + } + level := node[5] + entries := int(binary.LittleEndian.Uint16(node[6:])) + // The node header is 8+2*offSize wide, as in the group B-tree; the + // first key starts behind it. + nodeHeader := 8 + 2*f.offSize + // A v1 chunk key is the chunk size, the filter mask and one offset + // per dimension of the chunk *plus* the element size, which is why + // the key has one more offset than the dataset has axes; each entry + // is a key plus a child address. + keySize := 8 + 8*(rank+1) + entrySize := keySize + f.offSize + start := int(addr) + nodeHeader + for i := range entries { + p := start + i*entrySize + size := int(f.u64At(p, 4)) + mask := uint32(f.u64At(p+4, 4)) + for d := range rank { + offsets[d] = f.length(p + 8 + d*8) + } + child := f.at(p + keySize) + if level == 0 { + if err := place(child, offsets, mask, size); err != nil { + return err + } + continue + } + if err := f.chunkTreeWalk(child, layout, rank, nodes, place, depth+1, offsets); err != nil { + return err + } + } + return nil +} + +// runFilters undoes the filter pipeline for one chunk, in reverse +// order; a filter whose bit is set in the mask did not run on it. The +// returned bytes are either the stored chunk untouched or one of the +// stage's buffers, valid until the stage handles its next chunk. +func runFilters(data []byte, filters []hdf5Filter, mask uint32, dtype hdf5Type, chunkBytes int, stage *hdf5ChunkStage) ([]byte, error) { + const name = "LoadHDF5" + out := data + for i, filter := range slices.Backward(filters) { + if mask&(1< want { + return nil, base.Errf("the deflate stream inflates past the chunk size") + } + return out, nil +} + +// unshuffleInto undoes the shuffle filter into out, which holds the +// same length as data: the filter transposes the bytes of the +// elements, so the first block holds every element's first byte. The +// transpose runs over tiles of rows: a tile is a short contiguous run +// of out, and the reads that fill it walk one row of the source +// sequentially, so neither stream strides across the whole chunk per +// column. Every byte still lands where the flat transpose put it. +func unshuffleInto(out, data []byte, width int) { + n := len(data) / width + const tileRows = 64 + for j0 := 0; j0 < n; j0 += tileRows { + j1 := min(j0+tileRows, n) + for i := range width { + src := data[i*n+j0 : i*n+j1] + dst := out[j0*width+i:] + for k, b := range src { + dst[k*width] = b + } + } + } +} + +// checkFletcher32 verifies the fletcher32 checksum HDF5 appends to a +// filtered chunk. +func checkFletcher32(data []byte) error { + body, sum := data[:len(data)-4], binary.LittleEndian.Uint32(data[len(data)-4:]) + if got := fletcher32(body); got != sum { + return base.Errf("the fletcher32 checksum does not match: %08x, want %08x", got, sum) + } + return nil +} + +// arrayFromRaw builds the array for a dataset from its bytes. The +// decoded buffer is exactly one element per value, so the array takes +// it over instead of copying it a second time. Fixed-point data lands +// the core dtype its stored width and signedness declare, a boolean +// enumeration lands Bool, and a boolean cell outside the members 0 and +// 1 is refused rather than coerced. +func arrayFromRaw(raw []byte, dtype hdf5Type, shape []int) (*core.Array, error) { + switch { + case dtype.isFloat && dtype.width == 8: + vals := make([]float64, len(raw)/8) + for i := range vals { + vals[i] = math.Float64frombits(binary.LittleEndian.Uint64(raw[i*8:])) + } + return core.FloatsFromArray(vals, shape...) + case dtype.isFloat && dtype.width == 4: + vals := make([]float32, len(raw)/4) + for i := range vals { + vals[i] = math.Float32frombits(binary.LittleEndian.Uint32(raw[i*4:])) + } + return core.FromFloat32s(vals, shape...) + case dtype.isBool: + vals := make([]bool, len(raw)) + for i, b := range raw { + if b > 1 { + return nil, base.Errf("a boolean value of %d is outside the members 0 and 1", b) + } + vals[i] = b != 0 + } + return core.BoolsFromArray(vals, shape...) + case !dtype.signed && dtype.width == 8: + return nil, base.Errf("unsigned 64-bit integers have no exact core dtype") + } + switch { + case dtype.width == 1 && dtype.signed: + vals := make([]int8, len(raw)) + for i, b := range raw { + vals[i] = int8(b) + } + return core.Int8sFromArray(vals, shape...) + case dtype.width == 1: + // raw may be a view of the file buffer; the payload copies out + // of it before the array takes the copy over. + vals := make([]uint8, len(raw)) + copy(vals, raw) + return core.Uint8sFromArray(vals, shape...) + case dtype.width == 2 && dtype.signed: + vals := make([]int16, len(raw)/2) + for i := range vals { + vals[i] = int16(binary.LittleEndian.Uint16(raw[i*2:])) + } + return core.Int16sFromArray(vals, shape...) + case dtype.width == 2: + vals := make([]uint16, len(raw)/2) + for i := range vals { + vals[i] = binary.LittleEndian.Uint16(raw[i*2:]) + } + return core.Uint16sFromArray(vals, shape...) + case dtype.width == 4 && dtype.signed: + vals := make([]int32, len(raw)/4) + for i := range vals { + vals[i] = int32(binary.LittleEndian.Uint32(raw[i*4:])) + } + return core.Int32sFromArray(vals, shape...) + case dtype.width == 4: + vals := make([]uint32, len(raw)/4) + for i := range vals { + vals[i] = binary.LittleEndian.Uint32(raw[i*4:]) + } + return core.Uint32sFromArray(vals, shape...) + case dtype.width == 8: + // Signed 64-bit: unsigned was refused above, and decodeType + // admits no other fixed-point width. + vals := make([]int64, len(raw)/8) + for i := range vals { + vals[i] = int64(binary.LittleEndian.Uint64(raw[i*8:])) + } + return core.IntsFromArray(vals, shape...) + } + return nil, base.Errf("%d-byte fixed-point values have no core dtype", dtype.width) +} + +// decodeAttribute reads one attribute message: its name and a textual +// rendering of its value, the shape the package's attribute maps use. +func (f *hdf5File) decodeAttribute(m []byte) (string, string, bool) { + if len(m) < 8 || m[0] != 1 { + return "", "", false + } + nameSize := int(binary.LittleEndian.Uint16(m[2:])) + typeSize := int(binary.LittleEndian.Uint16(m[4:])) + spaceSize := int(binary.LittleEndian.Uint16(m[6:])) + p := 8 + if p+nameSize+typeSize+spaceSize > len(m) { + return "", "", false + } + name := strings.TrimRight(string(m[p:p+nameSize]), "\x00") + // The name field is padded so the datatype starts on an eight-byte + // boundary. The padding is inside the message, so the datatype is + // checked against the body's own length: the slice underneath is a + // view of the file buffer, whose capacity runs past the body. + p = alignUp(p+nameSize, 8) + if p+typeSize > len(m) { + return "", "", false + } + dtype, err := decodeType(m[p:p+typeSize], true) + if err != nil { + return "", "", false + } + // The datatype message is padded to an eight-byte boundary before + // the dataspace begins. + p = alignUp(p+typeSize, 8) + if p+spaceSize > len(m) { + return "", "", false + } + // The dataspace's dimension fields are sized by the file's own + // length size, exactly as the dataset path's are. + shape, err := decodeDataspace(m[p:p+spaceSize], f.lenSize) + if err != nil { + return "", "", false + } + p += spaceSize + // The value's byte count is bounded extent by extent against what is + // left of the message: a hostile dataspace must fail here, not wrap + // the product into a negative that would slip past a check made + // afterwards and then size an allocation. + size, err := hdf5ByteExtent(shape, dtype.width, uint64(len(m)-p)) + if err != nil { + return "", "", false + } + raw := m[p : p+int(size)] + if dtype.class == 3 { + // A fixed-length string: the bytes are the text. + return name, strings.TrimRight(string(raw), "\x00"), true + } + if dtype.class == 9 { + // A variable-length string: the attribute holds a descriptor + // naming a global heap object, and that object holds the text. + // The descriptor is four bytes of length, the collection + // address, then the object index. + if len(raw) < 8+f.offSize { + return "", "", false + } + length := int(binary.LittleEndian.Uint32(raw)) + heapAddr := uint64At(raw[4:], f.offSize) + index := binary.LittleEndian.Uint32(raw[4+f.offSize:]) + text, err := f.heapString(heapAddr, index, length) + if err != nil { + return "", "", false + } + return name, text, true + } + // The element count the numeric rendering below walks, derived from + // the validated byte count rather than from a second multiplication. + count := 0 + if dtype.width > 0 { + count = len(raw) / dtype.width + } + vals := make([]string, 0, count) + for i := range count { + cell := raw[i*dtype.width : (i+1)*dtype.width] + switch { + case dtype.isFloat && dtype.width == 8: + vals = append(vals, strconv.FormatFloat(math.Float64frombits(binary.LittleEndian.Uint64(cell)), 'g', -1, 64)) + case dtype.isFloat && dtype.width == 4: + vals = append(vals, strconv.FormatFloat(float64(math.Float32frombits(binary.LittleEndian.Uint32(cell))), 'g', -1, 32)) + case dtype.isBool: + // The boolean enumeration renders as its own values; a cell + // outside them is a file that contradicts its datatype, and + // the attribute is dropped rather than guessed at. + if cell[0] > 1 { + return "", "", false + } + vals = append(vals, strconv.FormatInt(int64(cell[0]), 10)) + case dtype.signed: + // The datatype's own signed bit decides how its stored bits + // read: an int8 attribute of -1 renders "-1", not 255. + u := uint64At(cell, dtype.width) + if width := uint(dtype.width); width < 64 { + mask := uint64(1)<<(8*width) - 1 + if u&(uint64(1)<<(8*width-1)) != 0 { + u |= ^mask // sign-extend into the 64-bit read + } + } + vals = append(vals, strconv.FormatInt(int64(u), 10)) + default: + vals = append(vals, strconv.FormatUint(uint64At(cell, dtype.width), 10)) + } + } + if len(vals) == 1 { + return name, vals[0], true + } + return name, "[" + strings.Join(vals, ", ") + "]", true +} + +// heapString reads a variable-length string out of a global heap +// collection: the descriptor names the collection and an object index, +// and the collection lists (index, reference count, size, bytes) +// records, each padded to eight bytes. +func (f *hdf5File) heapString(addr uint64, index uint32, length int) (string, error) { + const name = "LoadHDF5" + head := f.bytes(addr, 16) + if head == nil || !bytes.Equal(head[:4], []byte("GCOL")) { + return "", base.Errf("%s: no global heap collection at %d", name, addr) + } + size := f.length(int(addr) + 8) // the collection size, header included + collection := f.bytes(addr, size) + if collection == nil { + return "", base.Errf("%s: the global heap collection at %d lies outside the file", name, addr) + } + p := 8 + f.lenSize + for p+8+f.lenSize <= len(collection) { + objIndex := uint32(binary.LittleEndian.Uint16(collection[p:])) + objSize := uint64At(collection[p+8:], f.lenSize) + body := p + 8 + f.lenSize + // The bound compares in uint64: an int addition would wrap a + // hostile object size past the guard into a negative value and + // the slice below would panic. + if objSize > uint64(len(collection)-body) { + break + } + size := int(objSize) + if objIndex == index { + text := collection[body : body+size] + if length > 0 && length < len(text) { + text = text[:length] + } + return strings.TrimRight(string(text), "\x00"), nil + } + p = alignUp(body+size, 8) + } + return "", base.Errf("%s: the global heap object %d is not in the collection at %d", name, index, addr) +} + +// fletcher32 is the checksum HDF5's fletcher32 filter appends: a +// 32-bit Fletcher sum over 16-bit words, ones-complement folded. +func fletcher32(data []byte) uint32 { + sum1, sum2 := uint32(0xffff), uint32(0xffff) + // The filter runs in blocks of 359 words (718 bytes). + for len(data) > 0 { + n := min(len(data), 718) + block := data[:n] + if n%2 != 0 { + block = append(append([]byte{}, block...), 0) + } + for i := 0; i+1 < len(block); i += 2 { + sum1 += uint32(binary.BigEndian.Uint16(block[i:])) + sum2 += sum1 + } + sum1 = (sum1 & 0xffff) + (sum1 >> 16) + sum2 = (sum2 & 0xffff) + (sum2 >> 16) + data = data[n:] + } + sum1 = (sum1 & 0xffff) + (sum1 >> 16) + sum2 = (sum2 & 0xffff) + (sum2 >> 16) + return sum2<<16 | sum1 +} + +// hdf5Lookup3 is the Jenkins lookup3 hash in the little-endian, +// byte-wise form HDF5 stores as the checksum of superblock versions 2 +// and 3 and of every chunk of an object header version 2, with the +// zero initial value the library uses. The main loop mixes whole +// 12-byte blocks and the switch folds the tail, whose bytes fall into +// the three words highest first. +func hdf5Lookup3(key []byte) uint32 { + a := uint32(0xdeadbeef) + uint32(len(key)) + b := a + c := a + p := 0 + for ; len(key)-p > 12; p += 12 { + a += binary.LittleEndian.Uint32(key[p:]) + b += binary.LittleEndian.Uint32(key[p+4:]) + c += binary.LittleEndian.Uint32(key[p+8:]) + a, b, c = hdf5Lookup3Mix(a, b, c) + } + switch r := key[p:]; len(r) { + case 12: + c += uint32(r[11]) << 24 + fallthrough + case 11: + c += uint32(r[10]) << 16 + fallthrough + case 10: + c += uint32(r[9]) << 8 + fallthrough + case 9: + c += uint32(r[8]) + fallthrough + case 8: + b += uint32(r[7]) << 24 + fallthrough + case 7: + b += uint32(r[6]) << 16 + fallthrough + case 6: + b += uint32(r[5]) << 8 + fallthrough + case 5: + b += uint32(r[4]) + fallthrough + case 4: + a += uint32(r[3]) << 24 + fallthrough + case 3: + a += uint32(r[2]) << 16 + fallthrough + case 2: + a += uint32(r[1]) << 8 + fallthrough + case 1: + a += uint32(r[0]) + case 0: + return c + } + a, b, c = hdf5Lookup3Final(a, b, c) + return c +} + +// hdf5Lookup3Mix is lookup3's inner round. +func hdf5Lookup3Mix(a, b, c uint32) (uint32, uint32, uint32) { + a -= c + a ^= bits.RotateLeft32(c, 4) + c += b + b -= a + b ^= bits.RotateLeft32(a, 6) + a += c + c -= b + c ^= bits.RotateLeft32(b, 8) + b += a + a -= c + a ^= bits.RotateLeft32(c, 16) + c += b + b -= a + b ^= bits.RotateLeft32(a, 19) + a += c + c -= b + c ^= bits.RotateLeft32(b, 4) + b += a + return a, b, c +} + +// hdf5Lookup3Final is lookup3's closing avalanche. +func hdf5Lookup3Final(a, b, c uint32) (uint32, uint32, uint32) { + c ^= b + c -= bits.RotateLeft32(b, 14) + a ^= c + a -= bits.RotateLeft32(c, 11) + b ^= a + b -= bits.RotateLeft32(a, 25) + c ^= b + c -= bits.RotateLeft32(b, 16) + a ^= c + a -= bits.RotateLeft32(c, 4) + b ^= a + b -= bits.RotateLeft32(a, 14) + c ^= b + c -= bits.RotateLeft32(b, 24) + return a, b, c +} diff --git a/io/hdf5_bool_fixture_test.go b/io/hdf5_bool_fixture_test.go new file mode 100644 index 0000000..32dfa9c --- /dev/null +++ b/io/hdf5_bool_fixture_test.go @@ -0,0 +1,90 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "path/filepath" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The boolean fixture bool_enum.h5 under testdata/h5 was written by +// hand against the HDF5 1.8 specification, in the shape the reference +// library writes for an old-style group with a symbol table, and it +// validates against the reference library's own tools: h5dump reads +// the dataset back as an H5T_ENUM over H5T_STD_U8LE with the members +// FALSE = 0 and TRUE = 1 and the values TRUE, FALSE, TRUE, TRUE, +// FALSE. Its layout deliberately differs from what this package's +// writer produces: the messages of the dataset's object header sit in +// the order datatype, fill value, dataspace, layout, the dataspace +// carries no maximum dimensions and the fill value message declares a +// defined zero fill, so the pins below prove the reader's tolerance of +// a foreign layout rather than a round trip of its own bytes. + +// TestLoadHDF5ForeignBoolFixture pins the read side: the hand-crafted +// boolean enumeration of a foreign layout lands the dataset "flags" +// with the core Bool dtype and the values the reference library reads +// from the same bytes. +func TestLoadHDF5ForeignBoolFixture(t *testing.T) { + sets, err := LoadHDF5(h5Fixture(t, "bool_enum.h5")) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want 1", len(sets)) + } + d := sets[0] + if d.Path != "/flags" { + t.Fatalf("path = %q, want /flags", d.Path) + } + if s := d.Shape; len(s) != 1 || s[0] != 5 { + t.Fatalf("shape = %v, want [5]", s) + } + if dt := d.Values.Dtype(); dt != core.Bool { + t.Fatalf("dtype = %s, want bool", dt) + } + want := []bool{true, false, true, true, false} + if got := d.Values.RawBools()[:5]; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } +} + +// TestSaveHDF5BoolConvention pins the write side against the same +// logical values: SaveHDF5 stores them through the boolean enumeration +// convention and LoadHDF5 reads them back identically. The pin holds +// the two-sided agreement on the convention, not a byte match with the +// foreign fixture, whose layout the writer need not reproduce. +func TestSaveHDF5BoolConvention(t *testing.T) { + want := []bool{true, false, true, true, false} + values, err := core.FromBools(want, 5) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + path := filepath.Join(t.TempDir(), "bool_convention.h5") + if err := SaveHDF5(path, []HDF5Dataset{{Path: "/flags", Shape: []int{5}, Values: values}}, nil); err != nil { + t.Fatalf("SaveHDF5: %v", err) + } + sets, err := LoadHDF5(path) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want 1", len(sets)) + } + d := sets[0] + if d.Path != "/flags" { + t.Fatalf("path = %q, want /flags", d.Path) + } + if s := d.Shape; len(s) != 1 || s[0] != 5 { + t.Fatalf("shape = %v, want [5]", s) + } + if dt := d.Values.Dtype(); dt != core.Bool { + t.Fatalf("dtype = %s, want bool", dt) + } + if got := d.Values.RawBools()[:5]; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } +} diff --git a/io/hdf5_continuation_pins_test.go b/io/hdf5_continuation_pins_test.go new file mode 100644 index 0000000..3fee91a --- /dev/null +++ b/io/hdf5_continuation_pins_test.go @@ -0,0 +1,161 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "math" + "strings" + "testing" +) + +// Regression pins: HDF5 headers the reference library writes +// once they outgrow their first block, byte orders and unallocated +// storage the reader used to answer silently, and the CSV byte order +// mark every spreadsheet writes. + +// h5Link renders a continuation link message body: the block address +// and its length, eight bytes each in the hostile layout. +func h5Link(addr, length uint64) []byte { + b := make([]byte, 16) + binary.LittleEndian.PutUint64(b, addr) + binary.LittleEndian.PutUint64(b[8:], length) + return b +} + +// h5ContinuationBlock writes a flat message-list block at off and +// returns the end offset, everything eight-aligned. +func h5ContinuationBlock(f []byte, off int, msgs ...h5Msg) int { + for _, m := range msgs { + binary.LittleEndian.PutUint16(f[off:], m.typ) + binary.LittleEndian.PutUint16(f[off+2:], uint16(len(m.body))) + copy(f[off+8:], m.body) + off += alignUp(8+len(m.body), 8) + } + return off +} + +// TestLoadHDF5ChainedContinuation: the header walk followed exactly +// one continuation block and dropped every message of the second and +// later ones, so a legal attribute-rich file lost its datasets. Here +// the dataspace sits in the header, the datatype in the first +// continuation block and the layout in the second: only a walk that +// follows the whole chain can assemble the dataset. +func TestLoadHDF5ChainedContinuation(t *testing.T) { + const blockA, blockB, dataAt, n = 448, 640, 800, 1024 + f := h5HostileFile(n) + end := h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(2)}, + h5Msg{hdf5MsgContinuation, h5Link(blockA, 112)}, + ) + if end > blockA { + t.Fatalf("the header runs to %d, past the first block at %d", end, blockA) + } + endA := h5ContinuationBlock(f, blockA, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgContinuation, h5Link(blockB, 88)}, + ) + if endA > blockB { + t.Fatalf("block A runs to %d, past block B at %d", endA, blockB) + } + h5ContinuationBlock(f, blockB, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 16)}, + ) + binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(1.5)) + binary.LittleEndian.PutUint64(f[dataAt+8:], math.Float64bits(-2.5)) + sets, err := LoadHDF5(writeHostile(t, "chained.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want 1", len(sets)) + } + v0 := sets[0].Values.FloatAt(0) + if v0 != 1.5 || sets[0].Values.FloatAt(1) != -2.5 { + t.Fatalf("values = %v, %v, want 1.5 and -2.5", v0, sets[0].Values.FloatAt(1)) + } +} + +// TestLoadHDF5ContinuationCycle: a continuation block listing itself +// must be an error, not an infinite walk. +func TestLoadHDF5ContinuationCycle(t *testing.T) { + const blockA, n = 448, 640 + f := h5HostileFile(n) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(2)}, + h5Msg{hdf5MsgContinuation, h5Link(blockA, 40)}, + ) + // The block's only entry is a link back to itself. + h5ContinuationBlock(f, blockA, + h5Msg{hdf5MsgContinuation, h5Link(blockA, 40)}, + ) + if _, err := LoadHDF5(writeHostile(t, "cycle.h5", f)); err == nil || !strings.Contains(err.Error(), "twice") { + t.Fatalf("a self-referencing continuation block: err = %v", err) + } +} + +// TestLoadHDF5BigEndianRefusal: the byte-order bit of the datatype was +// parsed and never checked, so a big-endian dataset decoded as +// byte-swapped noise with no error. +func TestLoadHDF5BigEndianRefusal(t *testing.T) { + const dataAt, n = 448, 512 + f := h5HostileFile(n) + beType := h5FloatType(8) + beType[1] = 0x01 // class bit field: bit 0 set means big-endian + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(2)}, + h5Msg{hdf5MsgDatatype, beType}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 16)}, + ) + if _, err := LoadHDF5(writeHostile(t, "bigendian.h5", f)); err == nil || !strings.Contains(err.Error(), "big-endian") { + t.Fatalf("a big-endian dataset: err = %v", err) + } +} + +// TestLoadHDF5UnallocatedContiguous: the undefined storage address on +// a contiguous dataset refused the whole file; an empty dataset there +// is legal (nothing was ever allocated) and must load as empty, while +// a non-empty one names what is missing. +func TestLoadHDF5UnallocatedContiguous(t *testing.T) { + t.Run("empty", func(t *testing.T) { + f := h5HostileFile(512) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(0)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(math.MaxUint64, 0)}, + ) + sets, err := LoadHDF5(writeHostile(t, "empty.h5", f)) + if err != nil { + t.Fatalf("an empty unallocated dataset: %v", err) + } + if len(sets) != 1 || sets[0].Values.Len() != 0 { + t.Fatalf("datasets = %d, want one empty dataset", len(sets)) + } + }) + t.Run("non-empty", func(t *testing.T) { + f := h5HostileFile(512) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(2)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(math.MaxUint64, 16)}, + ) + if _, err := LoadHDF5(writeHostile(t, "noalloc.h5", f)); err == nil || !strings.Contains(err.Error(), "never allocated") { + t.Fatalf("a non-empty unallocated dataset: err = %v", err) + } + }) +} + +// TestLoadCSVSkipBOM: a leading UTF-8 byte order mark used to glue +// itself onto the first field and fail the whole load with a strconv +// error. +func TestLoadCSVSkipBOM(t *testing.T) { + const in = "\xEF\xBB\xBF1.5,2.5\n3,4\n" + a, err := LoadCSVReader(strings.NewReader(in), false) + if err != nil { + t.Fatalf("LoadCSVReader with a BOM: %v", err) + } + if a.FloatAt(0) != 1.5 || a.FloatAt(1) != 2.5 || a.FloatAt(2) != 3 || a.FloatAt(3) != 4 { + t.Fatalf("values = %v", a) + } +} diff --git a/io/hdf5_structure_pins_test.go b/io/hdf5_structure_pins_test.go new file mode 100644 index 0000000..ae62b0e --- /dev/null +++ b/io/hdf5_structure_pins_test.go @@ -0,0 +1,527 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "math" + "slices" + "strings" + "testing" + "time" +) + +// Regression pins: continuation links the reader sliced past the +// end of, a local heap header sized by the wrong field, the diamond of +// hard links that multiplied the walk, the read budget that only +// counted value bytes, the B-tree sizes pinned to eight-byte addresses, +// negative FITS column counts, the NAXIS prefix that swallowed user +// keywords, and the version 2 header features no test exercised. + +// h5SizedHostileFile returns an n-byte HDF5 file with the signature and +// a version 0 superblock of the given address and length sizes whose +// root object header sits at rootAt. h5HostileFile is the 8/8 case; +// the hostile files here need the others. +func h5SizedHostileFile(n, offSize, lenSize int, rootAt uint64) []byte { + f := make([]byte, n) + copy(f, hdf5Magic) + f[8] = 0 // superblock version 0 + f[13] = byte(offSize) + f[14] = byte(lenSize) + // Four addresses of offSize bytes from 24: base, free space, end of + // file, driver information. + putAddr := func(at int, v uint64) { + switch offSize { + case 4: + binary.LittleEndian.PutUint32(f[at:], uint32(v)) + case 8: + binary.LittleEndian.PutUint64(f[at:], v) + } + } + putAddr(24, 0) + putAddr(28, math.MaxUint64) + putAddr(32, uint64(n)) + putAddr(36, math.MaxUint64) + // The root symbol table entry: a link name offset of the length + // size, then the object header address of the address size. + entry := 24 + 4*offSize + putAddr(entry+lenSize, rootAt) + return f +} + +// h5HardLink renders a version 1 hard-link message body: a one-byte +// name length, the name and the object header address. +func h5HardLink(name string, addr uint64) []byte { + b := append([]byte{1, 0, byte(len(name))}, name...) + return binary.LittleEndian.AppendUint64(b, addr) +} + +// h5V3Superblock returns an n-byte file with a version 3 superblock +// (eight-byte addresses and lengths, the whole block under a lookup3 +// checksum) whose root object header sits at rootAt. +func h5V3Superblock(n int, rootAt uint64) []byte { + f := make([]byte, n) + copy(f, hdf5Magic) + f[8] = 3 + f[9] = 8 // address size + f[10] = 8 // length size + f[11] = 0 // consistency flags + binary.LittleEndian.PutUint64(f[12:], 0) // base address + binary.LittleEndian.PutUint64(f[20:], math.MaxUint64) // no extension + binary.LittleEndian.PutUint64(f[28:], uint64(n)) // end of file + binary.LittleEndian.PutUint64(f[36:], rootAt) + binary.LittleEndian.PutUint32(f[44:], hdf5Lookup3(f[:44])) + return f +} + +// h5WriteOHDRv2 writes a version 2 object header at at whose first +// message region is chunk, under the given flag byte, computing and +// storing the lookup3 checksum, and returns the offset just past the +// header. Flag-selected prefix fields are written as zeros. +func h5WriteOHDRv2(f []byte, at int, flags byte, chunk []byte) int { + copy(f[at:], hdf5ObjHdr2) + f[at+4] = 2 + f[at+5] = flags + p := at + 6 + if flags&0x20 != 0 { + p += 16 // access, modification, change and birth times + } + if flags&0x10 != 0 { + p += 4 // max compact and min dense attribute counts + } + width := 1 << (flags & 0x03) + switch width { + case 1: + f[p] = byte(len(chunk)) + case 2: + binary.LittleEndian.PutUint16(f[p:], uint16(len(chunk))) + case 4: + binary.LittleEndian.PutUint32(f[p:], uint32(len(chunk))) + case 8: + binary.LittleEndian.PutUint64(f[p:], uint64(len(chunk))) + } + p += width + copy(f[p:], chunk) + p += len(chunk) + binary.LittleEndian.PutUint32(f[p:], hdf5Lookup3(f[at:p])) + return p + 4 +} + +// h5V2Msg renders one version 2 object header message: a type byte, a +// two-byte size and a flag byte, widened by a two-byte creation order +// when the header tracks it, then the body. +func h5V2Msg(typ byte, order uint16, body []byte, ordered bool) []byte { + head := 4 + if ordered { + head = 6 + } + m := make([]byte, head+len(body)) + m[0] = typ + binary.LittleEndian.PutUint16(m[1:], uint16(len(body))) + if ordered { + binary.LittleEndian.PutUint16(m[4:], order) + } + copy(m[head:], body) + return m +} + +// h5V2Dataspace renders a version 2 dataspace message body: rank one, +// the given extent in eight bytes. +func h5V2Dataspace(dim uint64) []byte { + b := make([]byte, 12) + b[0] = 2 // version + b[1] = 1 // rank + binary.LittleEndian.PutUint64(b[4:], dim) + return b +} + +// TestLoadHDF5ShortV1Continuation pins a minimal hostile +// file: a continuation block that ends right after its eight-byte +// message header carries a zero-size continuation message, and the +// reader used to slice the link's offset and length out of bytes past +// the block (a [16:8] panic out of LoadHDF5). It must be a named +// error. +func TestLoadHDF5ShortV1Continuation(t *testing.T) { + f := h5HostileFile(512) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgContinuation, h5Link(136, 8)}, + ) + // The block: exactly eight bytes, one continuation header, no body. + binary.LittleEndian.PutUint16(f[136:], hdf5MsgContinuation) + binary.LittleEndian.PutUint16(f[138:], 0) + _, err := LoadHDF5(writeHostile(t, "v1short.h5", f)) + if err == nil { + t.Fatal("LoadHDF5 accepted a continuation message with no link body") + } + if !strings.Contains(err.Error(), "shorter than an offset and a length") { + t.Fatalf("error = %v, want the short-link refusal", err) + } +} + +// TestLoadHDF5ShortV2Continuation pins the version 2 twin: a header +// whose whole message region is one continuation message of size zero +// panicked the same way ([12:4]), because the stream read the link's +// offset and length past the region. +func TestLoadHDF5ShortV2Continuation(t *testing.T) { + f := h5V3Superblock(512, 96) + chunk := []byte{hdf5MsgContinuation, 0, 0, 0} // type 16, size 0, flags 0 + h5WriteOHDRv2(f, 96, 0, chunk) + _, err := LoadHDF5(writeHostile(t, "v2short.h5", f)) + if err == nil { + t.Fatal("LoadHDF5 accepted a version 2 continuation message with no link body") + } + if !strings.Contains(err.Error(), "shorter than an offset and a length") { + t.Fatalf("error = %v, want the short-link refusal", err) + } +} + +// TestLoadHDF5LocalHeapMixedSizes pins the local heap header size: the +// header holds two length-size fields and then one address-size field, +// but the reader sized it with three address-size fields, so a 4/8 file +// (twenty-byte header) was sliced at [24:] and panicked. A heap whose +// header lies about nothing must still be refused cleanly when what it +// points at is absent. +func TestLoadHDF5LocalHeapMixedSizes(t *testing.T) { + f := h5SizedHostileFile(512, 4, 8, 96) + body := make([]byte, 2*4) + binary.LittleEndian.PutUint32(body, math.MaxUint32) // B-tree: undefined + binary.LittleEndian.PutUint32(body[4:], 256) // local heap at 256 + h5ObjectHeader(f, 96, h5Msg{hdf5MsgSymbolTable, body}) + copy(f[256:], hdf5LocalHeap) + f[260] = 1 // version + _, err := LoadHDF5(writeHostile(t, "heap48.h5", f)) + if err != nil && strings.Contains(err.Error(), "runtime error") { + t.Fatalf("the 4/8 local heap panicked: %v", err) + } + _ = err // any named refusal is fine; the panic is the defect +} + +// TestWalkDiamondReadsOnce pins the diamond: two names on one group +// used to walk the group twice, and a diamond of depth d walked it 2^d +// times, which stopped the reader for hours on a file of a kilobyte. +// The object is read once, under the first path the traversal reaches, +// and a link that closes a cycle along the current path is still an +// error. +func TestWalkDiamondReadsOnce(t *testing.T) { + guard := time.AfterFunc(20*time.Second, func() { panic("diamond walk did not return") }) + defer guard.Stop() + + const childAt, datasetAt, dataAt = 256, 384, 512 + f := h5HostileFile(576) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgLink, h5HardLink("a", childAt)}, + h5Msg{hdf5MsgLink, h5HardLink("b", childAt)}, + ) + h5ObjectHeader(f, childAt, + h5Msg{hdf5MsgLink, h5HardLink("d", datasetAt)}, + ) + h5ObjectHeader(f, datasetAt, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)}, + ) + binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(2.5)) + sets, err := LoadHDF5(writeHostile(t, "diamond.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5 on a diamond of hard links: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want the single dataset once", len(sets)) + } + // Deterministic: the links of a group are visited in file order, so + // the first path wins and the listing is stable. + if sets[0].Path != "/a/d" { + t.Fatalf("path = %q, want /a/d, the first path to the object", sets[0].Path) + } + if got := sets[0].Values.FloatAt(0); got != 2.5 { + t.Fatalf("value = %v, want 2.5", got) + } + t.Run("cycle", func(t *testing.T) { + // The skip must never swallow a cycle: an object that closes a + // loop along the current path is refused, not skipped. + if _, err := LoadHDF5(writeHostile(t, "diamond-cycle.h5", selfLink())); err == nil || !strings.Contains(err.Error(), "hard-link cycle") { + t.Fatalf("a hard-link cycle: err = %v, want the cycle refusal", err) + } + }) +} + +// TestHDF5BudgetChargesDatasetEnvelope pins the aggregate budget: it +// used to deduct only the value bytes, so millions of empty datasets +// would allocate their result structures, paths and attribute maps +// without ever touching the budget. The walk charges a fixed envelope +// plus the dataset's path beside the values, so a budget below one +// envelope refuses even a file of empty datasets. +func TestHDF5BudgetChargesDatasetEnvelope(t *testing.T) { + const dataAt, n = 448, 512 + raw := h5HostileFile(n) + h5ObjectHeader(raw, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)}, + ) + binary.LittleEndian.PutUint64(raw[dataAt:], math.Float64bits(2.5)) + // The file itself is legal: at the default budget it loads. + sets, err := LoadHDF5(writeHostile(t, "envelope.h5", raw)) + if err != nil || len(sets) != 1 { + t.Fatalf("LoadHDF5 at the default budget: %d datasets, err = %v", len(sets), err) + } + // Below one envelope the same file refuses: the envelope plus the + // path is charged before any allocation happens. + f, err := newHDF5File(raw) + if err != nil { + t.Fatalf("newHDF5File: %v", err) + } + st := newHDF5WalkState() + st.budget = hdf5DatasetEnvelope / 2 + var out []HDF5Dataset + if err := f.walk(f.rootAddress, st, &out); err == nil || !strings.Contains(err.Error(), "budget") { + t.Fatalf("walk with a budget below one envelope: err = %v, want the budget refusal", err) + } +} + +// TestLoadHDF5SymbolTable4of4 pins the four-byte-address layout of the +// group structures: the version 1 B-tree header is 8+2*offSize wide +// (not 24), a symbol node entry is lenSize+offSize+24 (not 40) and the +// local heap header is 8+2*lenSize+offSize. A legal 4/4 file with a +// symbol-table group of two links used to die with "no symbol table +// node at 0"; the 8/8 fixtures are untouched by the arithmetic. +func TestLoadHDF5SymbolTable4of4(t *testing.T) { + const ( + treeAt = 200 + snodAt = 240 + heapAt = 320 + segAt = 352 + objA = 384 + objB = 512 + dataA = 640 + dataB = 648 + ) + f := h5SizedHostileFile(1024, 4, 4, 96) + body := make([]byte, 8) + binary.LittleEndian.PutUint32(body, treeAt) + binary.LittleEndian.PutUint32(body[4:], heapAt) + h5ObjectHeader(f, 96, h5Msg{hdf5MsgSymbolTable, body}) + + // The group B-tree: a sixteen-byte header (signature, type, level, + // one entry, two sibling addresses), key0 of four bytes, the child + // address, one trailing key. + copy(f[treeAt:], hdf5Tree) + f[treeAt+4] = 0 // group + f[treeAt+5] = 0 // leaf + binary.LittleEndian.PutUint16(f[treeAt+6:], 1) + binary.LittleEndian.PutUint32(f[treeAt+8:], math.MaxUint32) + binary.LittleEndian.PutUint32(f[treeAt+12:], math.MaxUint32) + binary.LittleEndian.PutUint32(f[treeAt+16:], 0) // key0 + binary.LittleEndian.PutUint32(f[treeAt+20:], snodAt) // child + binary.LittleEndian.PutUint32(f[treeAt+24:], 4) // trailing key + + // The symbol node: two entries of 32 bytes (heap offset, object + // address, cache type, reserved, scratch pad). + copy(f[snodAt:], hdf5SymbolNode) + f[snodAt+4] = 1 + binary.LittleEndian.PutUint16(f[snodAt+6:], 2) + binary.LittleEndian.PutUint32(f[snodAt+8:], 0) // "a" at heap offset 0 + binary.LittleEndian.PutUint32(f[snodAt+12:], objA) + binary.LittleEndian.PutUint32(f[snodAt+40:], 2) // "b" at heap offset 2 + binary.LittleEndian.PutUint32(f[snodAt+44:], objB) + + // The local heap: a twenty-byte header (signature, version, data + // segment size, free-list head, data segment address). + copy(f[heapAt:], hdf5LocalHeap) + f[heapAt+4] = 1 + binary.LittleEndian.PutUint32(f[heapAt+8:], 8) + binary.LittleEndian.PutUint32(f[heapAt+12:], math.MaxUint32) + binary.LittleEndian.PutUint32(f[heapAt+16:], segAt) + copy(f[segAt:], "a\x00b\x00\x00\x00\x00\x00") + + h5ObjectHeader(f, objA, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataA, 8)}, + ) + h5ObjectHeader(f, objB, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataB, 8)}, + ) + binary.LittleEndian.PutUint64(f[dataA:], math.Float64bits(4.5)) + binary.LittleEndian.PutUint64(f[dataB:], math.Float64bits(-2.5)) + + sets, err := LoadHDF5(writeHostile(t, "snod44.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5 refused a legal 4/4 symbol-table file: %v", err) + } + if len(sets) != 2 { + t.Fatalf("datasets = %d, want 2", len(sets)) + } + if sets[0].Path != "/a" || sets[1].Path != "/b" { + t.Fatalf("paths = %q, %q, want /a and /b", sets[0].Path, sets[1].Path) + } + if got := sets[0].Values.FloatAt(0); got != 4.5 { + t.Fatalf("/a = %v, want 4.5", got) + } + if got := sets[1].Values.FloatAt(0); got != -2.5 { + t.Fatalf("/b = %v, want -2.5", got) + } +} + +// TestLoadHDF5ContinuationChainDepth pins the version 1 depth guard: a +// chain of continuation blocks one longer than the cap must be refused +// like the version 2 walk refuses its own, not walked to the end. +func TestLoadHDF5ContinuationChainDepth(t *testing.T) { + const first = 224 + blocks := hdf5MaxHeaderBlocks + 2 + f := h5HostileFile(first + blocks*24) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgContinuation, h5Link(first, 24)}, + ) + for k := range blocks { + at := first + k*24 + binary.LittleEndian.PutUint16(f[at:], hdf5MsgContinuation) + if k == blocks-1 { + // The tail block carries nothing: the walk must be refused + // on reaching it, not on reading it. + binary.LittleEndian.PutUint16(f[at+2:], 0) + continue + } + binary.LittleEndian.PutUint16(f[at+2:], 16) + binary.LittleEndian.PutUint64(f[at+8:], uint64(at+24)) + binary.LittleEndian.PutUint64(f[at+16:], 24) + } + _, err := LoadHDF5(writeHostile(t, "chain.h5", f)) + if err == nil { + t.Fatalf("LoadHDF5 walked a chain of %d continuation blocks", blocks) + } + if !strings.Contains(err.Error(), "continuation blocks") { + t.Fatalf("error = %v, want the chain-depth refusal", err) + } +} + +// TestLoadHDF5V2HeaderFlags covers the version 2 header features the +// reference fixture does not carry: the four times (0x20), the attribute +// counts (0x10) and creation order tracking (0x04, which widens every +// message by its order field). A header with all three set must read +// like any other. +func TestLoadHDF5V2HeaderFlags(t *testing.T) { + const dataAt = 320 + f := h5V3Superblock(512, 96) + msgs := slices.Concat( + h5V2Msg(hdf5MsgDataspace, 1, h5V2Dataspace(2), true), + h5V2Msg(hdf5MsgDatatype, 2, h5FloatType(8), true), + h5V2Msg(hdf5MsgDataLayout, 3, h5ContiguousLayout(dataAt, 16), true), + ) + h5WriteOHDRv2(f, 96, 0x34, msgs) + binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(1.5)) + binary.LittleEndian.PutUint64(f[dataAt+8:], math.Float64bits(-2.5)) + + sets, err := LoadHDF5(writeHostile(t, "v2flags.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5 refused a flagged version 2 header: %v", err) + } + if len(sets) != 1 || sets[0].Path != "/" { + t.Fatalf("datasets = %v, want one dataset at /", sets) + } + if got := sets[0].Values.FloatAt(1); got != -2.5 { + t.Fatalf("value = %v, want -2.5", got) + } +} + +// TestLoadHDF5V2ContinuationChecksum covers the happy path of a version +// 2 continuation: the dataset's messages live in an OCHK block whose +// lookup3 checksum covers the signature and the messages alike, and a +// correct checksum must be accepted (the hostile files above only ever +// see a broken one). +func TestLoadHDF5V2ContinuationChecksum(t *testing.T) { + const blockAt, dataAt = 224, 320 + f := h5V3Superblock(512, 96) + msgs := slices.Concat( + h5V2Msg(hdf5MsgDataspace, 0, h5V2Dataspace(2), false), + h5V2Msg(hdf5MsgDatatype, 0, h5FloatType(8), false), + h5V2Msg(hdf5MsgDataLayout, 0, h5ContiguousLayout(dataAt, 16), false), + ) + block := append([]byte{}, hdf5Chunk2...) + block = append(block, msgs...) + block = binary.LittleEndian.AppendUint32(block, hdf5Lookup3(block)) + copy(f[blockAt:], block) + + link := binary.LittleEndian.AppendUint64( + binary.LittleEndian.AppendUint64([]byte{}, blockAt), uint64(len(block))) + chunk := make([]byte, 4+len(link)) + chunk[0] = hdf5MsgContinuation + binary.LittleEndian.PutUint16(chunk[1:], uint16(len(link))) + copy(chunk[4:], link) + h5WriteOHDRv2(f, 96, 0, chunk) + + binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(1.5)) + binary.LittleEndian.PutUint64(f[dataAt+8:], math.Float64bits(-2.5)) + + sets, err := LoadHDF5(writeHostile(t, "v2cont-ok.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5 refused a checksummed version 2 continuation: %v", err) + } + if len(sets) != 1 || sets[0].Path != "/" { + t.Fatalf("datasets = %v, want one dataset at /", sets) + } + if got := sets[0].Values.FloatAt(0); got != 1.5 { + t.Fatalf("value = %v, want 1.5", got) + } +} + +// TestLoadFITSTableNegativeTFIELDS pins the column count: a negative +// TFIELDS used to size the per-column slices with a negative length and +// panicked, in the binary and the ASCII branch alike. It is refused by +// name in both. +func TestLoadFITSTableNegativeTFIELDS(t *testing.T) { + for _, kind := range []string{"BINTABLE", "TABLE"} { + t.Run(kind, func(t *testing.T) { + hdr := cardBlock( + card("XTENSION= '"+kind+"'"), + card("BITPIX = 8"), + card("NAXIS = 2"), + card("NAXIS1 = 8"), + card("NAXIS2 = 1"), + card("PCOUNT = 0"), + card("GCOUNT = 1"), + card("TFIELDS = -1"), + card("END"), + ) + path := writeHostile(t, "tfields.fits", append(hdr, make([]byte, 2880)...)) + _, err := LoadFITSTable(path) + if err == nil { + t.Fatal("LoadFITSTable accepted TFIELDS = -1") + } + if !strings.Contains(err.Error(), "negative") { + t.Fatalf("error = %v, want the negative-count refusal", err) + } + }) + } +} + +// TestLoadFITSNaxisRefKeyword pins the NAXIS prefix: any keyword that +// started with NAXIS used to count as an axis, so a user keyword +// NAXISREF = 7 answered "NAXIS = 1 with 2 NAXISn cards" and refused a +// legal file. Only a number after the prefix is an axis, the rule the +// writer applies; NAXIS1 keeps counting. +func TestLoadFITSNaxisRefKeyword(t *testing.T) { + hdr := cardBlock( + card("SIMPLE = T"), + card("BITPIX = -64"), + card("NAXIS = 1"), + card("NAXIS1 = 2"), + card("NAXISREF= 7"), + card("END"), + ) + payload := binary.BigEndian.AppendUint64(nil, math.Float64bits(1.5)) + payload = binary.BigEndian.AppendUint64(payload, math.Float64bits(-0.5)) + a, headers, err := LoadFITS(writeHostile(t, "naxisref.fits", append(hdr, payload...))) + if err != nil { + t.Fatalf("LoadFITS refused a file with a NAXISREF keyword: %v", err) + } + if a.Len() != 2 || a.FloatAt(0) != 1.5 || a.FloatAt(1) != -0.5 { + t.Fatalf("values = %v, want 1.5 and -0.5", a) + } + if got := headers["NAXISREF"]; got != "7" { + t.Fatalf("NAXISREF = %q, want it reported as a user keyword", got) + } +} diff --git a/io/hdf5_test.go b/io/hdf5_test.go new file mode 100644 index 0000000..0514374 --- /dev/null +++ b/io/hdf5_test.go @@ -0,0 +1,748 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "math" + "os" + "path/filepath" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The HDF5 fixtures under testdata/h5 were written by the HDF5 +// reference library, and every value below was read back from them +// independently: the expected values in these tests are what the +// reference reports, not what this reader +// produces. +// +// fixture.h5: an int32 dataset stored contiguously, a float64 dataset +// chunked, gzip compressed and shuffled, a float32 dataset +// in a group, and string attributes on the root and the +// group (variable-length, so they live in a global heap) +// fletcher.h5: a float64 dataset chunked, gzip compressed, with the +// fletcher32 checksum filter on top +// latest.h5: written with libver="latest", so superblock version 3 +// and object header version 2, with a float64 dataset /d +// and one /g/e in a group + +func h5Fixture(t *testing.T, name string) string { + t.Helper() + return filepath.Join("testdata", "h5", name) +} + +// TestLoadHDF5Values pins the reader against the reference-written fixture. +func TestLoadHDF5Values(t *testing.T) { + sets, err := LoadHDF5(h5Fixture(t, "fixture.h5")) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 3 { + t.Fatalf("datasets = %d, want 3", len(sets)) + } + byPath := map[string]HDF5Dataset{} + for _, d := range sets { + byPath[d.Path] = d + } + // The paths come back sorted. + if paths := []string{sets[0].Path, sets[1].Path, sets[2].Path}; paths[0] != "/floats" || paths[1] != "/g/f32" || paths[2] != "/ints" { + t.Fatalf("paths = %v, want [/floats /g/f32 /ints]", paths) + } + + ints, ok := byPath["/ints"] + if !ok { + t.Fatal("/ints is missing") + } + if s := ints.Shape; len(s) != 2 || s[0] != 2 || s[1] != 3 { + t.Fatalf("/ints shape = %v, want [2 3]", s) + } + // The fixture's int32 dataset lands the native int32 dtype: the + // reader keeps the width the file stores instead of widening it. + if ints.Values.Dtype() != core.Int32 { + t.Fatalf("/ints dtype = %s, want int32", ints.Values.Dtype()) + } + for i, want := range []int32{1, 2, 3, 4, 5, 6} { + if got := ints.Values.RawInt32s()[i]; got != want { + t.Fatalf("/ints[%d] = %d, want %d", i, got, want) + } + } + + floats, ok := byPath["/floats"] + if !ok { + t.Fatal("/floats is missing") + } + if s := floats.Shape; len(s) != 1 || s[0] != 4 { + t.Fatalf("/floats shape = %v, want [4]", s) + } + if floats.Values.Dtype() != core.Float { + t.Fatalf("/floats dtype = %s, want float64", floats.Values.Dtype()) + } + for i, want := range []float64{1.5, 2.5, 3.5, 4.5} { + if got := floats.Values.RawFloats()[i]; got != want { + t.Fatalf("/floats[%d] = %v, want %v", i, got, want) + } + } + + f32, ok := byPath["/g/f32"] + if !ok { + t.Fatal("/g/f32 is missing") + } + if s := f32.Shape; len(s) != 2 || s[0] != 2 || s[1] != 2 { + t.Fatalf("/g/f32 shape = %v, want [2 2]", s) + } + if f32.Values.Dtype() != core.Float32 { + t.Fatalf("/g/f32 dtype = %s, want float32", f32.Values.Dtype()) + } + for i, want := range []float32{1, 2, 3, 4} { + if got := f32.Values.RawFloat32s()[i]; got != want { + t.Fatalf("/g/f32[%d] = %v, want %v", i, got, want) + } + } + + // The attributes: the root's title reaches every dataset, and the + // group's units reach the dataset inside it, the nearest group + // winning. + if got := ints.Attrs["title"]; got != "h5 fixture" { + t.Fatalf("/ints title = %q, want %q", got, "h5 fixture") + } + if _, ok := ints.Attrs["units"]; ok { + t.Fatalf("/ints picked up a group attribute it should not have: %v", ints.Attrs) + } + if got := f32.Attrs["units"]; got != "K" { + t.Fatalf("/g/f32 units = %q, want K", got) + } + if got := f32.Attrs["title"]; got != "h5 fixture" { + t.Fatalf("/g/f32 title = %q, want the inherited one", got) + } +} + +// TestLoadHDF5Fletcher32 pins the checksum filter: the chunk carries a +// fletcher32 sum that must verify before the chunk is used. +func TestLoadHDF5Fletcher32(t *testing.T) { + sets, err := LoadHDF5(h5Fixture(t, "fletcher.h5")) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want 1", len(sets)) + } + d := sets[0] + if s := d.Shape; len(s) != 1 || s[0] != 20 { + t.Fatalf("shape = %v, want [20]", s) + } + for i := range 20 { + if got := d.Values.RawFloats()[i]; got != float64(i) { + t.Fatalf("value %d = %v, want %d", i, got, i) + } + } +} + +// TestLoadHDF5Latest pins the "latest" file format against the +// reference-written fixture: superblock version 3, object headers +// version 2 with their lookup3 checksums, compact groups carrying link +// messages, and contiguous datasets. +func TestLoadHDF5Latest(t *testing.T) { + sets, err := LoadHDF5(h5Fixture(t, "latest.h5")) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 2 { + t.Fatalf("datasets = %d, want 2", len(sets)) + } + d, e := sets[0], sets[1] + if d.Path != "/d" || e.Path != "/g/e" { + t.Fatalf("paths = %q, %q, want /d and /g/e", d.Path, e.Path) + } + if s := d.Shape; len(s) != 1 || s[0] != 3 { + t.Fatalf("/d shape = %v, want [3]", s) + } + for i, want := range []float64{1, 2, 3} { + if got := d.Values.RawFloats()[i]; got != want { + t.Fatalf("/d[%d] = %v, want %v", i, got, want) + } + } + if s := e.Shape; len(s) != 1 || s[0] != 1 { + t.Fatalf("/g/e shape = %v, want [1]", s) + } + if got := e.Values.RawFloats()[0]; got != 4 { + t.Fatalf("/g/e[0] = %v, want 4", got) + } +} + +// TestHDF5Lookup3 pins the checksum against the sums the reference +// library wrote into the latest fixture: the superblock's and two +// object headers'. The literals are what the file stores, not what +// this implementation computes. +func TestHDF5Lookup3(t *testing.T) { + raw, err := os.ReadFile(h5Fixture(t, "latest.h5")) + if err != nil { + t.Fatal(err) + } + for _, c := range []struct { + name string + want uint32 + lo int + hi int + }{ + {"superblock", 0x39ff1913, 0, 44}, + {"root header", 0xb91c2db3, 48, 175}, + {"dataset header", 0x8d125cb5, 179, 443}, + } { + if got := hdf5Lookup3(raw[c.lo:c.hi]); got != c.want { + t.Errorf("%s: lookup3 = %#08x, want %#08x", c.name, got, c.want) + } + } +} + +// TestLoadHDF5Refusals pins the errors: a file that is not HDF5 at +// all, a truncated file, and corrupted latest-format checksums must +// each be refused with a message that says so, never read halfway. +func TestLoadHDF5Refusals(t *testing.T) { + dir := t.TempDir() + notHDF5 := filepath.Join(dir, "plain.bin") + if err := os.WriteFile(notHDF5, []byte("this is not an HDF5 file at all, not even close"), 0o644); err != nil { + t.Fatal(err) + } + if _, err := LoadHDF5(notHDF5); err == nil { + t.Fatal("expected an error for a file without the HDF5 signature") + } else if !strings.Contains(err.Error(), "signature") { + t.Fatalf("error = %v, want a signature refusal", err) + } + // A truncated copy of a valid file: the reader may accept it only + // when the structures it actually reads are complete, and it must + // never hand back a partial array. Whatever the cut, the call either + // errors or returns datasets whose element count matches their + // shape. + whole, err := os.ReadFile(h5Fixture(t, "fixture.h5")) + if err != nil { + t.Fatal(err) + } + for _, cut := range []int{8, 32, 100, 600, len(whole) / 2, len(whole) - 4} { + path := filepath.Join(dir, "cut.h5") + if err := os.WriteFile(path, whole[:cut], 0o644); err != nil { + t.Fatal(err) + } + sets, err := LoadHDF5(path) + if err != nil { + continue // refused, which is the expected answer + } + for _, d := range sets { + n := 1 + for _, s := range d.Shape { + n *= s + } + if d.Values.Len() != n { + t.Fatalf("a file truncated to %d bytes gave %q %d values for shape %v", + cut, d.Path, d.Values.Len(), d.Shape) + } + } + } + // The header itself must be refused: a superblock shorter than its + // fixed part cannot be read at all. + if _, err := LoadHDF5(writeCut(t, dir, whole, 40)); err == nil { + t.Fatal("expected an error for a file truncated inside the superblock") + } + // The latest format verifies its checksums: a flipped byte in the + // superblock and one in an object header must each refuse the file + // instead of reading past the corruption. + latest, err := os.ReadFile(h5Fixture(t, "latest.h5")) + if err != nil { + t.Fatal(err) + } + for _, c := range []struct { + name string + at int + }{ + // Byte 11 is the superblock's consistency flags, which the + // reader would otherwise ignore: only the checksum sees it. + {"superblock", 11}, + {"object header", 60}, + } { + corrupt := slices.Clone(latest) + corrupt[c.at] ^= 0xff + path := filepath.Join(dir, "corrupt.h5") + if err := os.WriteFile(path, corrupt, 0o644); err != nil { + t.Fatal(err) + } + if _, err := LoadHDF5(path); err == nil { + t.Fatalf("expected an error for a corrupted %s", c.name) + } else if !strings.Contains(err.Error(), "checksum") { + t.Fatalf("corrupted %s: error = %v, want a checksum refusal", c.name, err) + } + } +} + +// writeCut writes the first n bytes of data to a temp file and returns +// its path. +func writeCut(t *testing.T, dir string, data []byte, n int) string { + t.Helper() + path := filepath.Join(dir, "cut40.h5") + if err := os.WriteFile(path, data[:n], 0o644); err != nil { + t.Fatal(err) + } + return path +} + +// h5FixedType renders a version 1 fixed-point datatype message of the +// given element width and signedness. The message carries the bit +// offset and bit precision the HDF5 file format specification's +// fixed-point property table defines behind the eight-byte header, +// twelve bytes in total; the reader keys the landing on the header's +// size and signed bit. +func h5FixedType(size uint32, signed bool) []byte { + m := make([]byte, 12) + m[0] = 0x10 // version 1, class 0 (fixed-point) + if signed { + m[1] = 0x08 // class bit field: bit 3 marks two's complement + } + binary.LittleEndian.PutUint32(m[4:], size) + binary.LittleEndian.PutUint16(m[8:], 0) // bit offset + binary.LittleEndian.PutUint16(m[10:], uint16(8*size)) // bit precision + return m +} + +// h5EnumBoolType renders the boolean enumeration datatype message HDF5 +// writers carry booleans in, following the HDF5 file format +// specification's enumeration class layout: the member count in the +// class bit field, the base type as a complete fixed-point message, +// each member name NUL-terminated and padded from its own field start +// to a multiple of eight bytes, and the packed member values behind +// the names. +func h5EnumBoolType(names []string, values []byte) []byte { + // Version 1, class 8; member count; size 1; then the base type. + m := []byte{0x18, byte(len(names)), 0, 0, 1, 0, 0, 0} + m = append(m, h5FixedType(1, false)...) + for _, n := range names { + start := len(m) + m = append(m, n...) + m = append(m, 0) + for (len(m)-start)%8 != 0 { + m = append(m, 0) + } + } + m = append(m, values...) + return m +} + +// h5AttrMessage renders a version 1 attribute message: the name and +// every field boundary padded to the eight-byte grid the message +// format defines, then the value bytes. +func h5AttrMessage(name string, dtypeMsg []byte, dims []uint64, value []byte) []byte { + space := h5Dataspace(dims...) + nameSize := len(name) + 1 + dtypeAt := alignUp(8+nameSize, 8) + spaceAt := alignUp(dtypeAt+len(dtypeMsg), 8) + b := make([]byte, spaceAt+len(space)+len(value)) + b[0] = 1 + binary.LittleEndian.PutUint16(b[2:], uint16(nameSize)) + binary.LittleEndian.PutUint16(b[4:], uint16(len(dtypeMsg))) + binary.LittleEndian.PutUint16(b[6:], uint16(len(space))) + copy(b[8:], name) // the trailing NUL is the buffer's own zero + copy(b[dtypeAt:], dtypeMsg) + copy(b[spaceAt:], space) + copy(b[spaceAt+len(space):], value) + return b +} + +// h5ChunkTreeWidth writes a one-entry leaf chunk B-tree for a dataset +// of the given rank whose chunk elements are width bytes wide: the +// key's element-size slot must agree with the datatype, which the +// reader checks. +func h5ChunkTreeWidth(f []byte, off, rank int, width uint64, size uint32, chunkAt uint64) { + copy(f[off:], hdf5Tree) + f[off+4] = 1 // chunk tree + f[off+5] = 0 // leaf level + binary.LittleEndian.PutUint16(f[off+6:], 1) + p := off + 24 + binary.LittleEndian.PutUint32(f[p:], size) + // The filter mask stays zero; the chunk offsets stay zero. + binary.LittleEndian.PutUint64(f[p+8+8*rank:], width) + binary.LittleEndian.PutUint64(f[p+8+8*(rank+1):], chunkAt) +} + +// TestLoadHDF5NativeFixedPoint pins the fixed-point landings of the +// contiguous path: every stored width and signedness lands the core +// dtype that holds it exactly, extremes included, and int64 stays int. +func TestLoadHDF5NativeFixedPoint(t *testing.T) { + cases := []struct { + name string + dtype []byte + payload []byte + want core.Dtype + check func(t *testing.T, a *core.Array) + }{ + {"int8", h5FixedType(1, true), []byte{0x80, 0x00, 0x7f}, core.Int8, + func(t *testing.T, a *core.Array) { + if got, want := a.RawInt8s()[:3], []int8{-128, 0, 127}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"uint8", h5FixedType(1, false), []byte{0x00, 0x01, 0xff}, core.Uint8, + func(t *testing.T, a *core.Array) { + if got, want := a.RawUint8s()[:3], []uint8{0, 1, 255}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"int16", h5FixedType(2, true), []byte{0x00, 0x80, 0xff, 0xff, 0xff, 0x7f}, core.Int16, + func(t *testing.T, a *core.Array) { + if got, want := a.RawInt16s()[:3], []int16{-32768, -1, 32767}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"uint16", h5FixedType(2, false), []byte{0x00, 0x00, 0x00, 0x10, 0xff, 0xff}, core.Uint16, + func(t *testing.T, a *core.Array) { + if got, want := a.RawUint16s()[:3], []uint16{0, 4096, 65535}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"int32", h5FixedType(4, true), + []byte{0x00, 0x00, 0x00, 0x80, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0x7f}, core.Int32, + func(t *testing.T, a *core.Array) { + if got, want := a.RawInt32s()[:3], []int32{-2147483648, -1, 2147483647}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"uint32", h5FixedType(4, false), + []byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x40, 0xff, 0xff, 0xff, 0xff}, core.Uint32, + func(t *testing.T, a *core.Array) { + if got, want := a.RawUint32s()[:3], []uint32{0, 1 << 30, 4294967295}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"int64 stays int", h5FixedType(8, true), + []byte{0xfb, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, + 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0x01, 0, 0}, core.Int, + func(t *testing.T, a *core.Array) { + if got, want := a.RawInts()[:3], []int64{-5, 0, 1 << 40}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + f := h5HostileFile(512) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(3)}, + h5Msg{hdf5MsgDatatype, tc.dtype}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, uint64(len(tc.payload)))}, + ) + copy(f[448:], tc.payload) + sets, err := LoadHDF5(writeHostile(t, "native.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want 1", len(sets)) + } + d := sets[0] + if d.Values.Dtype() != tc.want { + t.Fatalf("dtype = %s, want %s", d.Values.Dtype(), tc.want) + } + if s := d.Shape; len(s) != 1 || s[0] != 3 { + t.Fatalf("shape = %v, want [3]", s) + } + tc.check(t, d.Values) + }) + } +} + +// TestLoadHDF5ChunkedNativeLandings pins the chunked dispatch: the +// per-cell decode lands the same native dtypes the contiguous path +// lands, through the chunk B-tree and the placement walk. +func TestLoadHDF5ChunkedNativeLandings(t *testing.T) { + cases := []struct { + name string + dtype []byte + width uint64 + payload []byte + want core.Dtype + check func(t *testing.T, a *core.Array) + }{ + {"uint8", h5FixedType(1, false), 1, []byte{0, 1, 255, 42}, core.Uint8, + func(t *testing.T, a *core.Array) { + if got, want := a.RawUint8s()[:4], []uint8{0, 1, 255, 42}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"int16", h5FixedType(2, true), 2, + []byte{0xfd, 0xff, 0x00, 0x80, 0xff, 0x7f, 0x07, 0x00}, core.Int16, + func(t *testing.T, a *core.Array) { + if got, want := a.RawInt16s()[:4], []int16{-3, -32768, 32767, 7}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + {"bool", h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), 1, + []byte{1, 0, 1, 1}, core.Bool, + func(t *testing.T, a *core.Array) { + if got, want := a.RawBools()[:4], []bool{true, false, true, true}; !slices.Equal(got, want) { + t.Fatalf("values = %v, want %v", got, want) + } + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + const btree, chunkAt = 256, 320 + f := h5HostileFile(512) + end := h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(4)}, + h5Msg{hdf5MsgDatatype, tc.dtype}, + h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(btree, 4, uint32(tc.width))}, + ) + if end > btree { + t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree) + } + h5ChunkTreeWidth(f, btree, 1, tc.width, uint32(len(tc.payload)), chunkAt) + copy(f[chunkAt:], tc.payload) + sets, err := LoadHDF5(writeHostile(t, "chunk-native.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want 1", len(sets)) + } + d := sets[0] + if d.Values.Dtype() != tc.want { + t.Fatalf("dtype = %s, want %s", d.Values.Dtype(), tc.want) + } + tc.check(t, d.Values) + }) + } +} + +// enumBoolLoad builds a one-dataset contiguous file around a datatype +// message and payload, and returns the load error or the dataset. +func enumBoolLoad(t *testing.T, dtypeMsg, payload []byte) ([]HDF5Dataset, error) { + t.Helper() + f := h5HostileFile(512) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(uint64(len(payload)))}, + h5Msg{hdf5MsgDatatype, dtypeMsg}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, uint64(len(payload)))}, + ) + copy(f[448:], payload) + return LoadHDF5(writeHostile(t, "enum.h5", f)) +} + +// TestLoadHDF5EnumBoolLandings pins the boolean enumeration landing: +// a one-byte unsigned base whose member values are a subset of {0, 1} +// lands core.Bool whatever the member names say, because the values, +// not the names, carry the semantics. +func TestLoadHDF5EnumBoolLandings(t *testing.T) { + cases := []struct { + name string + dtype []byte + values []byte + want []bool + }{ + {"members TRUE and FALSE", + h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), + []byte{1, 0, 1}, []bool{true, false, true}}, + {"names are irrelevant to the values", + h5EnumBoolType([]string{"present", "absent"}, []byte{0, 1}), + []byte{0, 1, 1}, []bool{false, true, true}}, + {"a single member of zero", + h5EnumBoolType([]string{"off"}, []byte{0}), + []byte{0, 0, 0}, []bool{false, false, false}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + sets, err := enumBoolLoad(t, tc.dtype, tc.values) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 { + t.Fatalf("datasets = %d, want 1", len(sets)) + } + d := sets[0] + if d.Values.Dtype() != core.Bool { + t.Fatalf("dtype = %s, want bool", d.Values.Dtype()) + } + if got := d.Values.RawBools()[:len(tc.want)]; !slices.Equal(got, tc.want) { + t.Fatalf("values = %v, want %v", got, tc.want) + } + }) + } +} + +// TestLoadHDF5EnumRefusals pins the loud refusals: every enumeration +// outside the boolean convention, every bit field, and a boolean +// payload cell outside the members are named errors, never a silent +// guess. The variants mutate the spec-shaped message, which also pins +// the field offsets the parser reads. +func TestLoadHDF5EnumRefusals(t *testing.T) { + signed := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[9] |= 0x08; return m } + bigEndian := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[9] |= 0x01; return m } + baseClass := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[8] = 0x11; return m } + baseSize := func() []byte { + m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}) + binary.LittleEndian.PutUint32(m[12:], 2) + return m + } + valueSize := func() []byte { + m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}) + binary.LittleEndian.PutUint32(m[4:], 2) + return m + } + noMembers := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[1] = 0; return m } + reserved := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[3] = 0x04; return m } + shortNames := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[1] = 3; return m } + bitField := func() []byte { m := h5FixedType(1, false); m[0] = 0x14; return m } + + cases := []struct { + name string + dtype []byte + payload []byte + want string + }{ + {"a member value outside {0, 1}", h5EnumBoolType([]string{"A", "B"}, []byte{0, 2}), []byte{0, 1}, "outside the boolean convention"}, + {"a signed base type", signed(), []byte{1, 0}, "signed base type"}, + {"a big-endian base type", bigEndian(), []byte{1, 0}, "big-endian"}, + {"a non-fixed-point base type", baseClass(), []byte{1, 0}, "base type of class 1"}, + {"a base type wider than one byte", baseSize(), []byte{1, 0}, "base type of 2 bytes"}, + {"values wider than one byte", valueSize(), []byte{1, 0}, "2-byte values"}, + {"no members at all", noMembers(), []byte{0}, "declares 0 members"}, + {"reserved bit field bits", reserved(), []byte{1, 0}, "unknown bit field bits"}, + {"more members than names", shortNames(), []byte{1, 0}, "ends inside"}, + {"a bit field datatype", bitField(), []byte{0}, "bit field"}, + {"a payload cell outside the members", + h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), []byte{1, 7, 0}, "outside the members 0 and 1"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + sets, err := enumBoolLoad(t, tc.dtype, tc.payload) + if err == nil { + t.Fatalf("LoadHDF5 accepted %s: %d datasets, %v", tc.name, len(sets), sets) + } + if !strings.Contains(err.Error(), tc.want) { + t.Fatalf("error = %v, want it to carry %q", err, tc.want) + } + }) + } + + // The same refusal on the chunked path, where the per-cell decode + // runs inside the placement walk. + t.Run("a chunked payload cell outside the members", func(t *testing.T) { + const btree, chunkAt = 256, 320 + f := h5HostileFile(512) + end := h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(4)}, + h5Msg{hdf5MsgDatatype, h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0})}, + h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(btree, 4, 1)}, + ) + if end > btree { + t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree) + } + h5ChunkTreeWidth(f, btree, 1, 1, 4, chunkAt) + copy(f[chunkAt:], []byte{1, 9, 0, 1}) + _, err := LoadHDF5(writeHostile(t, "enum-chunk.h5", f)) + if err == nil || !strings.Contains(err.Error(), "outside the members 0 and 1") { + t.Fatalf("chunked enum payload of 9: err = %v, want the member refusal", err) + } + }) +} + +// TestLoadHDF5Uint64Refused pins the unsigned 64-bit refusal at both +// sites that hold it: the dataset gate and the value decode, with the +// same text at each. +func TestLoadHDF5Uint64Refused(t *testing.T) { + const want = "unsigned 64-bit integers have no exact core dtype" + f := h5HostileFile(512) + h5ObjectHeader(f, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FixedType(8, false)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, 8)}, + ) + copy(f[448:], []byte{1, 0, 0, 0, 0, 0, 0, 0}) + _, err := LoadHDF5(writeHostile(t, "uint64.h5", f)) + if err == nil || !strings.Contains(err.Error(), want) { + t.Fatalf("LoadHDF5 on an unsigned 64-bit dataset: err = %v, want it to carry %q", err, want) + } + // The decode site, called directly: the same text, no widening. + if _, err := arrayFromRaw(make([]byte, 8), hdf5Type{class: 0, size: 8, width: 8}, []int{1}); err == nil || + !strings.Contains(err.Error(), want) { + t.Fatalf("arrayFromRaw on unsigned 64-bit bytes: err = %v, want it to carry %q", err, want) + } + // The chunked dispatch, through a dataset fixture: the chunk + // dimensions carry the element size in their last slot, matching + // the datatype, and the chunk B-tree address points past the end + // of the file, so the pin also records that the dataset gate + // refuses the datatype before any storage or tree is read. + cf := h5HostileFile(512) + h5ObjectHeader(cf, 96, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FixedType(8, false)}, + h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(1024, 1, 8)}, + ) + if _, err := LoadHDF5(writeHostile(t, "uint64-chunked.h5", cf)); err == nil || + !strings.Contains(err.Error(), want) { + t.Fatalf("LoadHDF5 on a chunked unsigned 64-bit dataset: err = %v, want it to carry %q", err, want) + } + // The chunkedArray dispatch itself, called directly with width 8 + // unsigned: the same refusal, reached before any chunk walk. + var fh hdf5File + if _, err := fh.chunkedArray("/u64", []int{1}, + hdf5Type{class: 0, size: 8, width: 8}, + hdf5Layout{class: 2, dims: []int{1}}, nil, 8); err == nil || + !strings.Contains(err.Error(), want) { + t.Fatalf("chunkedArray on unsigned 64-bit bytes: err = %v, want it to carry %q", err, want) + } +} + +// TestLoadHDF5AttributeSignedRendering pins the attribute text of +// numeric attributes: the datatype's own signed bit decides how its +// stored bits read, the boolean enumeration renders 0 and 1, unsigned +// values keep every digit, and a cell outside the boolean members +// drops the attribute instead of guessing. +func TestLoadHDF5AttributeSignedRendering(t *testing.T) { + const datasetAt, dataAt = 768, 832 + f := h5HostileFile(896) + msgs := []h5Msg{{hdf5MsgLink, h5HardLink("d", datasetAt)}} + // The dataspace carries the element count; the value bytes follow + // it as count many datatype-width cells. + add := func(name string, dtypeMsg, value []byte, elems uint64) { + msgs = append(msgs, h5Msg{hdf5MsgAttribute, h5AttrMessage(name, dtypeMsg, []uint64{elems}, value)}) + } + add("s8", h5FixedType(1, true), []byte{0xff}, 1) + add("u8", h5FixedType(1, false), []byte{0xff}, 1) + add("s16", h5FixedType(2, true), []byte{0xfe, 0xff}, 1) + add("u32", h5FixedType(4, false), []byte{0xff, 0xff, 0xff, 0xff}, 1) + add("flag", h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), []byte{1, 0}, 2) + add("bad", h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), []byte{1, 7}, 2) + end := h5ObjectHeader(f, 96, msgs...) + if end > datasetAt { + t.Fatalf("the root header runs to %d, past the dataset at %d", end, datasetAt) + } + h5ObjectHeader(f, datasetAt, + h5Msg{hdf5MsgDataspace, h5Dataspace(1)}, + h5Msg{hdf5MsgDatatype, h5FloatType(8)}, + h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)}, + ) + binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(2.5)) + sets, err := LoadHDF5(writeHostile(t, "attrs-signed.h5", f)) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if len(sets) != 1 || sets[0].Path != "/d" { + t.Fatalf("datasets = %v, want the linked /d", sets) + } + attrs := sets[0].Attrs + for k, want := range map[string]string{ + "s8": "-1", "u8": "255", "s16": "-2", "u32": "4294967295", "flag": "[1, 0]", + } { + if got := attrs[k]; got != want { + t.Fatalf("attr %s = %q, want %q", k, got, want) + } + } + if v, ok := attrs["bad"]; ok { + t.Fatalf("the attribute with a cell outside the members was accepted as %q", v) + } + if got := sets[0].Values.FloatAt(0); got != 2.5 { + t.Fatalf("value = %v, want 2.5", got) + } +} diff --git a/io/hdf5write.go b/io/hdf5write.go new file mode 100644 index 0000000..1f3c957 --- /dev/null +++ b/io/hdf5write.go @@ -0,0 +1,786 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +// HDF5 writing, the mirror of the reader in hdf5.go. The reader's +// verified decoders are the specification: every structure here is +// written in the shape the reader accepts and in the shape the +// reference library writes, as the fixtures under testdata/h5 pin it. +// +// Written: superblock version 0 (the classic layout, the default) and +// version 3 (the "latest" layout, whose superblock and object headers +// carry lookup3 checksums), object headers versions 1 and 2, groups +// stored as symbol tables (local heap, version 1 group B-tree, symbol +// table nodes) in the classic layout or as link messages in the latest +// one, datasets stored contiguously or in chunks through a version 1 +// chunk B-tree, the deflate and shuffle filters, fixed-point and +// floating-point datatypes of the usual widths, the boolean +// enumeration convention, fixed-length string datatypes, and +// attributes in the object header. +// +// Every address and length is eight bytes, as in the fixtures. The +// output is deterministic: children are written in sorted name order +// and nothing depends on map iteration. + +import ( + "bytes" + "compress/zlib" + "encoding/binary" + "maps" + "math" + "os" + "slices" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// HDF5WriteOptions tunes SaveHDF5 and SaveHDF5Text. The zero value +// writes the classic file layout with contiguous datasets, which every +// reader of the format understands. +type HDF5WriteOptions struct { + // Latest writes superblock version 3 with version 2 object + // headers: groups become link messages and every structure the + // format checksums carries a lookup3 sum. Latest files hold + // contiguous datasets only: the reference library stores filtered + // chunks of a latest file in a version 2 B-tree, which LoadHDF5 + // does not read, so combining Latest with a filter is refused. + Latest bool + + // Gzip applies the deflate filter to every numeric dataset at the + // given level: 0 (the default) disables it, -1 means the default + // level and 1 to 9 are the levels of the format. A filtered + // dataset is stored in chunks. + Gzip int + + // Shuffle applies the shuffle filter before deflate, which + // regroups the bytes of each element so compression sees the + // high-order bytes together. Shuffle alone also forces chunks. + Shuffle bool + + // ChunkBytes is the target size of one chunk in bytes for + // filtered datasets; 0 selects a default of 64 KiB. Datasets + // smaller than the target stay in one chunk. + ChunkBytes int +} + +// SaveHDF5 writes the datasets as an HDF5 file: the mirror of +// LoadHDF5. The paths build the group tree (the dataset "/g/f32" sits +// in the group "/g"), so the file reads back with the same paths, +// shapes, dtypes and values. Each dataset's Attrs are written on the +// dataset itself; the attributes of the root and of the groups come +// from groupAttrs, keyed by group path with the root keyed "/". +// +// Every dtype the writer stores lands the same dtype through LoadHDF5: +// bool through the HDF5 boolean enumeration convention, the narrow +// integers at their stored width and signedness, float32, float64 and +// int64 directly. Float16 and complex arrays are refused: LoadHDF5 +// decodes neither a two-byte floating-point nor a complex datatype, so +// the writer refuses them rather than write a file this package cannot +// read back. Attribute values are parsed back into typed attributes: a +// whole number becomes an int64 attribute, a decimal a float64 one, a +// bracketed list an int64 or float64 array, and anything else a +// fixed-length string, so a file written from LoadHDF5's own output +// reads back with the same attribute text. When several option values +// are passed the last one wins. +func SaveHDF5(path string, datasets []HDF5Dataset, groupAttrs map[string]map[string]string, opts ...HDF5WriteOptions) error { + const name = "SaveHDF5" + options := HDF5WriteOptions{} + for _, o := range opts { + options = o + } + if err := hdf5CheckOptions(name, options); err != nil { + return err + } + root, err := hdf5BuildPlan(name, datasets, nil, groupAttrs) + if err != nil { + return err + } + w := &hdf5Writer{latest: options.Latest, opts: options} + if err := w.write(root); err != nil { + return err + } + if err := os.WriteFile(path, w.buf, 0o644); err != nil { + return base.Errf("%s: %w", name, err) + } + return nil +} + +// HDF5TextDataset is one fixed-length string dataset for SaveHDF5Text: +// Text holds the elements in row-major order, each padded to the +// longest element in the file. The reader of this package refuses +// string datasets (it reads numeric arrays only), so these files are +// for other readers of the format. +type HDF5TextDataset struct { + Path string + Shape []int + Text []string +} + +// SaveHDF5Text writes fixed-length string datasets, the string side of +// the datatype message the reader refuses for data but accepts for +// attributes. The datasets are stored contiguously; the deflate and +// shuffle filters, which are chunked-storage filters, are refused +// here by name rather than silently dropped. When several option +// values are passed the last one wins. +func SaveHDF5Text(path string, texts []HDF5TextDataset, opts ...HDF5WriteOptions) error { + const name = "SaveHDF5Text" + options := HDF5WriteOptions{} + for _, o := range opts { + options = o + } + if options.Gzip != 0 || options.Shuffle { + return base.Errf("%s: the deflate and shuffle filters apply to numeric datasets; string datasets are written contiguously", name) + } + if err := hdf5CheckOptions(name, options); err != nil { + return err + } + root, err := hdf5BuildPlan(name, nil, texts, nil) + if err != nil { + return err + } + w := &hdf5Writer{latest: options.Latest, opts: options} + if err := w.write(root); err != nil { + return err + } + if err := os.WriteFile(path, w.buf, 0o644); err != nil { + return base.Errf("%s: %w", name, err) + } + return nil +} + +// hdf5CheckOptions rejects the option combinations the writer cannot +// honour, naming each one. +func hdf5CheckOptions(name string, opts HDF5WriteOptions) error { + switch opts.Gzip { + case 0, -1, 1, 2, 3, 4, 5, 6, 7, 8, 9: + default: + return base.Errf("%s: gzip level %d: use 0 to disable the filter, -1 for the default level or 1 to 9", name, opts.Gzip) + } + if opts.ChunkBytes < 0 { + return base.Errf("%s: a chunk target of %d bytes is negative", name, opts.ChunkBytes) + } + if opts.Latest && (opts.Gzip != 0 || opts.Shuffle) { + return base.Errf("%s: the latest format stores filtered chunks through a version 2 B-tree, which LoadHDF5 does not read; write filtered datasets to a classic file", name) + } + return nil +} + +// The B-tree and symbol node fanouts the classic layout declares in +// its superblock: a node of the format holds twice the K of its kind, +// so a symbol table node holds eight entries, a group B-tree node +// thirty-two children and a chunk B-tree node (whose K the format +// fixes at thirty-two for superblock version 0) sixty-four chunks. +const ( + hdf5GroupLeafK = 4 + hdf5GroupInnerK = 16 + hdf5IStoreK = 32 + hdf5ChunkTarget = 64 << 10 + hdf5MaxChunks = 4 << 20 + hdf5MaxAttrs = 4096 + hdf5MaxRank = 32 + hdf5MaxNameBytes = 4096 +) + +// hdf5OutSet is one dataset of the plan: the source of its payload, +// its shape and its datatype class (0 fixed-point, 1 floating-point, +// 3 string, 8 the boolean enumeration). +type hdf5OutSet struct { + path string + shape []int + class byte + width int + // signed marks a fixed-point payload as two's complement: the class + // bit field's bit 0x08, the bit the datatype message carries and + // the reader keys its landing on. + signed bool + // nbytes is the payload's byte size, which the plan has bounded. + nbytes int + // The source lives in exactly one of the payload fields below, the + // one its class, width and signedness name: a contiguous write + // encodes the values straight into the image and a chunked write + // stages them chunk by chunk, so no encoded copy of the payload is + // built. Every numeric payload is serialised little-endian at its + // own stored width, bool as one zero-or-one byte per element. + bools []bool + ints []int64 + i8s []int8 + u8s []uint8 + i16s []int16 + u16s []uint16 + i32s []int32 + u32s []uint32 + f32s []float32 + f64s []float64 + texts []string + attrs []hdf5Attr + // written state + addr uint64 +} + +// hdf5OutNode is one group of the plan: the root is the node whose +// path is "/". Groups sort their children by name before writing so +// the heap offsets and the B-tree order agree. +type hdf5OutNode struct { + path string + name string + attrs []hdf5Attr + groups []*hdf5OutNode + sets []*hdf5OutSet + // written state: the object header address and, in the classic + // layout, the group's B-tree and local heap. + addr uint64 + btree uint64 + heap uint64 +} + +// hdf5BuildPlan validates the paths, the dtypes and the attribute +// texts and lays the file out as a tree: the datasets under their +// groups, every attribute sorted by name. Anything the writer would +// refuse it refuses here, before a byte is written. +func hdf5BuildPlan(name string, datasets []HDF5Dataset, texts []HDF5TextDataset, groupAttrs map[string]map[string]string) (*hdf5OutNode, error) { + root := &hdf5OutNode{path: "/"} + groups := map[string]*hdf5OutNode{"/": root} + used := map[string]bool{"/": true} + var ensure func(path string) (*hdf5OutNode, error) + ensure = func(path string) (*hdf5OutNode, error) { + if g, ok := groups[path]; ok { + return g, nil + } + segs, err := hdf5CheckPath(name, path, "group") + if err != nil { + return nil, err + } + parentPath := "/" + if len(segs) > 1 { + parentPath = "/" + strings.Join(segs[:len(segs)-1], "/") + } + parent, err := ensure(parentPath) + if err != nil { + return nil, err + } + if used[path] { + return nil, base.Errf("%s: %q names both a dataset and a group", name, path) + } + g := &hdf5OutNode{path: path, name: segs[len(segs)-1]} + parent.groups = append(parent.groups, g) + groups[path] = g + used[path] = true + return g, nil + } + place := func(path string) (*hdf5OutNode, error) { + segs, err := hdf5CheckPath(name, path, "dataset") + if err != nil { + return nil, err + } + if used[path] { + return nil, base.Errf("%s: the path %q is written twice", name, path) + } + parentPath := "/" + if len(segs) > 1 { + parentPath = "/" + strings.Join(segs[:len(segs)-1], "/") + } + parent, err := ensure(parentPath) + if err != nil { + return nil, err + } + used[path] = true + return parent, nil + } + for i := range datasets { + d := &datasets[i] + parent, err := place(d.Path) + if err != nil { + return nil, err + } + s, err := hdf5PlanSet(name, d) + if err != nil { + return nil, err + } + parent.sets = append(parent.sets, s) + } + for i := range texts { + tx := &texts[i] + parent, err := place(tx.Path) + if err != nil { + return nil, err + } + s, err := hdf5PlanText(name, tx) + if err != nil { + return nil, err + } + parent.sets = append(parent.sets, s) + } + keys := slices.Sorted(maps.Keys(groupAttrs)) + for _, k := range keys { + g, ok := groups[k] + if !ok { + return nil, base.Errf("%s: the attribute path %q does not name a group of this file", name, k) + } + attrs, err := hdf5PlanAttrs(name, k, groupAttrs[k]) + if err != nil { + return nil, err + } + g.attrs = attrs + } + hdf5SortNode(root) + return root, nil +} + +// hdf5CheckPath splits an absolute object path into its segments and +// refuses what no file should carry: a relative path, the root path +// where an object is wanted, an empty segment, a segment holding a +// NUL byte or a path longer than the bound a sane file keeps. +func hdf5CheckPath(name, path, kind string) ([]string, error) { + if path == "" || path[0] != '/' { + return nil, base.Errf("%s: %s path %q is not absolute", name, kind, path) + } + if path == "/" { + return nil, base.Errf("%s: the root path does not name a %s", name, kind) + } + if len(path) > hdf5MaxNameBytes { + return nil, base.Errf("%s: %s path %q is longer than %d bytes", name, kind, path, hdf5MaxNameBytes) + } + segs := strings.Split(path[1:], "/") + for _, s := range segs { + if s == "" { + return nil, base.Errf("%s: %s path %q has an empty segment", name, kind, path) + } + if strings.IndexByte(s, 0) >= 0 { + return nil, base.Errf("%s: %s path %q holds a NUL byte", name, kind, path) + } + } + return segs, nil +} + +// hdf5SortNode orders the children of a group by name, recursively, so +// the written file is independent of the order the caller supplied. +func hdf5SortNode(g *hdf5OutNode) { + slices.SortFunc(g.groups, func(a, b *hdf5OutNode) int { return strings.Compare(a.name, b.name) }) + slices.SortFunc(g.sets, func(a, b *hdf5OutSet) int { return strings.Compare(a.path, b.path) }) + for _, sub := range g.groups { + hdf5SortNode(sub) + } +} + +// hdf5ImageEstimate bounds the byte size of the image the plan writes, +// the number the buffer is allocated from. The bound is loose on +// purpose: every structure the format wraps around a payload is +// charged a fixed frame, a deflated chunk cannot grow past its own +// bytes, and a chunked dataset is charged one chunk of padding plus a +// node per sixty-four chunks. Over-estimating costs the memory the +// write frees again; under-estimating costs one reallocation. +func hdf5ImageEstimate(g *hdf5OutNode, opts HDF5WriteOptions) int { + // The sum is kept in int64 and clamped to what an int holds, so no + // pile of frames can wrap it into a negative capacity. + return int(min(hdf5NodeEstimate(g, opts), maxInt)) +} + +// maxInt is the largest value an int holds on this platform. +const maxInt = int64(^uint(0) >> 1) + +// hdf5NodeEstimate sums one group's own frame, its attributes and the +// datasets and subgroups beneath it. +func hdf5NodeEstimate(g *hdf5OutNode, opts HDF5WriteOptions) int64 { + total := int64(4096+128*(len(g.groups)+len(g.sets))) + hdf5AttrsEstimate(g.attrs) + for _, s := range g.sets { + total += hdf5SetEstimate(s, opts) + } + for _, sub := range g.groups { + total += hdf5NodeEstimate(sub, opts) + } + return total +} + +// hdf5SetEstimate bounds the image bytes one dataset occupies: its +// stored payload, its object header and, when it is chunked, the +// padding, the chunk B-tree nodes and the filter that shrinks it. +func hdf5SetEstimate(s *hdf5OutSet, opts HDF5WriteOptions) int64 { + total := int64(s.nbytes+s.nbytes/64) + 4096 + if len(s.shape) == 0 || s.nbytes == 0 || s.class == 3 || (opts.Gzip == 0 && !opts.Shuffle) { + return total + } + chunkTarget := hdf5ChunkTargetOf(opts) + // The node count is a starting hint the write grows past by append, + // never a bound it must honour, so it is capped at what a 512-byte + // target would charge: a pathological chunk target of one or two + // bytes would otherwise charge a node, and its kilobyte of image, + // to every single data byte up front. + nodes := 1 + min(s.nbytes/max(chunkTarget/2, 1), s.nbytes/512+1, 1<<20) + return total + int64(chunkTarget) + int64(nodes)*int64(hdf5ChunkNodeSize(len(s.shape))) +} + +// hdf5AttrsEstimate bounds the attribute messages of one object: every +// value's bytes are at most four times the text they are parsed from, +// and the fields around them are charged a fixed frame. +func hdf5AttrsEstimate(attrs []hdf5Attr) int64 { + var total int64 + for _, a := range attrs { + total += int64(4*(len(a.name)+len(a.text)) + 256) + } + return total +} + +// hdf5PlanSet validates one numeric dataset and keeps its values for +// the write, which encodes them little-endian, the byte order the +// datatype message declares. +func hdf5PlanSet(name string, d *HDF5Dataset) (*hdf5OutSet, error) { + if d.Values == nil { + return nil, base.Errf("%s: dataset %q has no values", name, d.Path) + } + var class byte + var width int + var signed bool + var bools []bool + var ints []int64 + var i8s []int8 + var u8s []uint8 + var i16s []int16 + var u16s []uint16 + var i32s []int32 + var u32s []uint32 + var f32s []float32 + var f64s []float64 + switch d.Values.Dtype() { + case core.Bool: + class, width, bools = 8, 1, d.Values.RawBools() + case core.Int8: + class, width, signed, i8s = 0, 1, true, d.Values.RawInt8s() + case core.Uint8: + class, width, u8s = 0, 1, d.Values.RawUint8s() + case core.Int16: + class, width, signed, i16s = 0, 2, true, d.Values.RawInt16s() + case core.Uint16: + class, width, u16s = 0, 2, d.Values.RawUint16s() + case core.Int32: + class, width, signed, i32s = 0, 4, true, d.Values.RawInt32s() + case core.Uint32: + class, width, u32s = 0, 4, d.Values.RawUint32s() + case core.Int: + class, width, signed, ints = 0, 8, true, d.Values.RawInts() + case core.Float32: + class, width, f32s = 1, 4, d.Values.RawFloat32s() + case core.Float: + class, width, f64s = 1, 8, d.Values.RawFloats() + default: + return nil, base.Errf("%s: dataset %q: dtype %s is not supported; the writer stores bool, int8, uint8, int16, uint16, int32, uint32, int64, float32 and float64", name, d.Path, d.Values.Dtype()) + } + shape := d.Values.Shape() + if d.Shape != nil && !slices.Equal(d.Shape, shape) { + return nil, base.Errf("%s: dataset %q declares a shape of %v for values shaped %v", name, d.Path, d.Shape, shape) + } + if len(shape) > hdf5MaxRank { + return nil, base.Errf("%s: dataset %q has %d dimensions, the format allows %d", name, d.Path, len(shape), hdf5MaxRank) + } + // The payload's byte size is bounded here, in the plan, so the write + // can reserve it whole: this check is what stands between a wrapped + // shape and a reservation past the budget. + n, err := hdf5ByteExtent(shape, width, hdf5MaxDatasetBytes) + if err != nil { + return nil, base.Errf("%s: dataset %q: %w", name, d.Path, err) + } + attrs, err := hdf5PlanAttrs(name, d.Path, d.Attrs) + if err != nil { + return nil, err + } + return &hdf5OutSet{ + path: d.Path, shape: shape, class: class, width: width, signed: signed, + nbytes: int(n), bools: bools, ints: ints, i8s: i8s, u8s: u8s, + i16s: i16s, u16s: u16s, i32s: i32s, u32s: u32s, + f32s: f32s, f64s: f64s, attrs: attrs, + }, nil +} + +// hdf5PlanText validates one string dataset and pads its elements to +// the longest one, the fixed length the datatype message declares. +func hdf5PlanText(name string, tx *HDF5TextDataset) (*hdf5OutSet, error) { + if len(tx.Shape) > hdf5MaxRank { + return nil, base.Errf("%s: dataset %q has %d dimensions, the format allows %d", name, tx.Path, len(tx.Shape), hdf5MaxRank) + } + n := 1 + for _, d := range tx.Shape { + if d < 0 { + return nil, base.Errf("%s: dataset %q has the negative extent %d", name, tx.Path, d) + } + n *= d + } + if len(tx.Text) != n { + return nil, base.Errf("%s: dataset %q holds %d strings for a shape of %d elements", name, tx.Path, len(tx.Text), n) + } + width := 1 + for _, s := range tx.Text { + if strings.IndexByte(s, 0) >= 0 { + return nil, base.Errf("%s: dataset %q holds a string with a NUL byte, which a fixed-length element cannot carry", name, tx.Path) + } + width = max(width, len(s)) + } + // The same byte budget the numeric plan answers to: a shape whose + // extents wrap the element count onto len(nil) would otherwise pass + // the length check and write a header declaring data it does not + // store. + if _, err := hdf5ByteExtent(tx.Shape, width, hdf5MaxDatasetBytes); err != nil { + return nil, base.Errf("%s: dataset %q: %w", name, tx.Path, err) + } + // Every element occupies the fixed width; the write lays the + // strings into the image itself, where the padding behind each is + // the zero the buffer already holds. + return &hdf5OutSet{path: tx.Path, shape: tx.Shape, class: 3, width: width, nbytes: n * width, texts: tx.Text}, nil +} + +// encode lays the dataset's payload into dst, which holds exactly the +// nbytes the plan bounded: every numeric value little-endian at its own +// stored width, bool as one zero-or-one byte per element, the strings +// each into the fixed-width slot its index names. The zeros dst arrives +// with are the padding behind every string, so the encoder writes only +// the bytes the values themselves fill. A class or width the plan never +// produces is a loud error, never a silent zero payload. +func (s *hdf5OutSet) encode(dst []byte) error { + switch s.class { + case 0: + switch { + case s.width == 1 && s.signed: + for i, v := range s.i8s { + dst[i] = byte(v) + } + case s.width == 1: + for i, v := range s.u8s { + dst[i] = v + } + case s.width == 2 && s.signed: + for i, v := range s.i16s { + binary.LittleEndian.PutUint16(dst[i*2:], uint16(v)) + } + case s.width == 2: + for i, v := range s.u16s { + binary.LittleEndian.PutUint16(dst[i*2:], v) + } + case s.width == 4 && s.signed: + for i, v := range s.i32s { + binary.LittleEndian.PutUint32(dst[i*4:], uint32(v)) + } + case s.width == 4: + for i, v := range s.u32s { + binary.LittleEndian.PutUint32(dst[i*4:], v) + } + case s.width == 8: + for i, v := range s.ints { + binary.LittleEndian.PutUint64(dst[i*8:], uint64(v)) + } + default: + return s.payloadRefusal() + } + case 1: + switch s.width { + case 4: + for i, v := range s.f32s { + binary.LittleEndian.PutUint32(dst[i*4:], math.Float32bits(v)) + } + case 8: + for i, v := range s.f64s { + binary.LittleEndian.PutUint64(dst[i*8:], math.Float64bits(v)) + } + default: + return s.payloadRefusal() + } + case 3: + for i, t := range s.texts { + copy(dst[i*s.width:], t) + } + case 8: + // The enumeration's member values, one byte per element. + for i, v := range s.bools { + if v { + dst[i] = 1 + } else { + dst[i] = 0 + } + } + default: + return s.payloadRefusal() + } + return nil +} + +// payloadRefusal names the datatype an encoder cannot serialise. The +// plan produces none of them, so reaching one is a writer defect, and +// it fails loudly rather than emitting a payload of silent zeros. +func (s *hdf5OutSet) payloadRefusal() error { + return base.Errf("dataset %q: the writer cannot serialise datatype class %d of %d bytes per element", s.path, s.class, s.width) +} + +// hdf5Attr is one attribute of the plan: its name and the text its +// value is written from. +type hdf5Attr struct { + name string + text string +} + +// hdf5PlanAttrs validates and sorts the attributes of one object: a +// name must be non-empty and free of NUL bytes, the same constraint +// the reader's rendering can round-trip under. +func hdf5PlanAttrs(name, path string, attrs map[string]string) ([]hdf5Attr, error) { + if len(attrs) > hdf5MaxAttrs { + return nil, base.Errf("%s: %q carries %d attributes, past the %d the writer stores in one object header", name, path, len(attrs), hdf5MaxAttrs) + } + keys := slices.Sorted(maps.Keys(attrs)) + out := make([]hdf5Attr, 0, len(keys)) + for _, k := range keys { + if k == "" { + return nil, base.Errf("%s: %q carries an attribute with an empty name", name, path) + } + if len(k)+1 > 0xffff { + return nil, base.Errf("%s: %q carries the attribute %q whose name is past the %d bytes the attribute message counts", name, path, k, 0xffff) + } + if strings.IndexByte(k, 0) >= 0 || strings.IndexByte(attrs[k], 0) >= 0 { + return nil, base.Errf("%s: %q carries the attribute %q with a NUL byte in its name or value", name, path, k) + } + out = append(out, hdf5Attr{name: k, text: attrs[k]}) + } + return out, nil +} + +// hdf5Writer builds the file image: every address is an offset into +// buf, so structures written later can be referenced by structures +// written earlier through the patch at the end. The chunk staging +// fields are reused across every chunk of one write: the gather and +// shuffle buffers and the index vectors grow to the largest chunk the +// write lays out, and the deflater carries one compressor and one +// output buffer for the whole file. +type hdf5Writer struct { + buf []byte + latest bool + opts HDF5WriteOptions + sbAddr uint64 + + gatherScratch []byte + shuffleScratch []byte + chunkIdx []int + comp *zlib.Writer + compBuf bytes.Buffer + compLevel int +} + +// write lays the file out the way the reference library builds it: a +// placeholder superblock first (its fields name the end of the file +// and the root group, which are known only once everything is +// written), then the root group, whose header is allocated before its +// subtree and filled once the subtree has addresses, and the +// superblock itself last. +func (w *hdf5Writer) write(root *hdf5OutNode) error { + // The image is built into one buffer, so it is allocated once from + // the plan's own size: growing it as the structures are laid out + // would copy the whole file at every step. + w.buf = make([]byte, 0, hdf5ImageEstimate(root, w.opts)) + if w.latest { + w.sbAddr = w.alloc(hdf5Superblock3Size) + } else { + w.sbAddr = w.alloc(hdf5Superblock0Size) + } + if err := w.writeGroup(root); err != nil { + return err + } + w.finishSuperblock(root) + return nil +} + +// writeGroup dispatches to the group writer of the file's layout. +func (w *hdf5Writer) writeGroup(g *hdf5OutNode) error { + if w.latest { + return w.writeNewGroup(g) + } + return w.writeClassicGroup(g) +} + +// Sizes of the fixed parts the writer places first. Every address and +// length is eight bytes, as in the fixtures. +const ( + hdf5Superblock0Size = 24 + 4*8 + 8 + 8 + 4 + 4 + 16 // 96 + hdf5Superblock3Size = 12 + 4*8 + 4 // 48 +) + +// finishSuperblock fills the placeholder: the classic superblock +// names the root group through a symbol table entry whose cache holds +// the group's B-tree and local heap, the latest one names its object +// header directly and checksums the whole block with lookup3. +func (w *hdf5Writer) finishSuperblock(root *hdf5OutNode) { + copy(w.buf[w.sbAddr:], hdf5Magic) + if w.latest { + w.buf[w.sbAddr+8] = 3 + w.buf[w.sbAddr+9] = 8 + w.buf[w.sbAddr+10] = 8 + w.buf[w.sbAddr+11] = 0 // file consistency flags + w.set64(w.sbAddr+12, 0) + w.set64(w.sbAddr+20, math.MaxUint64) + w.set64(w.sbAddr+28, uint64(len(w.buf))) + w.set64(w.sbAddr+36, root.addr) + w.set32(w.sbAddr+44, hdf5Lookup3(w.buf[w.sbAddr:w.sbAddr+44])) + return + } + w.buf[w.sbAddr+8] = 0 // superblock version + w.buf[w.sbAddr+9] = 0 // free space storage version + w.buf[w.sbAddr+10] = 0 // root group symbol table entry version + w.buf[w.sbAddr+11] = 0 // reserved + w.buf[w.sbAddr+12] = 0 // shared header message format version + w.buf[w.sbAddr+13] = 8 // size of offsets + w.buf[w.sbAddr+14] = 8 // size of lengths + w.buf[w.sbAddr+15] = 0 // reserved + binary.LittleEndian.PutUint16(w.buf[w.sbAddr+16:], hdf5GroupLeafK) + binary.LittleEndian.PutUint16(w.buf[w.sbAddr+18:], hdf5GroupInnerK) + // The file consistency flags at +20 stay zero. + w.set64(w.sbAddr+24, 0) // base address + w.set64(w.sbAddr+32, math.MaxUint64) // free space information + w.set64(w.sbAddr+40, uint64(len(w.buf))) // end of file + w.set64(w.sbAddr+48, math.MaxUint64) // driver information + w.set64(w.sbAddr+56, 0) // root entry: link name offset + w.set64(w.sbAddr+64, root.addr) // root entry: object header address + w.set32(w.sbAddr+72, 1) // root entry: symbol table cache + w.set64(w.sbAddr+80, root.btree) // cache: B-tree address + w.set64(w.sbAddr+88, root.heap) // cache: local heap address +} + +func (w *hdf5Writer) alloc(n int) uint64 { + addr := uint64(len(w.buf)) + w.buf = append(w.buf, make([]byte, n)...) + return addr +} + +// reserve extends the image by n bytes without writing them and +// returns the address they start at; the caller fills the whole span +// in the same breath, so the payload passes through the writer once. +// The capacity beyond len(buf) always holds the zeros the buffer's +// allocations left there, which every fixed structure and string slot +// is padded from, and the encode that follows a reservation writes +// every byte the payload itself does not. +func (w *hdf5Writer) reserve(n int) uint64 { + addr := uint64(len(w.buf)) + if cap(w.buf)-len(w.buf) < n { + w.buf = append(w.buf, make([]byte, n)...) + return addr + } + w.buf = w.buf[:len(w.buf)+n] + return addr +} + +func (w *hdf5Writer) bytes(b []byte) uint64 { + addr := uint64(len(w.buf)) + w.buf = append(w.buf, b...) + return addr +} + +// pad8 aligns the image to the eight-byte boundary the format inserts +// between the structures of the classic layout. +func (w *hdf5Writer) pad8() { + if r := len(w.buf) % 8; r != 0 { + w.buf = append(w.buf, make([]byte, 8-r)...) + } +} + +func (w *hdf5Writer) set32(at uint64, v uint32) { + binary.LittleEndian.PutUint32(w.buf[at:], v) +} + +func (w *hdf5Writer) set64(at uint64, v uint64) { + binary.LittleEndian.PutUint64(w.buf[at:], v) +} diff --git a/io/hdf5write_estimate_test.go b/io/hdf5write_estimate_test.go new file mode 100644 index 0000000..b80eb9a --- /dev/null +++ b/io/hdf5write_estimate_test.go @@ -0,0 +1,23 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import "testing" + +// TestHDF5SetEstimateCapsTinyChunkTargets pins the node estimate's +// cap: the estimate is a starting buffer the write grows past by +// append, so a pathological chunk target of a couple of bytes must +// charge the dataset a bounded guess rather than a B-tree node, and +// its kilobyte of image, to every single data byte. +func TestHDF5SetEstimateCapsTinyChunkTargets(t *testing.T) { + s := &hdf5OutSet{path: "d", shape: []int{1 << 20}, class: 1, width: 8, nbytes: 8 << 20} + honest := hdf5SetEstimate(s, HDF5WriteOptions{ChunkBytes: 4096, Shuffle: true}) + if honest <= int64(s.nbytes) || honest > 100<<20 { + t.Fatalf("the honest estimate is %d bytes for a %d-byte payload, want a bounded cover", honest, s.nbytes) + } + tiny := hdf5SetEstimate(s, HDF5WriteOptions{ChunkBytes: 2, Shuffle: true}) + if tiny > 50<<20 { + t.Fatalf("a two-byte chunk target estimated %d bytes for a %d-byte payload", tiny, s.nbytes) + } +} diff --git a/io/hdf5write_group.go b/io/hdf5write_group.go new file mode 100644 index 0000000..e695df1 --- /dev/null +++ b/io/hdf5write_group.go @@ -0,0 +1,717 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +// Group and chunked-storage writing for the HDF5 writer. Both group +// layouts the reader accepts are written: the classic symbol table (a +// local heap of names, symbol table nodes and a version 1 group +// B-tree) and the latest-style compact group (link messages in the +// object header). Chunked storage hangs its chunks off a version 1 +// chunk B-tree, whose conventions the reference files pin: every node +// is allocated at the size the format's fanout implies, an interior +// node's key is the first key of the child's subtree, and a node's +// closing key is a sentinel ordered past its last chunk. + +import ( + "bytes" + "compress/zlib" + "encoding/binary" + "math" + "slices" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// hdf5GroupKid is one child pointer of a group B-tree node. The key of +// kid i is the heap offset of the last name child i-1 holds, so a +// node's entries are headed by the zero key (the heap's null name, +// below every real name) and each entry's key closes the range of the +// child before it, which is the convention the fixtures store. +type hdf5GroupKid struct { + key uint64 + child uint64 +} + +// hdf5GroupChild is one child entry of a group: its link name, its +// object header address and, until the children are written, the +// subgroup or dataset whose address it carries. +type hdf5GroupChild struct { + name string + addr uint64 + group *hdf5OutNode + set *hdf5OutSet +} + +// hdf5Kids merges a group's subgroups and datasets into one list in +// name order, the order the classic symbol table requires: the local +// heap assigns its offsets in list order and the reference library +// binary-searches both the B-tree keys and the symbol node records on +// those offsets, so a groups-first order would hide every child whose +// name interleaves with a dataset's. The latest layout's link list +// keeps the same order. +func hdf5Kids(g *hdf5OutNode) []hdf5GroupChild { + kids := make([]hdf5GroupChild, 0, len(g.groups)+len(g.sets)) + for _, sub := range g.groups { + kids = append(kids, hdf5GroupChild{name: sub.name, group: sub}) + } + for _, s := range g.sets { + kids = append(kids, hdf5GroupChild{name: baseName(s.path), set: s}) + } + slices.SortFunc(kids, func(a, b hdf5GroupChild) int { return strings.Compare(a.name, b.name) }) + return kids +} + +// kidAddr is a child's written object header address. +func kidAddr(k hdf5GroupChild) uint64 { + if k.group != nil { + return k.group.addr + } + return k.set.addr +} + +// writeClassicGroup writes a group in the classic layout, in the +// construction order the reference library uses: the object header is +// allocated first (over a placeholder, since the symbol table message +// can only name the B-tree and heap once they exist), then the local +// heap, then the children, and the symbol table nodes and B-tree last. +// A child group's symbol entry carries the B-tree and heap addresses +// in its cache, as the reference writes; a dataset's carries none. +func (w *hdf5Writer) writeClassicGroup(g *hdf5OutNode) error { + kids := hdf5Kids(g) + // The header placeholder: sized and shaped exactly like the final + // image, with the B-tree and heap addresses still zero. + msgs, err := w.classicGroupMsgs(g, 0, 0) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + image, err := hdf5HeaderV1(msgs) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + hdrAddr := w.bytes(image) + // The heap assigns every name its offset; the children are in + // sorted order, so offset order and name order agree and the + // B-tree's byte-offset keys sort as the names do. The heap carries + // slack behind the names, written as the free block the reference + // library leaves there: a next offset of one is the sentinel that + // ends the free list. + heapData := make([]byte, 8) // offset 0 is the null name no entry points at + offsets := make([]uint64, len(kids)) + for i, k := range kids { + offsets[i] = uint64(len(heapData)) + heapData = append(heapData, k.name...) + heapData = hdf5AppendAlign(append(heapData, 0)) + } + size := max(len(heapData)+16, 88) + free := len(heapData) + heapData = append(heapData, make([]byte, 16)...) // the free block: next, size + binary.LittleEndian.PutUint64(heapData[free:], 1) + binary.LittleEndian.PutUint64(heapData[free+8:], uint64(size-free)) + heapData = append(heapData, make([]byte, size-len(heapData))...) + dataAddr := w.bytes(heapData) + w.pad8() + heap := w.writeLocalHeap(dataAddr, size, free) + // The children: datasets, then the subgroups, each of which lays + // out its own subtree in the children's name order. + for _, s := range g.sets { + addr, err := w.writeDataset(s) + if err != nil { + // writeDataset names the dataset itself; a second wrap + // here would prefix it twice. + return err + } + s.addr = addr + } + for _, sub := range g.groups { + if err := w.writeGroup(sub); err != nil { + return err + } + } + for i := range kids { + kids[i].addr = kidAddr(kids[i]) + } + // Symbol table nodes, eight entries each: the node occupies the + // size the format's leaf K of four implies, however few entries are + // used, and an empty group still holds one. + leaves := make([]hdf5GroupKid, 0, max((len(kids)+2*hdf5GroupLeafK-1)/(2*hdf5GroupLeafK), 1)) + for start := 0; start < len(kids); start += 2 * hdf5GroupLeafK { + leaves = append(leaves, w.writeSymbolNode(kids[start:min(start+2*hdf5GroupLeafK, len(kids))], offsets[start:])) + } + if len(leaves) == 0 { + leaves = append(leaves, w.writeSymbolNode(nil, nil)) + } + btree := w.writeGroupTree(leaves) + // The header over the placeholder, now that the B-tree and heap + // have their final addresses. + msgs, err = w.classicGroupMsgs(g, btree, heap) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + image, err = hdf5HeaderV1(msgs) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + w.headerAt(hdrAddr, image) + g.addr, g.btree, g.heap = hdrAddr, btree, heap + return nil +} + +// classicGroupMsgs builds the message list of a classic group's object +// header: the symbol table message naming the B-tree and heap, then +// the attributes in name order. +func (w *hdf5Writer) classicGroupMsgs(g *hdf5OutNode, btree, heap uint64) ([]hdf5OutMsg, error) { + msgs := []hdf5OutMsg{{typ: hdf5MsgSymbolTable, data: w.appendAddrs(nil, btree, heap)}} + for _, a := range g.attrs { + m, err := hdf5AttrMessage(a) + if err != nil { + return nil, err + } + msgs = append(msgs, m) + } + return msgs, nil +} + +// writeSymbolNode writes one symbol table node of the given entries, +// allocated at eight slots. A child group's cache holds its B-tree and +// heap addresses; a dataset's cache stays empty. +func (w *hdf5Writer) writeSymbolNode(batch []hdf5GroupChild, offsets []uint64) hdf5GroupKid { + node := make([]byte, 8+2*hdf5GroupLeafK*(8+8+24)) + copy(node, hdf5SymbolNode) + node[4] = 1 + binary.LittleEndian.PutUint16(node[6:], uint16(len(batch))) + p := 8 + for i, k := range batch { + binary.LittleEndian.PutUint64(node[p:], offsets[i]) + binary.LittleEndian.PutUint64(node[p+8:], k.addr) + if k.group != nil { + binary.LittleEndian.PutUint32(node[p+16:], 1) + binary.LittleEndian.PutUint64(node[p+24:], k.group.btree) + binary.LittleEndian.PutUint64(node[p+32:], k.group.heap) + } + p += 8 + 8 + 24 + } + leaf := hdf5GroupKid{child: w.bytes(node)} + if len(batch) > 0 { + leaf.key = offsets[len(batch)-1] + } + return leaf +} + +// baseName returns the last segment of an absolute path. +func baseName(path string) string { + for i := len(path) - 1; i >= 0; i-- { + if path[i] == '/' { + return path[i+1:] + } + } + return path +} + +// writeNewGroup writes a group in the latest layout, the construction +// order the reference library uses: the object header is allocated +// first over a placeholder, then the children, and the header is +// written over the placeholder once every link's address is known. +// The link info and group info messages the reference writes head the +// message list; a group without links carries no link info message: +// beside no links the reader of this package takes that message for +// dense storage, which it refuses. +func (w *hdf5Writer) writeNewGroup(g *hdf5OutNode) error { + msgs, err := w.newGroupMsgs(g) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + image, err := hdf5HeaderV2(msgs) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + hdrAddr := w.bytes(image) + for _, s := range g.sets { + addr, err := w.writeDataset(s) + if err != nil { + // writeDataset names the dataset itself; a second wrap + // here would prefix it twice. + return err + } + s.addr = addr + } + for _, sub := range g.groups { + if err := w.writeGroup(sub); err != nil { + return err + } + } + msgs, err = w.newGroupMsgs(g) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + image, err = hdf5HeaderV2(msgs) + if err != nil { + return base.Errf("SaveHDF5: group %q: %w", g.path, err) + } + w.headerAt(hdrAddr, image) + g.addr = hdrAddr + return nil +} + +// newGroupMsgs builds the message list of a latest-style group's +// object header. Addresses unknown at placeholder time are zero; the +// message sizes are the same either way. +func (w *hdf5Writer) newGroupMsgs(g *hdf5OutNode) ([]hdf5OutMsg, error) { + var msgs []hdf5OutMsg + if len(g.groups)+len(g.sets) > 0 { + undef := bytes.Repeat([]byte{0xff}, 8) + info := append([]byte{0, 0}, undef...) + info = append(info, undef...) + msgs = append(msgs, + hdf5OutMsg{typ: hdf5MsgLinkInfo, data: info}, + hdf5OutMsg{typ: hdf5MsgGroupInfo, data: []byte{0, 0}}, + ) + } + addLink := func(name string, addr uint64) { + body := []byte{1, 0} // version 1, no creation order, one-byte length + if len(name) >= 256 { + body[1] = 0x01 // the name length widens to two bytes + body = binary.LittleEndian.AppendUint16(body, uint16(len(name))) + } else { + body = append(body, byte(len(name))) + } + body = append(body, name...) + body = binary.LittleEndian.AppendUint64(body, addr) + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgLink, data: body}) + } + for _, k := range hdf5Kids(g) { + addLink(k.name, kidAddr(k)) + } + for _, a := range g.attrs { + m, err := hdf5AttrMessage(a) + if err != nil { + // Unwrapped: the caller carries the group context and + // prefixes the entry point, so wrapping here would double + // both. + return nil, err + } + msgs = append(msgs, m) + } + return msgs, nil +} + +// writeLocalHeap writes a local heap header in front of the data +// segment written earlier: the data segment's size, the free list head +// and the segment address. +func (w *hdf5Writer) writeLocalHeap(dataAddr uint64, size, free int) uint64 { + head := append([]byte{}, hdf5LocalHeap...) + head = append(head, 0, 0, 0, 0) // version and reserved + head = binary.LittleEndian.AppendUint64(head, uint64(size)) + head = binary.LittleEndian.AppendUint64(head, uint64(free)) + head = binary.LittleEndian.AppendUint64(head, dataAddr) + return w.bytes(head) +} + +// writeGroupTree packs the symbol table nodes into version 1 group +// B-tree nodes of thirty-two children and returns the root's address. +func (w *hdf5Writer) writeGroupTree(leaves []hdf5GroupKid) uint64 { + kids := leaves + level := byte(0) + for len(kids) > 2*hdf5GroupInnerK { + var next []hdf5GroupKid + for start := 0; start < len(kids); start += 2 * hdf5GroupInnerK { + batch := kids[start:min(start+2*hdf5GroupInnerK, len(kids))] + addr := w.writeGroupNode(level, batch) + next = append(next, hdf5GroupKid{key: batch[len(batch)-1].key, child: addr}) + } + kids = next + level++ + } + return w.writeGroupNode(level, kids) +} + +// writeGroupNode writes one group B-tree node, allocated at the size +// the format's internal K of sixteen implies. The sibling addresses +// are the undefined value: a zero would be a defined address, and a +// reader would follow it as a sibling node. +func (w *hdf5Writer) writeGroupNode(level byte, kids []hdf5GroupKid) uint64 { + node := make([]byte, 8+2*8+2*hdf5GroupInnerK*(8+8)+8) + copy(node, hdf5Tree) + node[4] = 0 // a group B-tree + node[5] = level + binary.LittleEndian.PutUint16(node[6:], uint16(len(kids))) + binary.LittleEndian.PutUint64(node[8:], math.MaxUint64) // left sibling + binary.LittleEndian.PutUint64(node[16:], math.MaxUint64) // right sibling + p := 24 + binary.LittleEndian.PutUint64(node[p:], 0) // key[0]: the null name + p += 8 + for _, k := range kids { + binary.LittleEndian.PutUint64(node[p:], k.child) + p += 8 + binary.LittleEndian.PutUint64(node[p:], k.key) // the key closing this child + p += 8 + } + return w.bytes(node) +} + +// appendAddrs appends the given addresses as the file's eight-byte +// address fields. +func (w *hdf5Writer) appendAddrs(b []byte, addrs ...uint64) []byte { + for _, a := range addrs { + b = binary.LittleEndian.AppendUint64(b, a) + } + return b +} + +// hdf5ChunkKey is a version 1 chunk B-tree key: the chunk's stored +// size, its coordinates, one per dataset dimension plus the element +// size dimension the format keys last, and the address the chunk's +// bytes occupy. +type hdf5ChunkKey struct { + size uint32 + coords []uint64 + addr uint64 +} + +// hdf5ChunkKid is one child of a chunk B-tree node with the first and +// last keys of its subtree, which become the node's own keys. +type hdf5ChunkKid struct { + first, last hdf5ChunkKey + child uint64 +} + +// hdf5ChunkNodeSize is the size a chunk B-tree node occupies: the +// header, twice the storage K of thirty-two keys with their child +// pointers, and the closing key. +func hdf5ChunkNodeSize(rank int) int { + keySize := 8 + 8*(rank+1) + return 8 + 2*8 + 2*hdf5IStoreK*(keySize+8) + keySize +} + +// hdf5ChunkSentinel is the key that closes a node: the last chunk's +// coordinates with the element size dimension set to the element size, +// which orders it past every chunk of the node, as the reference +// files store it. Reference tooling reads either shape of the +// sentinel: recent releases write the chunk shape itself with the +// element size appended, which orders identically; this writer keeps +// the fixture convention. +func hdf5ChunkSentinel(k hdf5ChunkKey, width int) hdf5ChunkKey { + coords := append([]uint64{}, k.coords...) + coords[len(coords)-1] = uint64(width) + return hdf5ChunkKey{size: 0, coords: coords} +} + +// writeChunks chunks one dataset, applies the pipeline to every chunk +// and lays the chunks out through a chunk B-tree, returning the +// B-tree's address and the chunk shape. A chunk is stored complete: +// the edge chunks of a dataset whose extent is not a whole number of +// chunks are padded with the zero fill value, which is what chunked +// storage holds and what the reader requires. The staging buffers, the +// index vectors and the deflater are writer state reused across every +// chunk and every dataset of the write; each chunk fully overwrites +// what it stages, so no bytes travel between chunks. +func (w *hdf5Writer) writeChunks(s *hdf5OutSet) (uint64, []int, error) { + chunk := hdf5ChunkShape(s.shape, s.width, w.chunkTarget()) + grid := make([]int, len(s.shape)) + total := 1 + for i := range grid { + grid[i] = (s.shape[i] + chunk[i] - 1) / chunk[i] + total *= grid[i] + } + if total > hdf5MaxChunks { + return 0, nil, base.Errf("chunking to a %d byte target needs %d chunks, past the %d the writer lays out; raise the chunk target", w.chunkTarget(), total, hdf5MaxChunks) + } + if _, err := hdf5ByteExtent(chunk, s.width, hdf5MaxDatasetBytes); err != nil { + return 0, nil, base.Errf("the chunk shape %v: %w", chunk, err) + } + level := w.opts.Gzip + if level == -1 { + level = 6 + } + keys := make([]hdf5ChunkKey, 0, total) + coords := make([]int, len(grid)) + // Every key aliases one range of a single coordinate slab, which + // the chunk B-tree reads before the write ends. + coordSlab := make([]uint64, total*(len(chunk)+1)) + rank := len(s.shape) + chunkElems := 1 + for _, c := range chunk { + chunkElems *= c + } + chunkBytes := chunkElems * s.width + if cap(w.gatherScratch) < chunkBytes { + w.gatherScratch = make([]byte, chunkBytes) + } + if w.opts.Shuffle && s.width > 1 && cap(w.shuffleScratch) < chunkBytes { + w.shuffleScratch = make([]byte, chunkBytes) + } + if cap(w.chunkIdx) < 4*rank { + w.chunkIdx = make([]int, 4*rank) + } + origin := w.chunkIdx[:rank] + count := w.chunkIdx[rank : 2*rank] + srcStride := w.chunkIdx[2*rank : 3*rank] + dstStride := w.chunkIdx[3*rank : 4*rank] + cs, ds := 1, 1 + for i := rank - 1; i >= 0; i-- { + srcStride[i], dstStride[i] = cs, ds + cs *= s.shape[i] + ds *= chunk[i] + } + for ci := range total { + for d := range rank { + origin[d] = coords[d] * chunk[d] + } + data := w.gatherScratch[:chunkBytes] + if err := w.fillChunk(data, s, chunk, origin, count, srcStride, dstStride); err != nil { + // The dataset caller wraps with the dataset's own context; + // the bare encode error keeps it single. + return 0, nil, err + } + if w.opts.Shuffle && s.width > 1 { + hdf5ShuffleInto(w.shuffleScratch[:chunkBytes], data, s.width) + data = w.shuffleScratch[:chunkBytes] + } + if w.opts.Gzip != 0 { + compressed, err := w.deflateChunk(data, level) + if err != nil { + // The dataset caller wraps with the dataset's own + // context; the bare filter error keeps it single. + return 0, nil, err + } + data = compressed + } + addr := w.bytes(data) + w.pad8() + // The key's trailing coordinate is the element size dimension + // the format keys last, and the reference library carries 0 in + // it for every real chunk: only the closing sentinel holds the + // element size, which is what orders it past the node's chunks. + // A real key carrying the size would tie the sentinel and hide + // every chunk from a binary search. + key := hdf5ChunkKey{size: uint32(len(data)), coords: coordSlab[ci*(len(chunk)+1) : (ci+1)*(len(chunk)+1)], addr: addr} + for i := range chunk { + key.coords[i] = uint64(coords[i] * chunk[i]) + } + keys = append(keys, key) + // The odometer walks the chunk grid in row-major order, which + // is the order the B-tree's keys must be sorted in. + for d := len(coords) - 1; d >= 0; d-- { + coords[d]++ + if coords[d] < grid[d] { + break + } + coords[d] = 0 + } + } + return w.writeChunkTree(keys, len(s.shape), s.width), chunk, nil +} + +// fillChunk stages one chunk of the dataset into dst, which holds the +// chunk's full element extent: the cells inside the shape take the +// values encoded little-endian at their stored width, exactly the +// bytes encode produces from the same values, and the cells past any +// edge stay the zeros clear left. Every byte of dst is written on +// every call, so the reused gather buffer carries nothing between +// chunks. origin is the chunk's first cell in dataset coordinates, +// count walks the overlap of the chunk and the shape, and the strides +// are row-major element strides of the dataset and of the chunk. A +// class or width the plan never produces is the same loud refusal +// encode gives. +func (w *hdf5Writer) fillChunk(dst []byte, s *hdf5OutSet, chunk, origin, count, srcStride, dstStride []int) error { + clear(dst) + shape := s.shape + rank := len(shape) + width := s.width + for d := range rank { + count[d] = origin[d] + } + for { + src, loc := 0, 0 + for d := range rank { + src += count[d] * srcStride[d] + loc += (count[d] - origin[d]) * dstStride[d] + } + p := loc * width + switch { + case s.class == 0 && width == 1 && s.signed: + dst[p] = byte(s.i8s[src]) + case s.class == 0 && width == 1: + dst[p] = s.u8s[src] + case s.class == 0 && width == 2 && s.signed: + binary.LittleEndian.PutUint16(dst[p:], uint16(s.i16s[src])) + case s.class == 0 && width == 2: + binary.LittleEndian.PutUint16(dst[p:], s.u16s[src]) + case s.class == 0 && width == 4 && s.signed: + binary.LittleEndian.PutUint32(dst[p:], uint32(s.i32s[src])) + case s.class == 0 && width == 4: + binary.LittleEndian.PutUint32(dst[p:], s.u32s[src]) + case s.class == 0 && width == 8: + binary.LittleEndian.PutUint64(dst[p:], uint64(s.ints[src])) + case s.class == 8: + if s.bools[src] { + dst[p] = 1 + } else { + dst[p] = 0 + } + case s.class == 1 && width == 4: + binary.LittleEndian.PutUint32(dst[p:], math.Float32bits(s.f32s[src])) + case s.class == 1 && width == 8: + binary.LittleEndian.PutUint64(dst[p:], math.Float64bits(s.f64s[src])) + default: + return s.payloadRefusal() + } + d := rank - 1 + for d >= 0 { + count[d]++ + if count[d] < min(origin[d]+chunk[d], shape[d]) { + break + } + count[d] = origin[d] + d-- + } + if d < 0 { + return nil + } + } +} + +// chunkTarget is the configured chunk byte target, or the default. +func (w *hdf5Writer) chunkTarget() int { return hdf5ChunkTargetOf(w.opts) } + +// hdf5ChunkTargetOf resolves the chunk target the options ask for. +func hdf5ChunkTargetOf(opts HDF5WriteOptions) int { + if opts.ChunkBytes == 0 { + return hdf5ChunkTarget + } + return opts.ChunkBytes +} + +// hdf5ChunkShape picks a chunk shape that aims at the target bytes: a +// dataset that fits the target stays whole, otherwise the last axis is +// split first, which keeps the access the files of this package see +// (a row of a table, a trace of a signal) inside one chunk. +func hdf5ChunkShape(shape []int, width, target int) []int { + n := 1 + for _, d := range shape { + n *= d + } + if n*width <= target { + return append([]int{}, shape...) + } + row := width + for _, d := range shape[:len(shape)-1] { + row *= d + } + last := max(target/row, 1) + out := append([]int{}, shape[:len(shape)-1]...) + return append(out, min(last, shape[len(shape)-1])) +} + +// writeChunkTree packs the chunks into version 1 chunk B-tree nodes of +// sixty-four children and returns the root's address. An interior +// node's key for child i is the first key of child i's subtree, and a +// node's level counts its distance from the chunks. +func (w *hdf5Writer) writeChunkTree(keys []hdf5ChunkKey, rank, width int) uint64 { + kids := make([]hdf5ChunkKid, len(keys)) + for i, k := range keys { + kids[i] = hdf5ChunkKid{first: k, last: k, child: k.addr} + } + level := byte(0) + for len(kids) > 2*hdf5IStoreK { + var next []hdf5ChunkKid + for start := 0; start < len(kids); start += 2 * hdf5IStoreK { + batch := kids[start:min(start+2*hdf5IStoreK, len(kids))] + addr := w.writeChunkNode(level, batch, rank, width) + next = append(next, hdf5ChunkKid{first: batch[0].first, last: batch[len(batch)-1].last, child: addr}) + } + kids = next + level++ + } + return w.writeChunkNode(level, kids, rank, width) +} + +// writeChunkNode writes one chunk B-tree node: the first key of every +// child, a child pointer each, and the sentinel key that closes the +// node. The sibling addresses are the undefined value, as in the group +// B-tree. +func (w *hdf5Writer) writeChunkNode(level byte, kids []hdf5ChunkKid, rank, width int) uint64 { + node := make([]byte, hdf5ChunkNodeSize(rank)) + copy(node, hdf5Tree) + node[4] = 1 // a chunk B-tree + node[5] = level + binary.LittleEndian.PutUint16(node[6:], uint16(len(kids))) + binary.LittleEndian.PutUint64(node[8:], math.MaxUint64) // left sibling + binary.LittleEndian.PutUint64(node[16:], math.MaxUint64) // right sibling + p := 24 + writeKey := func(k hdf5ChunkKey) { + binary.LittleEndian.PutUint32(node[p:], k.size) + binary.LittleEndian.PutUint32(node[p+4:], 0) // the filter mask: every filter ran + for i, c := range k.coords { + binary.LittleEndian.PutUint64(node[p+8+8*i:], c) + } + p += 8 + 8*len(k.coords) + } + for _, k := range kids { + writeKey(k.first) + binary.LittleEndian.PutUint64(node[p:], k.child) + p += 8 + } + writeKey(hdf5ChunkSentinel(kids[len(kids)-1].last, width)) + return w.bytes(node) +} + +// filterEntries lists the configured pipeline in application order for +// the filter pipeline message: shuffle before deflate, the order that +// lays each element's high-order bytes together for the compressor. +// The reference stores the element width as the shuffle's one client +// value and the level as deflate's. +func (w *hdf5Writer) filterEntries(width int) []hdf5OutFilter { + var out []hdf5OutFilter + if w.opts.Shuffle { + out = append(out, hdf5OutFilter{id: hdf5FilterShuffle, name: "shuffle", values: []uint32{uint32(width)}}) + } + if w.opts.Gzip != 0 { + level := w.opts.Gzip + if level == -1 { + level = 6 + } + out = append(out, hdf5OutFilter{id: hdf5FilterDeflate, name: "deflate", values: []uint32{uint32(level)}}) + } + return out +} + +// hdf5ShuffleInto transposes the element bytes of data into out, which +// must hold the same length: the first block takes every element's +// first byte, the next every second byte. Every byte of out is written. +func hdf5ShuffleInto(out, data []byte, width int) { + n := len(data) / width + for i := range width { + for j := range n { + out[i*n+j] = data[j*width+i] + } + } +} + +// deflateChunk compresses one staged chunk into the zlib stream the +// deflate filter stores: a zlib header in front of the raw deflate +// stream, which is what the reader's inflate expects. The compressor +// and its output buffer are reused across the chunks of a write; +// Reset leaves the compressor in the state a fresh writer holds, so +// the stream carries the same bytes a per-chunk writer produced. The +// returned bytes are the buffer's, valid until the next call. +func (w *hdf5Writer) deflateChunk(data []byte, level int) ([]byte, error) { + if w.comp == nil || w.compLevel != level { + zw, err := zlib.NewWriterLevel(&w.compBuf, level) + if err != nil { + return nil, base.Errf("the deflate level %d is refused: %w", level, err) + } + w.comp, w.compLevel = zw, level + } else { + w.comp.Reset(&w.compBuf) + } + w.compBuf.Reset() + if _, err := w.comp.Write(data); err != nil { + return nil, base.Errf("deflate: %w", err) + } + if err := w.comp.Close(); err != nil { + return nil, base.Errf("deflate: %w", err) + } + return w.compBuf.Bytes(), nil +} diff --git a/io/hdf5write_object.go b/io/hdf5write_object.go new file mode 100644 index 0000000..af4c484 --- /dev/null +++ b/io/hdf5write_object.go @@ -0,0 +1,449 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +// Object header and message writing for the HDF5 writer: the messages +// here are shaped exactly as the reader's decoders in hdf5.go parse +// them, which the fixtures under testdata/h5 pin byte for byte where +// it matters. + +import ( + "encoding/binary" + "math" + "strconv" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// hdf5OutMsg is one message of an object header: a type, the payload +// the format defines for it. +type hdf5OutMsg struct { + typ uint16 + data []byte +} + +// hdf5HeaderV1 builds a version 1 object header image, the one the +// classic layout uses: the fixed head names the message count and the +// size of the message region, and every message is preceded by an +// eight-byte header. The message's stored size includes the padding +// that keeps the next message on an eight-byte boundary of the header, +// which is how the reference library writes every message class. +func hdf5HeaderV1(msgs []hdf5OutMsg) ([]byte, error) { + if len(msgs) > 0xffff { + return nil, base.Errf("SaveHDF5: an object header would carry %d messages, past the %d the format counts", len(msgs), 0xffff) + } + body := []byte{} + for _, m := range msgs { + data := hdf5PadField(m.data) + if len(data) > 0xffff { + return nil, base.Errf("SaveHDF5: a message of %d bytes is past the %d a version 1 header stores", len(data), 0xffff) + } + body = binary.LittleEndian.AppendUint16(body, m.typ) + body = binary.LittleEndian.AppendUint16(body, uint16(len(data))) + body = append(body, 0, 0, 0, 0) // message flags and reserved + body = append(body, data...) + } + head := []byte{1, 0} + head = binary.LittleEndian.AppendUint16(head, uint16(len(msgs))) + head = binary.LittleEndian.AppendUint32(head, 1) // reference count + head = binary.LittleEndian.AppendUint32(head, uint32(len(body))) + head = binary.LittleEndian.AppendUint32(head, 0) // padding to eight + return append(head, body...), nil +} + +// hdf5HeaderV2 builds a version 2 object header image, the one the +// latest layout uses: the messages are packed without alignment and +// the whole header, signature to last message, closes with a lookup3 +// checksum. The size field is one, two or four bytes, whichever holds +// it. +func hdf5HeaderV2(msgs []hdf5OutMsg) ([]byte, error) { + body := []byte{} + for _, m := range msgs { + if len(m.data) > 0xffff { + return nil, base.Errf("SaveHDF5: a message of %d bytes is past the %d a version 2 header stores", len(m.data), 0xffff) + } + body = append(body, byte(m.typ)) + body = binary.LittleEndian.AppendUint16(body, uint16(len(m.data))) + body = append(body, 0) // message flags + body = append(body, m.data...) + } + // The size of the size field is itself a mask: the stored value is + // the base-2 logarithm of the width, 0 for one byte, 1 for two and + // 2 for four, which the reader decodes as 1 << stored. + if len(body) > math.MaxUint32 { + return nil, base.Errf("SaveHDF5: an object header of %d bytes is past the %d the version 2 size field holds", len(body), uint64(math.MaxUint32)) + } + var logWidth byte + switch { + case len(body) > 0xffff: + logWidth = 2 + case len(body) > 0xff: + logWidth = 1 + } + width := byte(1) << logWidth + out := append([]byte{}, hdf5ObjHdr2...) + out = append(out, 2, logWidth) + for i := range width { + out = append(out, byte(len(body)>>(8*i))) + } + out = append(out, body...) + return binary.LittleEndian.AppendUint32(out, hdf5Lookup3(out)), nil +} + +// headerAt writes a built header image over a placeholder allocated +// earlier, keeping every address that already points here valid. +func (w *hdf5Writer) headerAt(addr uint64, image []byte) { + copy(w.buf[addr:], image) +} + +// hdf5AppendAlign appends the zero bytes that align b to its next +// eight-byte boundary. +func hdf5AppendAlign(b []byte) []byte { + if r := len(b) % 8; r != 0 { + return append(b, make([]byte, 8-r)...) + } + return b +} + +// hdf5PadField pads a name or a value field of a message to its next +// eight-byte boundary; a field already aligned stays as it is. +func hdf5PadField(b []byte) []byte { return hdf5AppendAlign(b) } + +// Datatype messages, version 1. The class bit field's first byte +// carries the byte order (cleared: little-endian) and, for fixed +// point, the signed bit; the floating-point classes carry the IEEE 754 +// layout the fixtures store, down to the exponent bias. +var ( + hdf5Float32Type = []byte{ + 0x11, // version 1, class 1 (floating-point) + 0x20, 0x1f, 0x00, // little-endian, the sign at bit 31 + 0x04, 0x00, 0x00, 0x00, // four-byte elements + 0x00, 0x00, // bit offset 0 + 0x20, 0x00, // 32 bits of precision + 0x17, 0x08, 0x00, 0x17, // exponent at 23 of 8, mantissa at 0 of 23 + 0x7f, 0x00, 0x00, 0x00, // exponent bias 127 + } + hdf5Float64Type = []byte{ + 0x11, // version 1, class 1 (floating-point) + 0x20, 0x3f, 0x00, // little-endian, the sign at bit 63 + 0x08, 0x00, 0x00, 0x00, // eight-byte elements + 0x00, 0x00, // bit offset 0 + 0x40, 0x00, // 64 bits of precision + 0x34, 0x0b, 0x00, 0x34, // exponent at 52 of 11, mantissa at 0 of 52 + 0xff, 0x03, 0x00, 0x00, // exponent bias 1023 + } +) + +// hdf5IntType writes a fixed-point datatype message of the given +// element width and signedness. The class bit field's byte order bit +// stays clear (little-endian) and bit 0x08 carries two's-complement +// signedness, the bit decodeType keys its landing on. The message is +// twelve bytes: the eight-byte header plus the bit offset and bit +// precision the format's fixed-point property table defines. +func hdf5IntType(width int, signed bool) []byte { + flags := byte(0) // little-endian, unsigned + if signed { + flags = 0x08 // bit 3: two's complement + } + b := []byte{0x10, flags, 0, 0} // version 1, class 0 (fixed-point) + b = binary.LittleEndian.AppendUint32(b, uint32(width)) + b = binary.LittleEndian.AppendUint16(b, 0) + b = binary.LittleEndian.AppendUint16(b, uint16(8*width)) + return b +} + +// hdf5BoolType writes the boolean enumeration datatype message, in the +// exact shape the reader's hdf5EnumBool admits: class 8 with a member +// count of two in the class bit field's low sixteen bits and the +// reserved byte zero, a base type that is a complete one-byte unsigned +// little-endian fixed-point message, the member names each NUL +// terminated and padded from its own field start to a multiple of eight +// bytes (the message version 1 convention) and the packed member values +// 0 and 1 behind the names. The member names carry no semantics; the +// values are what the landing reads, and the payload stores them, one +// byte per element. +func hdf5BoolType() []byte { + b := []byte{0x18, 0x02, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00} // version 1, class 8, two members, one-byte values + b = append(b, hdf5IntType(1, false)...) // the base type + for _, name := range []string{"FALSE", "TRUE"} { + start := len(b) + b = append(b, name...) + b = append(b, 0) + for (len(b)-start)%8 != 0 { + b = append(b, 0) + } + } + return append(b, 0, 1) // FALSE = 0, TRUE = 1 +} + +// hdf5StringType writes a fixed-length string datatype message padded +// with NUL bytes, the padding the reference writes for fixed strings. +func hdf5StringType(width int) []byte { + b := []byte{0x13, 0x01, 0, 0} // version 1, class 3, NUL-padded, ASCII + return binary.LittleEndian.AppendUint32(b, uint32(width)) +} + +// hdf5SpaceV1 writes a version 1 dataspace message: the shape, with +// the maximum dimensions (the same extents) behind it, as the +// reference writes. A scalar dataset declares rank 0. +func hdf5SpaceV1(shape []int) []byte { + if len(shape) == 0 { + return []byte{1, 0, 0, 0, 0, 0, 0, 0} + } + b := []byte{1, byte(len(shape)), 0x01, 0} + b = binary.LittleEndian.AppendUint32(b, 0) + for _, d := range shape { + b = binary.LittleEndian.AppendUint64(b, uint64(d)) + } + for _, d := range shape { + b = binary.LittleEndian.AppendUint64(b, uint64(d)) + } + return b +} + +// hdf5SpaceV2 writes a version 2 dataspace message, the one the latest +// layout uses: the fourth byte names the dataspace class, which the +// reference library sets to simple (one) for every ranked extent. +func hdf5SpaceV2(shape []int) []byte { + if len(shape) == 0 { + return []byte{2, 0, 0, 0} // a scalar: class 0, no dimensions + } + b := []byte{2, byte(len(shape)), 0x01, 1} + for _, d := range shape { + b = binary.LittleEndian.AppendUint64(b, uint64(d)) + } + for _, d := range shape { + b = binary.LittleEndian.AppendUint64(b, uint64(d)) + } + return b +} + +// hdf5FillValueMsg writes the fill value message: the version 2 form +// of the classic layout declares a defined, all-zero fill (allocated +// incrementally for contiguous storage and late for chunked, as the +// reference does), the version 3 form of the latest layout carries no +// fill value at all. +func hdf5FillValueMsg(latest bool, chunked bool) []byte { + if latest { + return []byte{3, 0x0a} + } + if chunked { + return []byte{2, 3, 2, 1, 0, 0, 0, 0} + } + return []byte{2, 2, 2, 1, 0, 0, 0, 0} +} + +// hdf5LayoutContiguous writes a version 3 contiguous layout message: +// the data's address and its byte extent. An empty dataset names no +// storage: the undefined address stands for it. +func hdf5LayoutContiguous(addr, size uint64) []byte { + b := []byte{3, 1} + b = binary.LittleEndian.AppendUint64(b, addr) + return binary.LittleEndian.AppendUint64(b, size) +} + +// hdf5LayoutChunked writes a version 3 chunked layout message: the +// B-tree address, then one more dimension than the dataset has, the +// chunk's shape followed by the element size. +func hdf5LayoutChunked(addr uint64, chunk []int, width int) []byte { + b := []byte{3, 2, byte(len(chunk) + 1)} + b = binary.LittleEndian.AppendUint64(b, addr) + for _, c := range chunk { + b = binary.LittleEndian.AppendUint32(b, uint32(c)) + } + return binary.LittleEndian.AppendUint32(b, uint32(width)) +} + +// hdf5OutFilter is one entry of a filter pipeline message: the +// identifier, the name the reference stores and the client values. +type hdf5OutFilter struct { + id uint16 + name string + values []uint32 +} + +// hdf5FilterMessage writes a version 1 filter pipeline message. Every +// entry carries the optional flag the reference sets, its name padded +// to eight bytes and its client values padded to the same boundary. +func hdf5FilterMessage(filters []hdf5OutFilter) []byte { + b := []byte{1, byte(len(filters)), 0, 0, 0, 0, 0, 0} + for _, f := range filters { + name := append([]byte(f.name), 0) + b = binary.LittleEndian.AppendUint16(b, f.id) + b = binary.LittleEndian.AppendUint16(b, uint16(len(name))) + b = binary.LittleEndian.AppendUint16(b, 1) + b = binary.LittleEndian.AppendUint16(b, uint16(len(f.values))) + b = hdf5PadField(append(b, name...)) + for _, v := range f.values { + b = binary.LittleEndian.AppendUint32(b, v) + } + b = hdf5AppendAlign(b) + } + return b +} + +// hdf5AttrMessage writes one version 1 attribute message: the name +// padded to eight bytes, the datatype padded to eight, the dataspace, +// then the value parsed out of its text. The reader renders every +// attribute it reads as text, so the writer parses text back into the +// typed attribute it names: a whole number becomes int64, a decimal +// float64, a bracketed list an int64 or float64 array and anything +// else a fixed-length string. +func hdf5AttrMessage(a hdf5Attr) (hdf5OutMsg, error) { + dtype, value, shape, err := hdf5AttrValue(a.name, a.text) + if err != nil { + return hdf5OutMsg{}, err + } + // The name size counts the name and its NUL terminator; the field + // itself is padded to eight bytes. + nameField := hdf5PadField(append([]byte(a.name), 0)) + space := hdf5SpaceV1(shape) + body := []byte{1, 0} + body = binary.LittleEndian.AppendUint16(body, uint16(len(a.name)+1)) + body = binary.LittleEndian.AppendUint16(body, uint16(len(dtype))) + body = binary.LittleEndian.AppendUint16(body, uint16(len(space))) + body = append(body, nameField...) + body = hdf5PadField(append(body, dtype...)) + body = append(body, space...) + body = append(body, value...) + return hdf5OutMsg{typ: hdf5MsgAttribute, data: body}, nil +} + +// hdf5AttrValue parses an attribute's text into a datatype message, +// the raw value bytes and the dataspace shape (nil for a scalar). It +// is the inverse of the reader's attribute rendering: FormatInt and +// FormatFloat output parse back to the values they were printed from, +// bit for bit, and every other text becomes a fixed-length string. +func hdf5AttrValue(name, text string) ([]byte, []byte, []int, error) { + if strings.HasPrefix(text, "[") && strings.HasSuffix(text, "]") { + inner := strings.TrimSpace(text[1 : len(text)-1]) + if inner == "" { + // An empty array: an int64 attribute of extent zero, which + // the reader renders back as "[]". + return hdf5IntType(8, true), nil, []int{0}, nil + } + parts := strings.Split(inner, ", ") + ints := make([]int64, 0, len(parts)) + floats := make([]float64, 0, len(parts)) + asFloat := false + for i, p := range parts { + if v, err := strconv.ParseInt(p, 10, 64); err == nil && !asFloat { + ints = append(ints, v) + floats = append(floats, float64(v)) + continue + } + v, err := strconv.ParseFloat(p, 64) + if err != nil { + return nil, nil, nil, base.Errf("the attribute %q holds the array value %q whose element %d is not a number", name, text, i) + } + asFloat = true + floats = append(floats, v) + } + if asFloat { + raw := make([]byte, 0, 8*len(floats)) + for _, v := range floats { + raw = binary.LittleEndian.AppendUint64(raw, math.Float64bits(v)) + } + return hdf5Float64Type, raw, []int{len(floats)}, nil + } + raw := make([]byte, 0, 8*len(ints)) + for _, v := range ints { + raw = binary.LittleEndian.AppendUint64(raw, uint64(v)) + } + return hdf5IntType(8, true), raw, []int{len(ints)}, nil + } + if v, err := strconv.ParseInt(text, 10, 64); err == nil { + return hdf5IntType(8, true), binary.LittleEndian.AppendUint64(nil, uint64(v)), nil, nil + } + if v, err := strconv.ParseFloat(text, 64); err == nil { + return hdf5Float64Type, binary.LittleEndian.AppendUint64(nil, math.Float64bits(v)), nil, nil + } + width := max(len(text), 1) + value := append([]byte(text), make([]byte, width-len(text))...) + return hdf5StringType(width), value, nil, nil +} + +// writeDataset writes one dataset: the dataspace, datatype and fill +// value messages, the filter pipeline and layout of its storage and +// its attributes, then the header of the layout the file uses. +func (w *hdf5Writer) writeDataset(s *hdf5OutSet) (uint64, error) { + msgs := make([]hdf5OutMsg, 0, 6) + if w.latest { + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataspace, data: hdf5SpaceV2(s.shape)}) + } else { + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataspace, data: hdf5SpaceV1(s.shape)}) + } + switch s.class { + case 0: + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5IntType(s.width, s.signed)}) + case 1: + if s.width == 4 { + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5Float32Type}) + } else { + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5Float64Type}) + } + case 3: + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5StringType(s.width)}) + case 8: + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5BoolType()}) + default: + // Every class the plan can produce has its message case above; + // an unlisted one would emit a file with no datatype message, + // which the reader refuses later with less context than this. + return 0, base.Errf("SaveHDF5: dataset %q: the writer emits no datatype message for class %d", s.path, s.class) + } + // Chunked storage is the filter pipeline's only carrier, so a + // dataset with filters goes chunked; strings keep their raw bytes, + // and a dataset with no elements has nothing for the filters to + // compress, so it stays contiguous at an undefined address either + // way. + chunked := len(s.shape) > 0 && s.nbytes > 0 && (w.opts.Gzip != 0 || w.opts.Shuffle) && s.class != 3 + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgFillValue, data: hdf5FillValueMsg(w.latest, chunked)}) + if chunked { + entries := w.filterEntries(s.width) + if len(entries) > 0 { + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgFilterPipeline, data: hdf5FilterMessage(entries)}) + } + tree, chunk, err := w.writeChunks(s) + if err != nil { + return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, err) + } + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataLayout, data: hdf5LayoutChunked(tree, chunk, s.width)}) + } else { + addr := uint64(math.MaxUint64) + var size uint64 + if s.nbytes > 0 { + // The payload is encoded where the format will read it: + // the reservation is its final address, so the values pass + // through the writer once instead of building a block and + // then copying it in. + addr = w.reserve(s.nbytes) + if encErr := s.encode(w.buf[addr:]); encErr != nil { + return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, encErr) + } + w.pad8() + size = uint64(s.nbytes) + } + msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataLayout, data: hdf5LayoutContiguous(addr, size)}) + } + for _, a := range s.attrs { + m, err := hdf5AttrMessage(a) + if err != nil { + return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, err) + } + msgs = append(msgs, m) + } + var image []byte + var err error + if w.latest { + image, err = hdf5HeaderV2(msgs) + } else { + image, err = hdf5HeaderV1(msgs) + } + if err != nil { + return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, err) + } + return w.bytes(image), nil +} diff --git a/io/hdf5write_test.go b/io/hdf5write_test.go new file mode 100644 index 0000000..561ff27 --- /dev/null +++ b/io/hdf5write_test.go @@ -0,0 +1,1276 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "bytes" + "encoding/binary" + "fmt" + "math" + "os" + "path/filepath" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The writer tests round-trip through the reader: every file the +// writer produces is read back with LoadHDF5, the verified decoders +// of hdf5.go, and compared element for element, integers exactly and +// floats bit for bit. The reference fixtures under testdata/h5 close +// the circle: read, written, read again. + +// hdf5WriteRead writes one file and reads it back, the shape every +// test below takes. +func hdf5WriteRead(t *testing.T, sets []HDF5Dataset, groupAttrs map[string]map[string]string, opts HDF5WriteOptions) []HDF5Dataset { + t.Helper() + path := filepath.Join(t.TempDir(), "written.h5") + if err := SaveHDF5(path, sets, groupAttrs, opts); err != nil { + t.Fatalf("SaveHDF5: %v", err) + } + back, err := LoadHDF5(path) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + return back +} + +// hdf5CompareSets compares two dataset listings by path: shapes, dtypes, +// values and attributes, integers exactly and floats bit for bit. Both +// listings are compared in the reader's path order. +func hdf5CompareSets(t *testing.T, want, got []HDF5Dataset) { + t.Helper() + byPath := func(a, b HDF5Dataset) int { return strings.Compare(a.Path, b.Path) } + slices.SortFunc(want, byPath) + slices.SortFunc(got, byPath) + if len(want) != len(got) { + t.Fatalf("datasets = %d, want %d", len(got), len(want)) + } + for i := range want { + w, g := want[i], got[i] + if g.Path != w.Path { + t.Fatalf("dataset %d path = %q, want %q", i, g.Path, w.Path) + } + if !slices.Equal(g.Shape, w.Shape) { + t.Fatalf("%s shape = %v, want %v", w.Path, g.Shape, w.Shape) + } + if g.Values.Dtype() != w.Values.Dtype() { + t.Fatalf("%s dtype = %s, want %s", w.Path, g.Values.Dtype(), w.Values.Dtype()) + } + switch w.Values.Dtype() { + case core.Int: + gi, wi := g.Values.RawInts(), w.Values.RawInts() + for j := range wi { + if gi[j] != wi[j] { + t.Fatalf("%s[%d] = %d, want %d", w.Path, j, gi[j], wi[j]) + } + } + case core.Float: + gf, wf := g.Values.RawFloats(), w.Values.RawFloats() + for j := range wf { + if math.Float64bits(gf[j]) != math.Float64bits(wf[j]) { + t.Fatalf("%s[%d] = %v, want %v (bit-exact)", w.Path, j, gf[j], wf[j]) + } + } + case core.Float32: + gf, wf := g.Values.RawFloat32s(), w.Values.RawFloat32s() + for j := range wf { + if math.Float32bits(gf[j]) != math.Float32bits(wf[j]) { + t.Fatalf("%s[%d] = %v, want %v (bit-exact)", w.Path, j, gf[j], wf[j]) + } + } + case core.Bool: + gb, wb := g.Values.RawBools(), w.Values.RawBools() + for j := range wb { + if gb[j] != wb[j] { + t.Fatalf("%s[%d] = %v, want %v", w.Path, j, gb[j], wb[j]) + } + } + case core.Int8: + gi, wi := g.Values.RawInt8s(), w.Values.RawInt8s() + for j := range wi { + if gi[j] != wi[j] { + t.Fatalf("%s[%d] = %d, want %d", w.Path, j, gi[j], wi[j]) + } + } + case core.Uint8: + gi, wi := g.Values.RawUint8s(), w.Values.RawUint8s() + for j := range wi { + if gi[j] != wi[j] { + t.Fatalf("%s[%d] = %d, want %d", w.Path, j, gi[j], wi[j]) + } + } + case core.Int16: + gi, wi := g.Values.RawInt16s(), w.Values.RawInt16s() + for j := range wi { + if gi[j] != wi[j] { + t.Fatalf("%s[%d] = %d, want %d", w.Path, j, gi[j], wi[j]) + } + } + case core.Uint16: + gi, wi := g.Values.RawUint16s(), w.Values.RawUint16s() + for j := range wi { + if gi[j] != wi[j] { + t.Fatalf("%s[%d] = %d, want %d", w.Path, j, gi[j], wi[j]) + } + } + case core.Int32: + gi, wi := g.Values.RawInt32s(), w.Values.RawInt32s() + for j := range wi { + if gi[j] != wi[j] { + t.Fatalf("%s[%d] = %d, want %d", w.Path, j, gi[j], wi[j]) + } + } + case core.Uint32: + gi, wi := g.Values.RawUint32s(), w.Values.RawUint32s() + for j := range wi { + if gi[j] != wi[j] { + t.Fatalf("%s[%d] = %d, want %d", w.Path, j, gi[j], wi[j]) + } + } + default: + t.Fatalf("%s: the comparison pins no values for dtype %s", w.Path, w.Values.Dtype()) + } + if len(g.Attrs) != len(w.Attrs) { + t.Fatalf("%s attrs = %v, want %v", w.Path, g.Attrs, w.Attrs) + } + for k, v := range w.Attrs { + if g.Attrs[k] != v { + t.Fatalf("%s attr %q = %q, want %q", w.Path, k, g.Attrs[k], v) + } + } + } +} + +// TestHDF5WriteRoundTrip writes every numeric dtype in every rank the +// model carries, contiguous, and reads the values back exactly. +func TestHDF5WriteRoundTrip(t *testing.T) { + floats, err := core.FromFloats([]float64{1.5, -2.5, math.Pi, math.MaxFloat64, math.SmallestNonzeroFloat64, 0, math.Inf(-1)}, 7) + if err != nil { + t.Fatal(err) + } + matrix, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + cube, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2) + if err != nil { + t.Fatal(err) + } + singles, err := core.FromFloat32s([]float32{1, -1.5, 3.25, 1e-20, math.MaxFloat32}, 5) + if err != nil { + t.Fatal(err) + } + ints, err := core.FromInts([]int64{0, 1, -1, math.MaxInt64, math.MinInt64 + 1, 1 << 40}, 2, 3) + if err != nil { + t.Fatal(err) + } + empty, err := core.FromInts([]int64{}, 0, 3) + if err != nil { + t.Fatal(err) + } + want := []HDF5Dataset{ + {Path: "/empty", Shape: empty.Shape(), Values: empty}, + {Path: "/floats", Shape: floats.Shape(), Values: floats}, + {Path: "/ints", Shape: ints.Shape(), Values: ints}, + {Path: "/matrix", Shape: matrix.Shape(), Values: matrix}, + {Path: "/singles", Shape: singles.Shape(), Values: singles}, + {Path: "/cube", Shape: cube.Shape(), Values: cube}, + } + got := hdf5WriteRead(t, want, nil, HDF5WriteOptions{}) + hdf5CompareSets(t, want, got) + // The same again through the latest layout, which stores every + // dataset contiguously with checksummed headers. + got = hdf5WriteRead(t, want, nil, HDF5WriteOptions{Latest: true}) + hdf5CompareSets(t, want, got) +} + +// TestHDF5WriteGroupsAndAttrs writes nested groups with attributes on +// the root, the groups and the datasets, in both layouts, and checks +// what the reader reports: the dataset's own attributes and the ones +// inherited from the groups above it. +func TestHDF5WriteGroupsAndAttrs(t *testing.T) { + for _, latest := range []bool{false, true} { + inner, err := core.FromFloats([]float64{1, 2}, 2) + if err != nil { + t.Fatal(err) + } + outer, err := core.FromInts([]int64{7}, 1) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + {Path: "/a/b/inner", Shape: inner.Shape(), Values: inner, Attrs: map[string]string{"units": "m/s"}}, + {Path: "/outer", Shape: outer.Shape(), Values: outer, Attrs: map[string]string{"note": "top level"}}, + } + groupAttrs := map[string]map[string]string{ + "/": {"title": "grouped"}, + "/a": {"units": "group units"}, + "/a/b": {"depth": "2"}, + } + got := hdf5WriteRead(t, sets, groupAttrs, HDF5WriteOptions{Latest: latest}) + byPath := map[string]HDF5Dataset{} + for _, d := range got { + byPath[d.Path] = d + } + innerGot, ok := byPath["/a/b/inner"] + if !ok { + t.Fatalf("latest=%v: /a/b/inner is missing in %v", latest, hdf5Paths(got)) + } + // The dataset's own attribute and everything inherited from the + // groups above it, the nearest group winning. + for k, want := range map[string]string{"units": "m/s", "title": "grouped", "depth": "2"} { + if got := innerGot.Attrs[k]; got != want { + t.Fatalf("latest=%v: /a/b/inner attr %q = %q, want %q", latest, k, got, want) + } + } + outerGot := byPath["/outer"] + for k, want := range map[string]string{"note": "top level", "title": "grouped"} { + if got := outerGot.Attrs[k]; got != want { + t.Fatalf("latest=%v: /outer attr %q = %q, want %q", latest, k, got, want) + } + } + } +} + +// hdf5Paths lists the paths of a dataset listing, for error messages. +func hdf5Paths(sets []HDF5Dataset) []string { + out := make([]string, len(sets)) + for i, d := range sets { + out[i] = d.Path + } + return out +} + +// TestHDF5WriteChunked writes chunked datasets through the deflate +// filter, the only route this writer takes to chunked storage: whole +// dataset chunks, edge chunks hanging past the shape, and enough +// chunks to force a multi-level chunk B-tree. +func TestHDF5WriteChunked(t *testing.T) { + // ChunkBytes 64 splits the float64 axis into eight-element chunks; + // 598 elements make 75 chunks, which is past the sixty-four + // children one chunk B-tree node holds, and the last chunk hangs + // six elements past the shape. + want := make([]float64, 598) + for i := range want { + want[i] = float64(i) * 1.5 + } + values, err := core.FromFloats(want, 598) + if err != nil { + t.Fatal(err) + } + // A three-dimensional dataset whose extent is not a whole number of + // chunks along the trailing axis: the edge chunks are stored padded. + rows := make([]float64, 0, 5*3*7) + for i := range 5 * 3 * 7 { + rows = append(rows, float64(i)) + } + matrix, err := core.FromFloats(rows, 5, 3, 7) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + {Path: "/long", Shape: values.Shape(), Values: values}, + {Path: "/box", Shape: matrix.Shape(), Values: matrix}, + } + got := hdf5WriteRead(t, sets, nil, HDF5WriteOptions{Gzip: 1, ChunkBytes: 64}) + hdf5CompareSets(t, sets, got) +} + +// TestHDF5WriteFilters round-trips the deflate and shuffle filters at +// several levels, alone and together, on data whose chunking crosses +// chunk B-tree leaves. +func TestHDF5WriteFilters(t *testing.T) { + want := make([]float64, 1000) + for i := range want { + // Runs of equal values compress; the tail differs so the + // round-trip is not trivially constant. + if i < 900 { + want[i] = float64(i / 10) + } else { + want[i] = math.Sin(float64(i)) + } + } + values, err := core.FromFloats(want, 1000) + if err != nil { + t.Fatal(err) + } + ints, err := core.FromInts([]int64{5, 3, 5, 3, 5, 3, 5, 3}, 8) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + {Path: "/gzip_shuffle", Shape: values.Shape(), Values: values}, + {Path: "/ints", Shape: ints.Shape(), Values: ints}, + } + for _, opts := range []HDF5WriteOptions{ + {Gzip: 1, ChunkBytes: 128}, + {Gzip: 6, ChunkBytes: 128}, + {Gzip: 9, ChunkBytes: 128}, + {Gzip: -1, ChunkBytes: 128}, + {Shuffle: true, ChunkBytes: 128}, + {Shuffle: true, Gzip: 4, ChunkBytes: 64}, + {Shuffle: true, Gzip: 6}, + } { + got := hdf5WriteRead(t, sets, nil, opts) + hdf5CompareSets(t, sets, got) + } + // The deflate stage genuinely compresses: eight hundred bytes of + // repeating values deflate to a fraction of themselves. (The file + // as a whole carries the B-tree nodes the format allocates, which + // dwarf the data of a test-sized dataset.) + flat := bytes.Repeat([]byte{0, 0, 0, 0, 0, 0, 0xF0, 0x3F}, 100) + small, err := (&hdf5Writer{}).deflateChunk(flat, 6) + if err != nil { + t.Fatal(err) + } + if len(small) >= len(flat) { + t.Fatalf("the deflate filter stored %d bytes for %d of data", len(small), len(flat)) + } +} + +// TestHDF5WriteWideGroup writes a group with more children than two +// symbol table nodes hold, which forces a multi-level group B-tree, +// and reads every child back by path. +func TestHDF5WriteWideGroup(t *testing.T) { + // 260 children: 33 symbol table nodes, which is past the + // thirty-two children one group B-tree node holds, so the tree + // grows a second level. + var want []HDF5Dataset + for i := range 260 { + v, err := core.FromInts([]int64{int64(i)}, 1) + if err != nil { + t.Fatal(err) + } + want = append(want, HDF5Dataset{Path: fmt.Sprintf("/g/d%03d", i), Shape: v.Shape(), Values: v}) + } + for _, latest := range []bool{false, true} { + got := hdf5WriteRead(t, want, nil, HDF5WriteOptions{Latest: latest}) + hdf5CompareSets(t, want, got) + } +} + +// TestHDF5WriteAttrs pins the attribute round-trip: whole numbers +// return as int64 attributes, decimals as float64 ones, bracketed +// lists as arrays, and any other text as a fixed-length string. +func TestHDF5WriteAttrs(t *testing.T) { + values, err := core.FromInts([]int64{1}, 1) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + {Path: "/x", Shape: values.Shape(), Values: values, Attrs: map[string]string{ + "n": "7", + "neg": "-12", + "x": "2.5", + "tiny": "1e-300", + "arr": "[1, 2, 3]", + "floats": "[0.5, -1.25, 2.75]", + "label": "hello world", + "short": "K", + "unicode": "vlno moe", + "spacey": " padded ", + }}, + } + got := hdf5WriteRead(t, sets, nil, HDF5WriteOptions{}) + for k, want := range sets[0].Attrs { + if got[0].Attrs[k] != want { + t.Fatalf("attr %q = %q, want %q", k, got[0].Attrs[k], want) + } + } + // A float64 attribute of a whole value reads back as "3" and is + // rewritten as the int64 attribute "3": the value is preserved, + // the type is the writer's inference. + floaty, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + sets = []HDF5Dataset{{Path: "/x", Shape: floaty.Shape(), Values: floaty, Attrs: map[string]string{"v": "3"}}} + got = hdf5WriteRead(t, sets, nil, HDF5WriteOptions{}) + if got[0].Attrs["v"] != "3" { + t.Fatalf("attr v = %q, want %q", got[0].Attrs["v"], "3") + } + // Non-numeric array elements are refused by name. + sets[0].Attrs = map[string]string{"bad": "[1, x]"} + path := filepath.Join(t.TempDir(), "bad.h5") + if err := SaveHDF5(path, sets, nil); err == nil { + t.Fatal("expected an error for a non-numeric array attribute") + } else if !strings.Contains(err.Error(), "not a number") { + t.Fatalf("error = %v, want a not-a-number refusal", err) + } +} + +// TestHDF5WriteText writes fixed-length string datasets, checks the +// datatype message they produce against the reader's own decoder, and +// confirms the reader refuses them exactly as its documentation says. +func TestHDF5WriteText(t *testing.T) { + path := filepath.Join(t.TempDir(), "text.h5") + texts := []HDF5TextDataset{ + {Path: "/words", Shape: []int{4}, Text: []string{"forty two", "hi", "", "K"}}, + {Path: "/grid", Shape: []int{2, 2}, Text: []string{"a", "bb", "ccc", "dddd"}}, + } + if err := SaveHDF5Text(path, texts, HDF5WriteOptions{}); err != nil { + t.Fatalf("SaveHDF5Text: %v", err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + // The words dataset stores four elements of nine bytes: the + // datatype message is the fixed-length string class the reader + // decodes for attributes, and its width is the longest element. + pattern := []byte{0x13, 0x01, 0, 0} + pattern = binary.LittleEndian.AppendUint32(pattern, 9) + at := bytes.Index(raw, pattern) + if at < 0 { + t.Fatalf("no nine-byte fixed-string datatype message in %d bytes", len(raw)) + } + dtype, err := decodeType(raw[at:at+8], true) + if err != nil { + t.Fatalf("decodeType: %v", err) + } + if dtype.class != 3 || dtype.width != 9 { + t.Fatalf("datatype = class %d width %d, want class 3 width 9", dtype.class, dtype.width) + } + // The reader refuses string datasets by name. + if _, err := LoadHDF5(path); err == nil { + t.Fatal("expected the reader to refuse a string dataset") + } else if !strings.Contains(err.Error(), "string datasets are not supported") { + t.Fatalf("error = %v, want the string dataset refusal", err) + } + // The latest layout carries them too, still refused by the reader. + path = filepath.Join(t.TempDir(), "text3.h5") + if err := SaveHDF5Text(path, texts[:1], HDF5WriteOptions{Latest: true}); err != nil { + t.Fatalf("SaveHDF5Text: %v", err) + } + if _, err := LoadHDF5(path); err == nil { + t.Fatal("expected the reader to refuse a latest-format string dataset") + } + // Filters are refused by name rather than silently dropped. + if err := SaveHDF5Text(filepath.Join(t.TempDir(), "z.h5"), texts, HDF5WriteOptions{Gzip: 6}); err == nil { + t.Fatal("expected an error for filters on string datasets") + } + // A element count that does not fill the shape is refused. + if err := SaveHDF5Text(filepath.Join(t.TempDir(), "s.h5"), []HDF5TextDataset{{Path: "/s", Shape: []int{3}, Text: []string{"a", "b"}}}, HDF5WriteOptions{}); err == nil { + t.Fatal("expected an error for a short text list") + } +} + +// TestHDF5WriteLatest pins the latest layout against its checksums: +// the file reads back, and a single flipped byte inside either +// checksummed region, the superblock or an object header, refuses the +// read. The classic layout carries no such checksums, and a flipped +// data byte there reads back corrupted but without error. +func TestHDF5WriteLatest(t *testing.T) { + values, err := core.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + t.Fatal(err) + } + deep, err := core.FromInts([]int64{9, 8}, 2) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + // The root attribute reaches every dataset on read, the + // reader's inheritance rule, so the expected listings carry it. + {Path: "/g/d", Shape: values.Shape(), Values: values, Attrs: map[string]string{"title": "latest"}}, + {Path: "/g/h/e", Shape: deep.Shape(), Values: deep, Attrs: map[string]string{"title": "latest"}}, + } + path := filepath.Join(t.TempDir(), "latest.h5") + if err := SaveHDF5(path, sets, map[string]map[string]string{"/": {"title": "latest"}}, HDF5WriteOptions{Latest: true}); err != nil { + t.Fatal(err) + } + got, err := LoadHDF5(path) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + hdf5CompareSets(t, sets, got) + // The superblock carries a lookup3 checksum over its first forty + // bytes; byte 11 is the consistency flags, which nothing else + // reads. The root object header, found the way the reader finds + // it, checksums signature to last message, and the writer lays the + // datasets' data out in front of it, so a flip in the first bytes + // past the superblock lands in unchecksummed data instead. + whole, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + if got := hdf5Lookup3(whole[:44]); got != binary.LittleEndian.Uint32(whole[44:]) { + t.Fatalf("the written superblock fails its own checksum: %#08x", got) + } + f, err := newHDF5File(whole) + if err != nil { + t.Fatal(err) + } + for _, c := range []struct { + name string + at int + }{ + {"superblock", 11}, + {"object header", int(f.rootAddress) + 8}, + } { + corrupt := slices.Clone(whole) + corrupt[c.at] ^= 0xff + bad := filepath.Join(t.TempDir(), "corrupt.h5") + if err := os.WriteFile(bad, corrupt, 0o644); err != nil { + t.Fatal(err) + } + if _, err := LoadHDF5(bad); err == nil { + t.Fatalf("expected an error for a corrupted %s", c.name) + } else if !strings.Contains(err.Error(), "checksum") { + t.Fatalf("corrupted %s: error = %v, want a checksum refusal", c.name, err) + } + } + // The classic layout carries no checksums: a flipped data byte + // reads back changed, silently. + classic := filepath.Join(t.TempDir(), "classic.h5") + if err := SaveHDF5(classic, sets[:1], nil); err != nil { + t.Fatal(err) + } + whole, err = os.ReadFile(classic) + if err != nil { + t.Fatal(err) + } + // The dataset holds 1, 2, 3: the little-endian bytes of the value 1 + // sit in the data region and nowhere else in a file this small. + one := binary.LittleEndian.AppendUint64(nil, math.Float64bits(1)) + at := bytes.Index(whole, one) + if at < 0 { + t.Fatal("the value 1 is not in the file") + } + whole[at] ^= 0xff + flipped := filepath.Join(t.TempDir(), "flipped.h5") + if err := os.WriteFile(flipped, whole, 0o644); err != nil { + t.Fatal(err) + } + back, err := LoadHDF5(flipped) + if err != nil { + t.Fatalf("LoadHDF5: %v", err) + } + if back[0].Values.RawFloats()[0] == 1 { + t.Fatal("the flipped byte did not reach the values") + } +} + +// TestHDF5WriteFixtureRoundTrip reads each reference fixture, writes +// it back and reads again: paths, shapes, dtypes, values and +// attributes must agree exactly, in the dtypes the fixtures store +// natively. fletcher.h5 is stored chunked and compressed; the +// round-trip rewrites it contiguous, so only the values and attributes +// are the contract there. fixture.h5's int32 dataset is written as +// int32, which the writer stores at its own width and the reader lands +// int32 again. +func TestHDF5WriteFixtureRoundTrip(t *testing.T) { + for _, c := range []struct { + name string + latest bool + }{ + {"fixture.h5", false}, + {"fletcher.h5", false}, + {"latest.h5", true}, + } { + sets, err := LoadHDF5(h5Fixture(t, c.name)) + if err != nil { + t.Fatalf("%s: %v", c.name, err) + } + got := hdf5WriteRead(t, sets, nil, HDF5WriteOptions{Latest: c.latest}) + hdf5CompareSets(t, sets, got) + } +} + +// TestHDF5WriteRefusals pins the errors: every malformed input is +// refused with a message that names the defect, before any bytes are +// written. +func TestHDF5WriteRefusals(t *testing.T) { + dir := t.TempDir() + values, err := core.FromInts([]int64{1}, 1) + if err != nil { + t.Fatal(err) + } + good := HDF5Dataset{Path: "/x", Shape: values.Shape(), Values: values} + bad := func(sets []HDF5Dataset, attrs map[string]map[string]string, opts HDF5WriteOptions) error { + return SaveHDF5(filepath.Join(dir, "bad.h5"), sets, attrs, opts) + } + for _, c := range []struct { + name string + err error + text string + }{ + {"relative path", bad([]HDF5Dataset{{Path: "x", Shape: values.Shape(), Values: values}}, nil, HDF5WriteOptions{}), "absolute"}, + {"root path", bad([]HDF5Dataset{{Path: "/", Shape: values.Shape(), Values: values}}, nil, HDF5WriteOptions{}), "root"}, + {"empty segment", bad([]HDF5Dataset{{Path: "/a//b", Shape: values.Shape(), Values: values}}, nil, HDF5WriteOptions{}), "empty segment"}, + {"trailing slash", bad([]HDF5Dataset{{Path: "/a/", Shape: values.Shape(), Values: values}}, nil, HDF5WriteOptions{}), "empty segment"}, + {"duplicate", bad([]HDF5Dataset{good, good}, nil, HDF5WriteOptions{}), "twice"}, + {"group conflict", bad([]HDF5Dataset{good, {Path: "/x/y", Shape: values.Shape(), Values: values}}, nil, HDF5WriteOptions{}), "both a dataset and a group"}, + {"no values", bad([]HDF5Dataset{{Path: "/x"}}, nil, HDF5WriteOptions{}), "no values"}, + {"unknown attribute group", bad([]HDF5Dataset{good}, map[string]map[string]string{"/nope": {"a": "1"}}, HDF5WriteOptions{}), "does not name a group"}, + {"empty attribute name", bad([]HDF5Dataset{good}, map[string]map[string]string{"/": {"": "1"}}, HDF5WriteOptions{}), "empty name"}, + {"nul in path", bad([]HDF5Dataset{{Path: "/a\x00b", Shape: values.Shape(), Values: values}}, nil, HDF5WriteOptions{}), "NUL"}, + {"nul in attribute", bad([]HDF5Dataset{good}, map[string]map[string]string{"/": {"a\x00": "1"}}, HDF5WriteOptions{}), "NUL"}, + {"latest with gzip", bad([]HDF5Dataset{good}, nil, HDF5WriteOptions{Latest: true, Gzip: 6}), "version 2 B-tree"}, + {"latest with shuffle", bad([]HDF5Dataset{good}, nil, HDF5WriteOptions{Latest: true, Shuffle: true}), "version 2 B-tree"}, + {"gzip level", bad([]HDF5Dataset{good}, nil, HDF5WriteOptions{Gzip: 10}), "gzip level"}, + {"negative chunk target", bad([]HDF5Dataset{good}, nil, HDF5WriteOptions{ChunkBytes: -1}), "negative"}, + {"disagreeing shape", bad([]HDF5Dataset{{Path: "/x", Shape: []int{2, 1}, Values: values}}, nil, HDF5WriteOptions{}), "declares a shape"}, + } { + if c.err == nil { + t.Errorf("%s: expected an error, got none", c.name) + continue + } + if !strings.Contains(c.err.Error(), c.text) { + t.Errorf("%s: error = %v, want it to name %q", c.name, c.err, c.text) + } + } + // A complex dataset is refused: the format carries it, the core + // does not. + complexes, err := core.FromComplexes([]complex128{1}, 1) + if err != nil { + t.Fatal(err) + } + err = bad([]HDF5Dataset{{Path: "/z", Shape: complexes.Shape(), Values: complexes}}, nil, HDF5WriteOptions{}) + if err == nil || !strings.Contains(err.Error(), "not supported") { + t.Fatalf("complex dtype: error = %v, want an unsupported-dtype refusal", err) + } + // A chunk target so small that the chunk count explodes is refused + // with the lever named; the shape is laid out directly because + // materialising that many elements would cost more than the + // refusal is worth. + chunky := &hdf5Writer{opts: HDF5WriteOptions{ChunkBytes: 8}} + _, _, cerr := chunky.writeChunks(&hdf5OutSet{path: "/w", shape: []int{hdf5MaxChunks + 1}, width: 8}) + if cerr == nil || !strings.Contains(cerr.Error(), "chunks") { + t.Fatalf("tiny chunks: error = %v, want a chunk-count refusal", cerr) + } + // A text shape whose extents wrap the element count onto zero + // matches an empty Text slice exactly, so the length check alone + // cannot refuse it: the byte budget must, before a header goes out + // declaring 2^64 elements. + terr := SaveHDF5Text(filepath.Join(dir, "bad-text.h5"), + []HDF5TextDataset{{Path: "/t", Shape: []int{4294967296, 4294967296}, Text: nil}}) + if terr == nil || !strings.Contains(terr.Error(), "budget") { + t.Fatalf("wrapped text extents: error = %v, want the byte-budget refusal", terr) + } + // Nothing above may have left a valid file behind: the refusals + // happen before a byte is written. + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + for _, e := range entries { + t.Errorf("a refused write left %q behind", e.Name()) + } +} + +// TestHDF5WriteEmptyFilteredRoundTrip keeps a dataset with no elements +// out of the chunked path: the filters have nothing to compress and a +// zero extent divides the chunk grid into zero-sized chunks, so the +// dataset is stored contiguous at an undefined address in both +// layouts. +func TestHDF5WriteEmptyFilteredRoundTrip(t *testing.T) { + empty, err := core.FromFloats(nil, 0, 3) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{{Path: "/empty", Shape: empty.Shape(), Values: empty}} + got := hdf5WriteRead(t, sets, nil, HDF5WriteOptions{Gzip: 6, Shuffle: true}) + hdf5CompareSets(t, sets, got) + // The latest layout stores the empty dataset contiguous too; the + // filtered latest combination is refused by design and pinned in + // the refusals test. + got = hdf5WriteRead(t, sets, nil, HDF5WriteOptions{Latest: true}) + hdf5CompareSets(t, sets, got) +} + +// TestHDF5WriteLatestGroupBadAttrRefused pins the error the latest +// layout returns for a group attribute whose text does not parse: the +// refusal must come out of the write, not be overwritten by the +// header encode that follows it. +func TestHDF5WriteLatestGroupBadAttrRefused(t *testing.T) { + values, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{{Path: "/leaf", Shape: values.Shape(), Values: values}} + attrs := map[string]map[string]string{"/": {"bad": "[1, x]"}} + path := filepath.Join(t.TempDir(), "bad.h5") + err = SaveHDF5(path, sets, attrs, HDF5WriteOptions{Latest: true}) + if err == nil { + t.Fatal("SaveHDF5 with a malformed group attribute answered nil") + } + if !strings.Contains(err.Error(), "bad") { + t.Fatalf("error = %v, want it to name the attribute", err) + } +} + +// TestHDF5WriteInterleavedNamesRoundTrip pins the child order the +// classic symbol table requires: the heap offsets and the B-tree keys +// sort in one merged name order, so a subgroup whose name interleaves +// with the datasets' is reachable. Round-trips both layouts. +func TestHDF5WriteInterleavedNamesRoundTrip(t *testing.T) { + values, err := core.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + {Path: "/alpha", Shape: values.Shape(), Values: values}, + {Path: "/box", Shape: values.Shape(), Values: values}, + {Path: "/deep/nested", Shape: values.Shape(), Values: values}, + {Path: "/f32", Shape: values.Shape(), Values: values}, + {Path: "/zoom/in", Shape: values.Shape(), Values: values}, + } + for _, latest := range []bool{false, true} { + got := hdf5WriteRead(t, sets, nil, HDF5WriteOptions{Latest: latest}) + hdf5CompareSets(t, sets, got) + } +} + +// TestHDF5WriteChunkTreeKeyConvention parses a written chunk B-tree +// node and pins the key convention the reference files store: every +// real chunk key carries 0 in the trailing element-size coordinate and +// only the node's closing sentinel carries the element size, which is +// what orders it past the node's chunks. Coordinates ascend through +// the node. +func TestHDF5WriteChunkTreeKeyConvention(t *testing.T) { + want := make([]float64, 90) + for i := range want { + want[i] = float64(i) + } + values, err := core.FromFloats(want, 10, 9) + if err != nil { + t.Fatal(err) + } + w := &hdf5Writer{opts: HDF5WriteOptions{ChunkBytes: 4 * 8}} + s, err := hdf5PlanSet("SaveHDF5", &HDF5Dataset{Path: "/grid", Shape: values.Shape(), Values: values}) + if err != nil { + t.Fatal(err) + } + tree, chunk, err := w.writeChunks(s) + if err != nil { + t.Fatal(err) + } + // One node: a 3x3 chunk grid sits far under the fanout. + node := w.buf[tree:] + if string(node[:4]) != "TREE" { + t.Fatalf("chunk tree signature = %q, want TREE", node[:4]) + } + if node[5] != 0 { + t.Fatalf("chunk tree level = %d, want a leaf", node[5]) + } + entries := int(binary.LittleEndian.Uint16(node[6:])) + if entries != 9 { + t.Fatalf("chunk tree entries = %d, want 9 (a 3x3 grid of %v)", entries, chunk) + } + keySize := 8 + 8*(len(s.shape)+1) + p := 24 + for i := 0; i <= entries; i++ { + coords := make([]uint64, len(s.shape)+1) + for j := range coords { + coords[j] = binary.LittleEndian.Uint64(node[p+8+8*j:]) + } + last := i == entries + if want := uint64(s.width); last { + if coords[len(coords)-1] != want { + t.Fatalf("sentinel trailing coordinate = %d, want the element size %d", coords[len(coords)-1], want) + } + } else { + if coords[len(coords)-1] != 0 { + t.Fatalf("chunk %d trailing coordinate = %d, want 0", i, coords[len(coords)-1]) + } + if i > 0 && coords[len(coords)-2] <= binary.LittleEndian.Uint64(node[p-keySize+8+8*(len(coords)-2):]) && coords[0] == binary.LittleEndian.Uint64(node[p-keySize+8:]) { + t.Fatalf("chunk %d coordinates do not advance past chunk %d", i, i-1) + } + } + p += keySize + if !last { + p += 8 + } + } +} + +// TestHDF5WriteDeterministicBytes pins the writer's determinism claim: +// the same content supplied in different dataset and attribute orders +// writes byte for bit the same file. +func TestHDF5WriteDeterministicBytes(t *testing.T) { + a, err := core.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + t.Fatal(err) + } + b, err := core.FromInts([]int64{7, 8}, 2) + if err != nil { + t.Fatal(err) + } + files := make([][]byte, 2) + for i, sets := range [][]HDF5Dataset{ + {{Path: "/a", Shape: a.Shape(), Values: a}, {Path: "/b", Shape: b.Shape(), Values: b}}, + {{Path: "/b", Shape: b.Shape(), Values: b}, {Path: "/a", Shape: a.Shape(), Values: a}}, + } { + path := filepath.Join(t.TempDir(), "det.h5") + if err := SaveHDF5(path, sets, map[string]map[string]string{"/": {"x": "1", "y": "2"}}, HDF5WriteOptions{}); err != nil { + t.Fatal(err) + } + files[i], err = os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + } + if !bytes.Equal(files[0], files[1]) { + t.Fatalf("two writes of the same content in different orders differ in %d of %d bytes", countByteDiffs(files[0], files[1]), len(files[0])) + } +} + +func countByteDiffs(a, b []byte) int { + n := 0 + for i := range a { + if a[i] != b[i] { + n++ + } + } + return n +} + +// TestHDF5WriteChildOrderBytes pins the merged child order at the byte +// level, the pin the round-trip test cannot be: the reader walks the +// symbol nodes linearly and stays blind to their order, while the +// reference library binary-searches both the symbol node records and +// the B-tree keys on the local heap offsets. Ten root children whose +// names interleave groups and datasets span two symbol node leaves: +// their entries must read in one ascending heap-offset order matching +// the alphabetical names, and the group B-tree must route to the two +// leaves in order with closing keys on each leaf's last offset. +func TestHDF5WriteChildOrderBytes(t *testing.T) { + values, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + {Path: "/ann", Shape: values.Shape(), Values: values}, + {Path: "/bee/x", Shape: values.Shape(), Values: values}, + {Path: "/cee", Shape: values.Shape(), Values: values}, + {Path: "/dee/x", Shape: values.Shape(), Values: values}, + {Path: "/eff", Shape: values.Shape(), Values: values}, + {Path: "/gee/x", Shape: values.Shape(), Values: values}, + {Path: "/its/x", Shape: values.Shape(), Values: values}, + {Path: "/jay", Shape: values.Shape(), Values: values}, + {Path: "/mill", Shape: values.Shape(), Values: values}, + {Path: "/zoo/x", Shape: values.Shape(), Values: values}, + } + path := filepath.Join(t.TempDir(), "order.h5") + if err := SaveHDF5(path, sets, nil, HDF5WriteOptions{}); err != nil { + t.Fatal(err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + // The root group's leaves: every subgroup carries a symbol node of + // exactly one child, so the leaves of the root are the nodes with + // more. Eight children fit one leaf of eight slots, so ten spill + // into two, eight and two. + var leaves []int + for at := range hdf5ScanSignature(raw, []byte("SNOD")) { + if int(binary.LittleEndian.Uint16(raw[at+6:])) > 1 { + leaves = append(leaves, at) + } + } + if len(leaves) != 2 { + t.Fatalf("the root group spans %d symbol node leaves, want 2", len(leaves)) + } + // The local heap's data segment: the header names its address. + heap := raw[hdf5FirstSignature(raw, []byte("HEAP")):] + dataAddr := binary.LittleEndian.Uint64(heap[24:]) + nameAt := func(off uint64) string { + start := int(dataAddr + off) + end := start + for raw[end] != 0 { + end++ + } + return string(raw[start:end]) + } + var offsets []uint64 + var names []string + for _, at := range leaves { + entries := int(binary.LittleEndian.Uint16(raw[at+6:])) + for i := range entries { + off := binary.LittleEndian.Uint64(raw[at+8+40*i:]) + if len(offsets) > 0 && off <= offsets[len(offsets)-1] { + t.Fatalf("heap offsets not ascending across the leaves: %d then %d", offsets[len(offsets)-1], off) + } + offsets = append(offsets, off) + names = append(names, nameAt(off)) + } + } + want := []string{"ann", "bee", "cee", "dee", "eff", "gee", "its", "jay", "mill", "zoo"} + if !slices.Equal(names, want) { + t.Fatalf("children in heap-offset order = %v, want %v", names, want) + } + // The root group's B-tree: one child pointer per leaf, keyed by the + // leaf's last offset. It is the tree with two children (each + // subgroup carries a tree of its own, with one). + treeAt := -1 + for i := range hdf5ScanSignature(raw, []byte("TREE")) { + if raw[i+4] == 0 && int(binary.LittleEndian.Uint16(raw[i+6:])) == len(leaves) { + treeAt = i + break + } + } + if treeAt < 0 { + t.Fatal("no group B-tree carries the two root leaves") + } + prefix := make([]int, len(leaves)+1) + for i, at := range leaves { + prefix[i+1] = prefix[i] + int(binary.LittleEndian.Uint16(raw[at+6:])) + } + kids := int(binary.LittleEndian.Uint16(raw[treeAt+6:])) + p := treeAt + 24 + // The node reads key0, then child and closing key per child. + if key := binary.LittleEndian.Uint64(raw[p:]); key != 0 { + t.Fatalf("the first B-tree key = %d, want the null name 0", key) + } + p += 8 + prev := uint64(0) + for i := range kids { + child := int(binary.LittleEndian.Uint64(raw[p:])) + p += 8 + if child != leaves[i] { + t.Fatalf("B-tree child %d points at %d, want the symbol node at %d", i, child, leaves[i]) + } + key := binary.LittleEndian.Uint64(raw[p:]) + p += 8 + last := offsets[prefix[i+1]-1] + if key != last { + t.Fatalf("the closing key %d of child %d, want its leaf's last offset %d", key, i, last) + } + if key <= prev { + t.Fatalf("B-tree keys not ascending at child %d: %d then %d", i, prev, key) + } + prev = key + } +} + +// hdf5ScanSignature yields the offsets of every signature block. +func hdf5ScanSignature(raw, sig []byte) func(yield func(int) bool) { + return func(yield func(int) bool) { + for i := 0; i+4 <= len(raw); i++ { + if string(raw[i:i+4]) == string(sig) { + if !yield(i) { + return + } + } + } + } +} + +// hdf5FirstSignature returns the offset of the first signature block. +func hdf5FirstSignature(raw, sig []byte) int { + for at := range hdf5ScanSignature(raw, sig) { + return at + } + return -1 +} + +// TestHDF5WriteHeaderBoundRefusals pins the four bounds the header +// encoders and the plan enforce: a version 1 message past the +// countable length, a single version 2 message past its own countable +// length, an attribute name past the length the attribute message +// counts, and a text dataset past the rank the format allows. +func TestHDF5WriteHeaderBoundRefusals(t *testing.T) { + dir := t.TempDir() + values, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + huge := strings.Repeat("x", 70000) + for _, c := range []struct { + name string + write func() error + text string + }{ + {"classic message past 0xffff", func() error { + return SaveHDF5(filepath.Join(dir, "v1.h5"), + []HDF5Dataset{{Path: "/d", Shape: values.Shape(), Values: values, Attrs: map[string]string{"big": huge}}}, + nil, HDF5WriteOptions{}) + }, "a version 1 header stores"}, + {"latest message past 0xffff", func() error { + return SaveHDF5(filepath.Join(dir, "v2.h5"), + []HDF5Dataset{{Path: "/d", Shape: values.Shape(), Values: values, Attrs: map[string]string{"big": huge}}}, + nil, HDF5WriteOptions{Latest: true}) + }, "a version 2 header stores"}, + {"attribute name past 0xffff", func() error { + return SaveHDF5(filepath.Join(dir, "name.h5"), + []HDF5Dataset{{Path: "/d", Shape: values.Shape(), Values: values, Attrs: map[string]string{huge: "1"}}}, + nil, HDF5WriteOptions{}) + }, "attribute message counts"}, + {"text dataset past the rank", func() error { + shape := make([]int, hdf5MaxRank+1) + return SaveHDF5Text(filepath.Join(dir, "rank.h5"), + []HDF5TextDataset{{Path: "/t", Shape: shape, Text: nil}}, HDF5WriteOptions{}) + }, "dimensions"}, + } { + err := c.write() + if err == nil { + t.Errorf("%s: expected an error, got none", c.name) + continue + } + if !strings.Contains(err.Error(), c.text) { + t.Errorf("%s: error = %v, want it to name %q", c.name, err, c.text) + } + } + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + for _, e := range entries { + t.Errorf("a refused write left %q behind", e.Name()) + } +} + +// hdf5NarrowSets builds one dataset per dtype the narrow write round +// added, shaped [5, 3] with the values at each dtype's extremes beside +// deterministic filler. A small chunk target splits the trailing axis +// of this shape into one- or two-column chunks, so the round-trip +// configs below exercise multi-chunk grids and, for the one-byte +// dtypes, edge chunks that hang past the shape. +func hdf5NarrowSets(t *testing.T) []HDF5Dataset { + t.Helper() + must := func(values *core.Array, err error) *core.Array { + if err != nil { + t.Fatal(err) + } + return values + } + mk := []HDF5Dataset{} + add := func(path string, values *core.Array) { + mk = append(mk, HDF5Dataset{Path: path, Shape: values.Shape(), Values: values}) + } + add("/bool", must(core.FromBools([]bool{ + false, true, false, + true, true, false, + false, true, true, + false, true, false, + false, true, true}, 5, 3))) + add("/int8", must(core.FromInt8s([]int8{ + -128, 127, 0, + -1, 7, -7, + 42, -42, 100, + -100, 3, -3, + 1, -56, 64}, 5, 3))) + add("/uint8", must(core.FromUint8s([]uint8{ + 0, 255, 1, + 7, 42, 100, + 3, 200, 64, + 128, 12, 99, + 254, 17, 8}, 5, 3))) + add("/int16", must(core.FromInt16s([]int16{ + -32768, 32767, 0, + -1, 256, -256, + 1000, -1000, 3, + -3, 32766, -32767, + 17, 42, -7}, 5, 3))) + add("/uint16", must(core.FromUint16s([]uint16{ + 0, 65535, 1, + 256, 1000, 3, + 32768, 32767, 65534, + 17, 42, 999, + 8, 12, 5}, 5, 3))) + add("/int32", must(core.FromInt32s([]int32{ + -2147483648, 2147483647, 0, + -1, 65536, -65536, + 1000000, -1000000, 3, + -3, 2147483646, -2147483647, + 17, 42, -7}, 5, 3))) + add("/uint32", must(core.FromUint32s([]uint32{ + 0, 4294967295, 1, + 65536, 1000000, 3, + 2147483648, 2147483647, 4294967294, + 17, 42, 999999, + 8, 12, 5}, 5, 3))) + return mk +} + +// TestHDF5WriteNarrowRoundTrip writes every dtype the narrow write +// round added, contiguous in both layouts and chunked through the +// deflate and shuffle filters at several chunk targets, and pins the +// round trip: the landing dtype equals the written dtype (int16 reads +// back int16, bool lands bool) and every value returns exactly, at +// each dtype's extremes. +func TestHDF5WriteNarrowRoundTrip(t *testing.T) { + sets := hdf5NarrowSets(t) + for _, opts := range []HDF5WriteOptions{ + {}, + {Latest: true}, + {Gzip: 1, ChunkBytes: 12}, + {Shuffle: true, ChunkBytes: 12}, + {Gzip: 6, Shuffle: true, ChunkBytes: 32}, + } { + got := hdf5WriteRead(t, sets, nil, opts) + hdf5CompareSets(t, sets, got) + } + // Zero-length narrow datasets, shape [0]: one contiguous and one + // chunked config. The round trip has no values to compare, so the + // pin is the path, the shape and the landing dtype per dataset. + empty := hdf5NarrowEmptySets(t) + for _, opts := range []HDF5WriteOptions{ + {}, + {Gzip: 1, ChunkBytes: 12}, + } { + got := hdf5WriteRead(t, empty, nil, opts) + hdf5CompareSets(t, empty, got) + } +} + +// hdf5NarrowEmptySets builds one shape-[0] dataset per narrow dtype: +// the zero-extent round trip the narrow write path must also carry. +func hdf5NarrowEmptySets(t *testing.T) []HDF5Dataset { + t.Helper() + must := func(values *core.Array, err error) *core.Array { + if err != nil { + t.Fatal(err) + } + return values + } + mk := []HDF5Dataset{} + add := func(path string, values *core.Array) { + mk = append(mk, HDF5Dataset{Path: path, Shape: values.Shape(), Values: values}) + } + add("/bool0", must(core.FromBools([]bool{}, 0))) + add("/int80", must(core.FromInt8s([]int8{}, 0))) + add("/uint80", must(core.FromUint8s([]uint8{}, 0))) + add("/int160", must(core.FromInt16s([]int16{}, 0))) + add("/uint160", must(core.FromUint16s([]uint16{}, 0))) + add("/int320", must(core.FromInt32s([]int32{}, 0))) + add("/uint320", must(core.FromUint32s([]uint32{}, 0))) + return mk +} + +// TestHDF5WriteBoolMessage pins the boolean enumeration datatype +// message the writer emits, byte for byte, against the reader's own +// decoder: class 8 with two members in the class bit field and the +// reserved byte zero, a complete one-byte unsigned little-endian +// fixed-point base, the member names each NUL terminated and padded in +// their own field to a multiple of eight bytes, and the packed member +// values 0 and 1 behind the names. The payload bytes are the member +// values themselves. +func TestHDF5WriteBoolMessage(t *testing.T) { + values, err := core.FromBools([]bool{false, true, false}, 3) + if err != nil { + t.Fatal(err) + } + path := filepath.Join(t.TempDir(), "bool.h5") + if err := SaveHDF5(path, []HDF5Dataset{{Path: "/mask", Shape: values.Shape(), Values: values}}, nil, HDF5WriteOptions{}); err != nil { + t.Fatalf("SaveHDF5: %v", err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + pattern := []byte{ + 0x18, 0x02, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, // class 8, two members, one-byte values + 0x10, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x08, 0x00, // the unsigned base + 'F', 'A', 'L', 'S', 'E', 0, 0, 0, + 'T', 'R', 'U', 'E', 0, 0, 0, 0, + 0x00, 0x01, // FALSE = 0, TRUE = 1 + } + at := bytes.Index(raw, pattern) + if at < 0 { + t.Fatalf("no boolean enumeration datatype message in %d bytes", len(raw)) + } + dtype, err := decodeType(raw[at:at+len(pattern)], false) + if err != nil { + t.Fatalf("decodeType on the written message: %v", err) + } + if dtype.class != 8 || !dtype.isBool || dtype.width != 1 { + t.Fatalf("datatype = class %d width %d isBool %v, want the one-byte boolean convention", dtype.class, dtype.width, dtype.isBool) + } + // The payload: the member values, one byte per element, which the + // round-trip test reads back as the written booleans. + if bytes.Index(raw, []byte{0x00, 0x01, 0x00}) < 0 { + t.Fatal("the boolean payload 0, 1, 0 is not in the file") + } +} + +// TestHDF5WriteNarrowRefusals pins the refusal for a dtype outside the +// stored set: float16 and complex are named against the full list the +// writer stores, and the refusal happens before any bytes are written. +// The loud defaults of the payload and message switches are reached +// directly: a datatype class the plan never produces fails loudly +// rather than writing silent zeros or a file without a datatype +// message. +func TestHDF5WriteNarrowRefusals(t *testing.T) { + const stored = "the writer stores bool, int8, uint8, int16, uint16, int32, uint32, int64, float32 and float64" + dir := t.TempDir() + halves, err := core.FromFloat16s([]float64{1.5}, 1) + if err != nil { + t.Fatal(err) + } + complexes, err := core.FromComplexes([]complex128{1}, 1) + if err != nil { + t.Fatal(err) + } + for _, c := range []struct { + name string + values *core.Array + }{ + {"float16", halves}, + {"complex", complexes}, + } { + path := filepath.Join(dir, c.name+".h5") + err := SaveHDF5(path, []HDF5Dataset{{Path: "/x", Shape: c.values.Shape(), Values: c.values}}, nil, HDF5WriteOptions{}) + if err == nil { + t.Fatalf("%s: expected an unsupported-dtype refusal", c.name) + } + if want := "dtype " + c.name + " is not supported; " + stored; !strings.Contains(err.Error(), want) { + t.Fatalf("%s: error = %v, want it to carry %q", c.name, err, want) + } + if _, statErr := os.Stat(path); !os.IsNotExist(statErr) { + t.Fatalf("%s: the refused write left %q behind", c.name, path) + } + } + unknown := &hdf5OutSet{path: "/bad", class: 2, width: 4, nbytes: 4} + if err := unknown.encode(make([]byte, 4)); err == nil || !strings.Contains(err.Error(), "cannot serialise") { + t.Fatalf("encode of class 2: err = %v, want the serialisation refusal", err) + } + w := &hdf5Writer{} + if _, err := w.writeDataset(unknown); err == nil || !strings.Contains(err.Error(), "no datatype message") { + t.Fatalf("writeDataset of class 2: err = %v, want the missing-message refusal", err) + } + odd := &hdf5OutSet{path: "/odd", class: 0, width: 3, nbytes: 3} + if err := odd.encode(make([]byte, 3)); err == nil || !strings.Contains(err.Error(), "cannot serialise") { + t.Fatalf("encode of a three-byte fixed-point payload: err = %v, want the serialisation refusal", err) + } +} + +// TestHDF5WriteNarrowDeterministicBytes pins the determinism claim for +// the narrow dtypes: the same content supplied in different dataset +// orders writes byte for bit the same file, contiguous in both +// layouts, chunked through shuffle alone and chunked through shuffle +// and deflate. +func TestHDF5WriteNarrowDeterministicBytes(t *testing.T) { + sets := hdf5NarrowSets(t) + reordered := slices.Clone(sets) + slices.Reverse(reordered) + for _, opts := range []HDF5WriteOptions{ + {}, + {Latest: true}, + {Shuffle: true, ChunkBytes: 12}, + {Gzip: 4, Shuffle: true, ChunkBytes: 12}, + } { + files := make([][]byte, 2) + for i, in := range [][]HDF5Dataset{sets, reordered} { + path := filepath.Join(t.TempDir(), "det.h5") + if err := SaveHDF5(path, in, nil, opts); err != nil { + t.Fatal(err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + files[i] = raw + } + if !bytes.Equal(files[0], files[1]) { + t.Fatalf("opts %v: two writes of the same narrow content differ in %d of %d bytes", + opts, countByteDiffs(files[0], files[1]), len(files[0])) + } + } +} diff --git a/io/helpers_test.go b/io/helpers_test.go new file mode 100644 index 0000000..6ad7b17 --- /dev/null +++ b/io/helpers_test.go @@ -0,0 +1,34 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +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 +} + +// mustComplexes builds a complex array, failing the test on a bad shape. +func mustComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array { + t.Helper() + a, err := core.FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err) + } + return a +} diff --git a/io/hostile_test.go b/io/hostile_test.go new file mode 100644 index 0000000..9f9e7a1 --- /dev/null +++ b/io/hostile_test.go @@ -0,0 +1,267 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "os" + "path/filepath" + "strings" + "testing" +) + +// Hostile inputs: a data file is untrusted input, so every size field +// the header carries must be checked against the bytes actually +// present before it is used. These cases pin the contract that a +// malformed file is an error, never a panic and never an allocation +// sized by the header alone. + +// hostileNCName writes a 4-byte-padded NetCDF name. +func hostileNCName(b []byte, name string) []byte { + b = binary.BigEndian.AppendUint32(b, uint32(len(name))) + b = append(b, name...) + if pad := (4 - len(name)%4) % 4; pad != 0 { + b = append(b, make([]byte, pad)...) + } + return b +} + +// ncHostileHeader builds a minimal CDF-1 file: magic, numrecs, a +// dimension list with the given lengths, an absent attribute list, and +// one NC_DOUBLE variable over those dimensions. Nothing follows the +// header, so any declared data is truncated by construction. +func ncHostileHeader(names []string, lengths []uint32) []byte { + b := []byte{'C', 'D', 'F', 1} + b = binary.BigEndian.AppendUint32(b, 0) // numrecs + b = binary.BigEndian.AppendUint32(b, ncTagDimension) + b = binary.BigEndian.AppendUint32(b, uint32(len(lengths))) + for i, l := range lengths { + b = hostileNCName(b, names[i]) + b = binary.BigEndian.AppendUint32(b, l) + } + b = binary.BigEndian.AppendUint32(b, 0) // absent attribute list + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, ncTagVariable) + b = binary.BigEndian.AppendUint32(b, 1) + b = hostileNCName(b, "v") + b = binary.BigEndian.AppendUint32(b, uint32(len(lengths))) + for i := range lengths { + b = binary.BigEndian.AppendUint32(b, uint32(i)) + } + b = binary.BigEndian.AppendUint32(b, 0) // absent variable attributes + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, ncTypeDouble) + b = binary.BigEndian.AppendUint32(b, 0) // vsize + b = binary.BigEndian.AppendUint32(b, 0) // begin + return b +} + +// writeHostile drops the bytes into a temp file and returns the path. +func writeHostile(t *testing.T, name string, data []byte) string { + t.Helper() + path := filepath.Join(t.TempDir(), name) + if err := os.WriteFile(path, data, 0o644); err != nil { + t.Fatal(err) + } + return path +} + +// TestLoadNetCDFHostileCounts covers the three ways a declared count +// can be used to demand memory the file cannot back: a product that +// wraps to a small positive number, a product that wraps negative, and +// record counts that are simply larger than the file. +func TestLoadNetCDFHostileCounts(t *testing.T) { + cases := []struct { + name string + names []string + lengths []uint32 + }{ + { + // 2^21 * 2^20 * 2^20 = 2^61, which times the 8-byte + // element width wraps to zero in a 64-bit multiply. + name: "product wraps to zero", + names: []string{"a", "b", "c"}, + lengths: []uint32{1 << 21, 1 << 20, 1 << 20}, + }, + { + // 4294967295 squared wraps to a negative product. + name: "product wraps negative", + names: []string{"a", "b"}, + lengths: []uint32{4294967295, 4294967295}, + }, + { + name: "one dimension longer than the file", + names: []string{"a"}, + lengths: []uint32{1 << 30}, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + path := writeHostile(t, "hostile.nc", ncHostileHeader(tc.names, tc.lengths)) + if _, _, _, err := LoadNetCDF(path); err == nil { + t.Fatal("expected an error for a header whose data cannot fit the file") + } + }) + } +} + +// TestLoadNetCDFHostileRecordCounts pins the list caps: a header that +// declares millions of dimensions, attributes or variables in a file +// of a few bytes must fail before the list is allocated. +func TestLoadNetCDFHostileRecordCounts(t *testing.T) { + const declared = 2000000 + + dimCount := []byte{'C', 'D', 'F', 1} + dimCount = binary.BigEndian.AppendUint32(dimCount, 0) + dimCount = binary.BigEndian.AppendUint32(dimCount, ncTagDimension) + dimCount = binary.BigEndian.AppendUint32(dimCount, declared) + + attrCount := []byte{'C', 'D', 'F', 1} + attrCount = binary.BigEndian.AppendUint32(attrCount, 0) + attrCount = binary.BigEndian.AppendUint32(attrCount, 0) // absent dimensions + attrCount = binary.BigEndian.AppendUint32(attrCount, 0) + attrCount = binary.BigEndian.AppendUint32(attrCount, ncTagAttribute) + attrCount = binary.BigEndian.AppendUint32(attrCount, declared) + + varCount := []byte{'C', 'D', 'F', 1} + varCount = binary.BigEndian.AppendUint32(varCount, 0) + varCount = binary.BigEndian.AppendUint32(varCount, 0) // absent dimensions + varCount = binary.BigEndian.AppendUint32(varCount, 0) + varCount = binary.BigEndian.AppendUint32(varCount, 0) // absent attributes + varCount = binary.BigEndian.AppendUint32(varCount, 0) + varCount = binary.BigEndian.AppendUint32(varCount, ncTagVariable) + varCount = binary.BigEndian.AppendUint32(varCount, declared) + + cases := []struct { + name string + data []byte + }{ + {"dimensions", dimCount}, + {"attributes", attrCount}, + {"variables", varCount}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + path := writeHostile(t, "counts.nc", tc.data) + _, _, _, err := LoadNetCDF(path) + if err == nil { + t.Fatalf("a %d-byte file declaring %d %s must be refused", len(tc.data), declared, tc.name) + } + if !strings.Contains(err.Error(), "remaining bytes") { + t.Fatalf("error = %v, want the declared-count bound", err) + } + }) + } +} + +// TestLoadNetCDFTruncatedVariable pins the data-block check: a header +// that points its variable past the end of the file fails without +// reading beyond it. +func TestLoadNetCDFTruncatedVariable(t *testing.T) { + data := ncHostileHeader([]string{"a"}, []uint32{4}) + binary.BigEndian.PutUint32(data[len(data)-4:], 1<<30) // begin past EOF + path := writeHostile(t, "short.nc", data) + if _, _, _, err := LoadNetCDF(path); err == nil { + t.Fatal("expected an error for a variable whose data lies past the end of the file") + } +} + +// TestLoadFITSTableTruncatedHostile pins the table data guard: a file +// that ends inside the header block has no data at all, whether or not +// the column forms give it a row size to divide by. +func TestLoadFITSTableTruncatedHostile(t *testing.T) { + // No data block follows the header: the cards stop at 640 bytes. + header := func(form string) []byte { + var b []byte + for _, body := range []string{ + "XTENSION= 'BINTABLE'", "BITPIX = 8", "NAXIS = 2", + "NAXIS1 = 4", "NAXIS2 = 2", "TFIELDS = 1", + "TFORM1 = '" + form + strings.Repeat(" ", 8-len(form)) + "'", "END", + } { + b = append(b, card(body)...) + } + return b + } + for _, form := range []string{"X", "D"} { + t.Run(form, func(t *testing.T) { + data := header(form) + if len(data) != 640 { + t.Fatalf("the test header is %d bytes, want 640", len(data)) + } + path := writeHostile(t, "trunc.fits", data) + if _, err := LoadFITSTable(path); err == nil { + t.Fatal("expected an error for a table whose data block is missing") + } + }) + } +} + +// TestLoadNetCDFUnusedLongDimension pins the other side of the bound: a +// header may declare a dimension longer than the data any variable +// uses, because the format allows it, so the reader must accept the +// file rather than refuse a legal one. +func TestLoadNetCDFUnusedLongDimension(t *testing.T) { + // One dimension of a million elements, one variable using none of + // it: the count product is 1, so nothing is read. + var b []byte + b = append(b, 'C', 'D', 'F', 1) + b = binary.BigEndian.AppendUint32(b, 0) // numrecs + b = binary.BigEndian.AppendUint32(b, 10) // NC_DIMENSION + b = binary.BigEndian.AppendUint32(b, 1) + b = hostileNCName(b, "big") + b = binary.BigEndian.AppendUint32(b, 1<<20) + b = binary.BigEndian.AppendUint32(b, 0) // absent attributes + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, 11) // NC_VARIABLE + b = binary.BigEndian.AppendUint32(b, 1) + b = hostileNCName(b, "scalar") + b = binary.BigEndian.AppendUint32(b, 0) // rank 0 + b = binary.BigEndian.AppendUint32(b, 0) // absent attributes + b = binary.BigEndian.AppendUint32(b, 0) + b = binary.BigEndian.AppendUint32(b, 6) // NC_DOUBLE + b = binary.BigEndian.AppendUint32(b, 8) // vsize + b = binary.BigEndian.AppendUint32(b, 0) // begin + b = append(b, 0, 0, 0, 0, 0, 0, 0, 0) // one double at offset 0 + path := writeHostile(t, "unused.nc", b) + dims, vars, _, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF refused a legal file: %v", err) + } + if len(dims) != 1 || dims[0].Length != 1<<20 { + t.Fatalf("dims = %+v, want one dimension of 2^20", dims) + } + if len(vars) != 1 || vars[0].Values.Len() != 1 { + t.Fatalf("vars = %+v, want one scalar", vars) + } +} + +// TestLoadHDF5ChunkTreeCycle pins the visited set on the chunk B-tree +// walk: an inner node that lists itself among its children must be +// refused. Without the set the walk multiplies into entries^depth node +// visits before the depth guard can fire, so a few hundred crafted +// bytes never terminate. +func TestLoadHDF5ChunkTreeCycle(t *testing.T) { + const entries = 512 + rank := 1 + keySize := 8 + 8*(rank+1) + entrySize := keySize + 8 + buf := make([]byte, 24+entries*entrySize) + copy(buf, []byte("TREE")) + buf[4] = 1 // a chunk node + buf[5] = 1 // inner level: every child is recursed, not placed + binary.LittleEndian.PutUint16(buf[6:], entries) + for i := range entries { + p := 24 + i*entrySize + binary.LittleEndian.PutUint32(buf[p:], 16) // chunk size + binary.LittleEndian.PutUint32(buf[p+4:], 0) // filter mask + // The key offsets stay zero; the child address points at + // this very node. + binary.LittleEndian.PutUint64(buf[p+keySize:], 0) + } + f := &hdf5File{data: buf, offSize: 8, lenSize: 8} + place := func(uint64, []uint64, uint32, int) error { return nil } + if err := f.chunkTree(0, hdf5Layout{}, rank, map[uint64]bool{}, place, 0); err == nil || !strings.Contains(err.Error(), "revisits") { + t.Fatalf("chunkTree on a self-referencing node: %v", err) + } +} diff --git a/io/mmap.go b/io/mmap.go new file mode 100644 index 0000000..99d624a --- /dev/null +++ b/io/mmap.go @@ -0,0 +1,181 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "math" + "os" + "syscall" + "unsafe" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Memory-mapped arrays. A file of native-endian numbers can back an +// Array directly: the operating system maps the file's pages into the +// address space and the array reads them in place, so a data cube far +// larger than RAM opens instantly and only the touched pages ever +// reach memory. The mapping is read-only, which matches the arrays' +// immutability contract exactly, and the caller releases it with the +// returned function once the numbers are no longer needed. + +// MapFloats maps n float64 values of path, starting at byte offset, +// into a read-only one-dimensional array. The values must have been +// written in the machine's native byte order (binary.NativeEndian). +// The array is a live view of the mapping: release unmaps it, and any +// use of the array afterwards is a use-after-free, so release comes +// strictly last. A negative offset, a non-positive count, a missing +// file or a file that does not hold all n values is an error. +func MapFloats(path string, offset int64, n int) (a *core.Array, release func() error, err error) { + const name = "MapFloats" + length, err := countBytes(name, n, 8) + if err != nil { + return nil, nil, err + } + raw, release, err := mapRegion(path, offset, length, 8, name) + if err != nil { + return nil, nil, err + } + // SAFETY: countBytes checked that n is positive and that n*8 does not + // overflow, and mapRegion checked that the file holds that many bytes + // from an offset that is a multiple of 8 and returned a mapping whose + // base is page-aligned with exactly that intra-page offset, so the + // pointer is 8-aligned and n float64 values lie inside the mapping. + // The mapping stays alive until the caller runs release. + values := unsafe.Slice((*float64)(unsafe.Pointer(&raw[0])), n) + a, err = core.FromFloatSlice(values, n) + if err != nil { + _ = release() + return nil, nil, base.Errf("%s: %w", name, err) + } + return a, release, nil +} + +// MapFloat32s maps n float32 values of path into a read-only array, +// with the same contract as MapFloats. +func MapFloat32s(path string, offset int64, n int) (a *core.Array, release func() error, err error) { + const name = "MapFloat32s" + length, err := countBytes(name, n, 4) + if err != nil { + return nil, nil, err + } + raw, release, err := mapRegion(path, offset, length, 4, name) + if err != nil { + return nil, nil, err + } + // SAFETY: as in MapFloats, with the float32 alignment of 4. + values := unsafe.Slice((*float32)(unsafe.Pointer(&raw[0])), n) + a, err = core.FromFloat32Slice(values, n) + if err != nil { + _ = release() + return nil, nil, base.Errf("%s: %w", name, err) + } + return a, release, nil +} + +// MapInts maps n int64 values of path into a read-only array, with +// the same contract as MapFloats. +func MapInts(path string, offset int64, n int) (a *core.Array, release func() error, err error) { + const name = "MapInts" + length, err := countBytes(name, n, 8) + if err != nil { + return nil, nil, err + } + raw, release, err := mapRegion(path, offset, length, 8, name) + if err != nil { + return nil, nil, err + } + // SAFETY: as in MapFloats; an int64 has the same size and alignment. + values := unsafe.Slice((*int64)(unsafe.Pointer(&raw[0])), n) + a, err = core.IntsFromArray(values, n) + if err != nil { + _ = release() + return nil, nil, base.Errf("%s: %w", name, err) + } + return a, release, nil +} + +// countBytes converts an element count into a byte length, refusing the +// counts whose byte size does not fit an int64. Multiplying first and +// checking the product afterwards is the failure the review found: for +// n = 2^61+3 float64 values the product wraps to 24, which passes every +// file-size check, and the typed view is then built with the original n, +// which unsafe.Slice rejects with a panic instead of an error. +func countBytes(name string, n int, width int64) (int64, error) { + if n <= 0 { + return 0, base.Errf("%s: the element count must be positive, got %d", name, n) + } + if int64(n) > math.MaxInt64/width { + return 0, base.Errf("%s: %d elements of %d bytes take more bytes than a length can address", + name, n, width) + } + return int64(n) * width, nil +} + +// mapRegion maps length bytes of path at offset read-only and returns +// the byte slice viewing them plus a release that unmaps. The count, +// the offset and the element alignment are checked against the file's +// size before the mapping, so a short file or a misaligned request +// fails with an error instead of a fault or a misaligned array. +func mapRegion(path string, offset, length, align int64, name string) ([]byte, func() error, error) { + if offset < 0 { + return nil, nil, base.Errf("%s: offset must not be negative, got %d", name, offset) + } + // The caller turns the bytes into a typed slice with unsafe, which + // requires the address to satisfy the element's alignment: an + // offset that is not a multiple of the element size would hand back + // a misaligned array. + if offset%align != 0 { + return nil, nil, base.Errf("%s: offset %d is not a multiple of %d, the element size", name, offset, align) + } + if length <= 0 { + return nil, nil, base.Errf("%s: length must be positive, got %d", name, length) + } + file, err := os.Open(path) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + defer file.Close() + info, err := file.Stat() + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + if info.IsDir() { + return nil, nil, base.Errf("%s: %s is a directory", name, path) + } + if size := info.Size(); offset > size || length > size-offset { + return nil, nil, base.Errf("%s: %s holds %d bytes, %d needed from offset %d", + name, path, size, length, offset) + } + // The mapping call demands a page-aligned offset; misaligned + // requests map from the page boundary below and the returned slice + // skips the intra-page part. + page := int64(os.Getpagesize()) + pageBase := offset / page * page + raw, err := syscall.Mmap(int(file.Fd()), pageBase, int(length+offset-pageBase), syscall.PROT_READ, syscall.MAP_SHARED) + if err != nil { + return nil, nil, base.Errf("%s: mapping %s failed: %w", name, path, err) + } + view := raw[offset-pageBase:] + return view, func() error { + if err := syscall.Munmap(raw); err != nil { + return base.Errf("%s: unmapping %s failed: %w", name, path, err) + } + return nil + }, nil +} + +// SaveNativeFloats writes floats to path in the machine's native byte +// order, the format MapFloats reads back. It exists so a mapping test +// or tool can produce its own data without reaching for encoding +// details; a zero offset aligns it with MapFloats' contract. +func SaveNativeFloats(path string, values []float64) error { + buf := make([]byte, 8*len(values)) + for i, v := range values { + binary.NativeEndian.PutUint64(buf[i*8:], math.Float64bits(v)) + } + return os.WriteFile(path, buf, 0o644) +} diff --git a/io/mmap_test.go b/io/mmap_test.go new file mode 100644 index 0000000..ce8f88a --- /dev/null +++ b/io/mmap_test.go @@ -0,0 +1,165 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "math" + "os" + "path/filepath" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestMapFloatsRoundTrip writes floats natively, maps them back, and +// checks spot values far apart, which is exactly the sparse-touch +// pattern big-file mapping exists for. +func TestMapFloatsRoundTrip(t *testing.T) { + const n = 4096 + path := filepath.Join(t.TempDir(), "values.bin") + values := make([]float64, n) + for i := range n { + values[i] = math.Sin(float64(i)) * 1e6 + } + if err := SaveNativeFloats(path, values); err != nil { + t.Fatalf("SaveNativeFloats: %v", err) + } + a, release, err := MapFloats(path, 0, n) + if err != nil { + t.Fatalf("MapFloats: %v", err) + } + if a.Len() != n || a.Dtype() != core.Float { + t.Fatalf("mapped array %s of length %d", a.Dtype(), a.Len()) + } + for _, i := range []int{0, 1, 999, 2048, n - 1} { + if a.FloatAt(i) != values[i] { + t.Fatalf("value %d = %.14g, want %.14g", i, a.FloatAt(i), values[i]) + } + } + // A reshape keeps the view: strides read straight from the mapping. + square, err := core.Reshape(a, 64, 64) + if err != nil { + t.Fatalf("Reshape: %v", err) + } + if square.FloatAt(3*64+7) != values[3*64+7] { + t.Fatal("the reshaped view lost the mapping") + } + if err := release(); err != nil { + t.Fatalf("release: %v", err) + } +} + +// TestMapFloat32sAndInts covers the other two element types. +func TestMapFloat32sAndInts(t *testing.T) { + dir := t.TempDir() + const n = 100 + f32path := filepath.Join(dir, "f32.bin") + var buf []byte + var word [4]byte + want32 := make([]float32, n) + for i := range n { + want32[i] = float32(i) / 7 + binary.NativeEndian.PutUint32(word[:], math.Float32bits(want32[i])) + buf = append(buf, word[:]...) + } + if err := os.WriteFile(f32path, buf, 0o644); err != nil { + t.Fatal(err) + } + a, release, err := MapFloat32s(f32path, 0, n) + if err != nil { + t.Fatalf("MapFloat32s: %v", err) + } + if a.RawFloat32s()[42] != want32[42] { + t.Fatalf("f32[42] = %v, want %v", a.RawFloat32s()[42], want32[42]) + } + if err := release(); err != nil { + t.Fatal(err) + } + + ipath := filepath.Join(dir, "i64.bin") + var ibuf []byte + var iword [8]byte + wantInt := make([]int64, n) + for i := range n { + wantInt[i] = int64(i) * 1_000_000 + binary.NativeEndian.PutUint64(iword[:], uint64(wantInt[i])) + ibuf = append(ibuf, iword[:]...) + } + if err := os.WriteFile(ipath, ibuf, 0o644); err != nil { + t.Fatal(err) + } + b, releaseInt, err := MapInts(ipath, 0, n) + if err != nil { + t.Fatalf("MapInts: %v", err) + } + if b.RawInts()[13] != wantInt[13] { + t.Fatalf("i64[13] = %d, want %d", b.RawInts()[13], wantInt[13]) + } + if err := releaseInt(); err != nil { + t.Fatal(err) + } +} + +// TestMapFloatsOffset maps from a non-zero offset inside a file that +// carries a small header first. The header is a multiple of the +// element size: the typed view the reader builds is built with unsafe +// and must be aligned. +func TestMapFloatsOffset(t *testing.T) { + path := filepath.Join(t.TempDir(), "headered.bin") + header := []byte("TENSOR HEADER 24 BYTES!!") + if len(header)%8 != 0 { + t.Fatalf("the test header is %d bytes, want a multiple of 8", len(header)) + } + values := []float64{3.5, 1.25, -9} + var buf []byte + buf = append(buf, header...) + var word [8]byte + for _, v := range values { + binary.NativeEndian.PutUint64(word[:], math.Float64bits(v)) + buf = append(buf, word[:]...) + } + if err := os.WriteFile(path, buf, 0o644); err != nil { + t.Fatal(err) + } + a, release, err := MapFloats(path, int64(len(header)), 3) + if err != nil { + t.Fatalf("MapFloats: %v", err) + } + for i, want := range values { + if a.FloatAt(i) != want { + t.Fatalf("value %d = %.14g, want %.14g", i, a.FloatAt(i), want) + } + } + if err := release(); err != nil { + t.Fatal(err) + } +} + +// TestMapErrors pins the validation contract: short files, negative +// offsets, zero counts and directories are errors, never mappings. +func TestMapErrors(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "small.bin") + if err := SaveNativeFloats(path, []float64{1, 2, 3}); err != nil { + t.Fatal(err) + } + if _, _, err := MapFloats(path, 4, 1); err == nil { + t.Fatal("expected an error for an offset that misaligns the elements") + } + if _, _, err := MapFloats(path, 0, 4); err == nil { + t.Fatal("expected an error when the file is too short") + } + if _, _, err := MapFloats(path, -8, 2); err == nil { + t.Fatal("expected an error for a negative offset") + } + if _, _, err := MapFloats(path, 0, 0); err == nil { + t.Fatal("expected an error for a zero count") + } + if _, _, err := MapFloats(dir, 0, 1); err == nil { + t.Fatal("expected an error for a directory") + } + if _, _, err := MapFloats(filepath.Join(dir, "missing.bin"), 0, 1); err == nil { + t.Fatal("expected an error for a missing file") + } +} diff --git a/io/netcdf.go b/io/netcdf.go new file mode 100644 index 0000000..56a9ca8 --- /dev/null +++ b/io/netcdf.go @@ -0,0 +1,1077 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "maps" + "math" + "os" + "slices" + "strconv" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// NetCDF classic I/O. The Network Common Data Form is the archival +// format of climate and ocean science: a self-describing header of +// named dimensions, attributes and variables over big-endian binary +// payloads. This module speaks the classic model, CDF-1 on write and +// CDF-1 or CDF-2 on read. A record dimension (the unlimited first +// axis) interleaves its slabs across the record variables; a file +// carrying one is read, and written back, with the interleaving +// preserved. Variables land the core dtype their classic type code +// carries: NC_BYTE as int8, NC_CHAR as uint8 raw bytes, NC_SHORT as +// int16, NC_INT as int32, and NC_FLOAT and NC_DOUBLE as float64. +// CHAR carries bytes at the array level, never text; whether those +// bytes spell text is the caller's question, not the array's. + +// NetCDF type codes as the file stores them. +const ( + ncTypeByte = 1 + ncTypeChar = 2 + ncTypeShort = 3 + ncTypeInt = 4 + ncTypeFloat = 5 + ncTypeDouble = 6 +) + +// Section list tags as the file stores them. +const ( + ncTagDimension = 10 + ncTagVariable = 11 + ncTagAttribute = 12 +) + +// NetCDFDim is one named dimension of a NetCDF classic file. A length +// of zero is the record dimension (the unlimited one): only it may be +// zero, it must lead the list, and its extent is the number of records, +// which a record variable carries on its first axis. Reading a file and +// writing it back preserves the record dimension. +type NetCDFDim struct { + Name string + Length int +} + +// NetCDFVar is one named variable of a NetCDF classic file. Values +// are row-major with the slowest dimension first, exactly as the file +// stores them, and Dims names the dimensions in that same order. +type NetCDFVar struct { + Name string + Dims []string + Values *core.Array + Attrs map[string]string +} + +// SaveNetCDF writes a NetCDF classic (CDF-1) file: dimensions with +// positive lengths plus at most one record dimension (length zero, +// declared first), float64 as NC_DOUBLE, float32 as NC_FLOAT and int64 +// as NC_INT (the classic model's widest integer, so values outside the +// int32 range are an error rather than a silent truncation), and text +// attributes as NC_CHAR. A variable whose first dimension is the record +// dimension is a record variable: it must hold a whole number of +// records, every record variable must agree on that number, and the +// file interleaves one slab of each per record, padded to four bytes, +// the layout the classic model defines. Attribute keys are written in +// sorted order so the same inputs give the same bytes. +func SaveNetCDF(path string, dims []NetCDFDim, vars []NetCDFVar, attrs map[string]string) error { + const name = "SaveNetCDF" + recordDim := -1 + seen := make(map[string]bool, len(dims)) + for i, d := range dims { + if err := ncCheckName(d.Name); err != nil { + return base.Errf("%s: dimension %d: %w", name, i+1, err) + } + if seen[d.Name] { + return base.Errf("%s: dimension %q is declared twice", name, d.Name) + } + seen[d.Name] = true + if d.Length < 0 { + return base.Errf("%s: dimension %q has negative length %d", name, d.Name, d.Length) + } + if d.Length == 0 { + // The record dimension: at most one, and it leads the list, + // as the classic model requires. + if i != 0 { + return base.Errf("%s: dimension %q has length 0 but is not the first: the record dimension leads the list", + name, d.Name) + } + recordDim = i + } + } + dimIndex := make(map[string]int, len(dims)) + for i, d := range dims { + dimIndex[d.Name] = i + } + seenVar := make(map[string]bool, len(vars)) + recordVars := make([]bool, len(vars)) + perRecord := make([]int, len(vars)) + numrecs := 0 + numrecsSet := false + for i := range vars { + v := &vars[i] + if err := ncCheckName(v.Name); err != nil { + return base.Errf("%s: variable %d: %w", name, i+1, err) + } + if seenVar[v.Name] { + return base.Errf("%s: variable %q is declared twice", name, v.Name) + } + seenVar[v.Name] = true + for j, dn := range v.Dims { + di, ok := dimIndex[dn] + if !ok { + return base.Errf("%s: variable %q refers to undeclared dimension %q", name, v.Name, dn) + } + if dims[di].Length == 0 && j != 0 { + return base.Errf("%s: variable %q carries the record dimension %q in position %d: a record variable's record axis leads its dimensions", + name, v.Name, dn, j) + } + } + if _, _, err := ncExternal(v.Values.Dtype()); err != nil { + return base.Errf("%s: variable %q: %w", name, v.Name, err) + } + if v.Values.Dtype() == core.Int { + // NC_INT is a signed 32-bit value and the classic model has + // no wider integer, so a value outside that range is an + // error rather than a silent truncation. + for _, x := range v.Values.RawInts() { + if x < math.MinInt32 || x > math.MaxInt32 { + return base.Errf("%s: variable %q holds %d, outside the int32 range the classic model stores", + name, v.Name, x) + } + } + } + isRecord := recordDim >= 0 && len(v.Dims) > 0 && v.Dims[0] == dims[recordDim].Name + if isRecord { + per := 1 + for _, dn := range v.Dims[1:] { + per *= dims[dimIndex[dn]].Length + } + if per <= 0 || v.Values.Len()%per != 0 { + return base.Errf("%s: record variable %q holds %d values, not a whole number of records of %d", + name, v.Name, v.Values.Len(), per) + } + recs := v.Values.Len() / per + if numrecsSet && recs != numrecs { + return base.Errf("%s: record variable %q holds %d records, the others %d", + name, v.Name, recs, numrecs) + } + numrecs, numrecsSet = recs, true + perRecord[i], recordVars[i] = per, true + continue + } + want := 1 + for _, dn := range v.Dims { + want *= dims[dimIndex[dn]].Length + } + if v.Values.Len() != want { + return base.Errf("%s: variable %q holds %d values, its dimensions hold %d", + name, v.Name, v.Values.Len(), want) + } + } + + // The image's size follows from the plan before a byte is written: + // the data section exactly, the header from the names and the + // attribute texts with a fixed frame each. The buffer is allocated + // once from that, so the appends below fill it instead of growing + // and copying the whole file as they go. + buf := make([]byte, 0, ncImageEstimate(dims, vars, attrs, perRecord, recordVars, numrecs)) + buf = append(buf, 'C', 'D', 'F', 1) + buf = binary.BigEndian.AppendUint32(buf, uint32(numrecs)) + // dim_list + if len(dims) == 0 { + buf = append(buf, 0, 0, 0, 0, 0, 0, 0, 0) // ABSENT + } else { + buf = binary.BigEndian.AppendUint32(buf, ncTagDimension) + buf = binary.BigEndian.AppendUint32(buf, uint32(len(dims))) + for _, d := range dims { + buf = ncAppendName(buf, d.Name) + buf = binary.BigEndian.AppendUint32(buf, uint32(d.Length)) + } + } + + // gatt_list + if len(attrs) == 0 { + buf = append(buf, 0, 0, 0, 0, 0, 0, 0, 0) // ABSENT + } else { + buf = binary.BigEndian.AppendUint32(buf, ncTagAttribute) + buf = binary.BigEndian.AppendUint32(buf, uint32(len(attrs))) + for _, k := range slices.Sorted(maps.Keys(attrs)) { + buf = ncAppendName(buf, k) + buf = binary.BigEndian.AppendUint32(buf, ncTypeChar) + buf = binary.BigEndian.AppendUint32(buf, uint32(len(attrs[k]))) + buf = append(buf, attrs[k]...) + buf = ncPad4(buf) + } + } + + // var_list. The begin field is the last word of each record, so + // its position is known as the record is appended; the absolute + // data offsets follow once the whole header is in place. + buf = binary.BigEndian.AppendUint32(buf, ncTagVariable) + buf = binary.BigEndian.AppendUint32(buf, uint32(len(vars))) + beginAt := make([]int, len(vars)) + for i := range vars { + v := &vars[i] + buf = ncAppendName(buf, v.Name) + buf = binary.BigEndian.AppendUint32(buf, uint32(len(v.Dims))) + for _, dn := range v.Dims { + buf = binary.BigEndian.AppendUint32(buf, uint32(dimIndex[dn])) + } + if len(v.Attrs) == 0 { + buf = append(buf, 0, 0, 0, 0, 0, 0, 0, 0) // ABSENT + } else { + buf = binary.BigEndian.AppendUint32(buf, ncTagAttribute) + buf = binary.BigEndian.AppendUint32(buf, uint32(len(v.Attrs))) + for _, k := range slices.Sorted(maps.Keys(v.Attrs)) { + buf = ncAppendName(buf, k) + buf = binary.BigEndian.AppendUint32(buf, ncTypeChar) + buf = binary.BigEndian.AppendUint32(buf, uint32(len(v.Attrs[k]))) + buf = append(buf, v.Attrs[k]...) + buf = ncPad4(buf) + } + } + _, size, _ := ncExternal(v.Values.Dtype()) + // A record variable's vsize is one record's worth, padded to a + // four-byte boundary; a fixed variable's is its whole payload. + bytes := v.Values.Len() * size + if recordVars[i] { + bytes = int(ncPad4Size(int64(perRecord[i]) * int64(size))) + } + if uint64(bytes) > math.MaxUint32 { + return base.Errf("%s: variable %q needs %d bytes, past what a CDF-1 size word can hold", + name, v.Name, bytes) + } + buf = binary.BigEndian.AppendUint32(buf, uint32(ncExternalType(v.Values.Dtype()))) + buf = binary.BigEndian.AppendUint32(buf, uint32(bytes)) + beginAt[i] = len(buf) + buf = binary.BigEndian.AppendUint32(buf, 0) // patched below + } + + // data: the fixed variables first, in header order, each padded to + // a four-byte boundary, then the records. Each record carries the + // slab of every record variable, again in header order, each slab + // padded to four bytes, so a record variable's bytes are strided by + // the whole record size. + offset := uint32(len(buf)) + for i := range vars { + if recordVars[i] { + continue + } + v := &vars[i] + binary.BigEndian.PutUint32(buf[beginAt[i]:], offset) + start := len(buf) + payload, err := ncAppendPayload(buf, v.Values, 0, v.Values.Len()) + if err != nil { + return base.Errf("%s: variable %q: %w", name, v.Name, err) + } + buf = payload + buf = ncPad4(buf) + offset += uint32(len(buf) - start) + } + // A record variable's begin points at its own slab inside the first + // record, which is what readers (and the model's own arithmetic) + // expect: record r's slab of that variable then sits a whole record + // size further on, at begin + r*recordSize. + slabOffset := offset + for i := range vars { + if !recordVars[i] { + continue + } + binary.BigEndian.PutUint32(buf[beginAt[i]:], slabOffset) + slabOffset += uint32(ncPad4Size(int64(perRecord[i]) * int64(elemSize(vars[i].Values)))) + } + for rec := range numrecs { + for i := range vars { + if !recordVars[i] { + continue + } + start := len(buf) + lo := rec * perRecord[i] + payload, err := ncAppendPayload(buf, vars[i].Values, lo, lo+perRecord[i]) + if err != nil { + return base.Errf("%s: variable %q: %w", name, vars[i].Name, err) + } + buf = payload + for (len(buf)-start)%4 != 0 { + buf = append(buf, 0) + } + } + } + // CDF-1 addresses its offsets with a signed 32-bit word, so a file + // at or past 2 GiB cannot be addressed; refuse rather than wrap the + // begin fields. + if uint64(len(buf)) > 1<<31-1 { + return base.Errf("%s: the file would be %d bytes, past the 2 GiB a CDF-1 offset can address", name, len(buf)) + } + return os.WriteFile(path, buf, 0o644) +} + +// ncImageEstimate bounds the byte size of the file image SaveNetCDF +// builds. The data section is exact: each payload's four-byte-padded +// extent, with a record variable's slab repeated once per record. The +// header is charged a fixed frame per dimension, variable, dimension +// reference and attribute plus twice the bytes of every name and +// attribute text, which is past what the tagged lists cost. The sum is +// kept in int64 and clamped to the 2 GiB a CDF-1 offset can address, +// so no declaration can wrap it into a small or negative capacity. +func ncImageEstimate(dims []NetCDFDim, vars []NetCDFVar, attrs map[string]string, perRecord []int, recordVars []bool, numrecs int) int { + const budget = int64(1) << 31 + total := int64(2048) + add := func(n int64) { + if total += n; total > budget { + total = budget + } + } + nameBytes := func(k, v string) int64 { return int64(2*(len(k)+len(v))) + 64 } + for i := range vars { + size := int64(elemSize(vars[i].Values)) + if recordVars[i] { + add(ncPad4Size(int64(perRecord[i])*size) * int64(numrecs)) + } else { + add(ncPad4Size(int64(vars[i].Values.Len()) * size)) + } + } + for _, d := range dims { + add(nameBytes(d.Name, "")) + } + for k, v := range attrs { + add(nameBytes(k, v)) + } + for i := range vars { + add(nameBytes(vars[i].Name, "") + 4*int64(len(vars[i].Dims)) + 128) + for k, v := range vars[i].Attrs { + add(nameBytes(k, v)) + } + } + return int(total) +} + +// elemSize returns the external width of an array's dtype; ncExternal +// has already accepted it by the time the writer calls this. +func elemSize(a *core.Array) int { + _, size, err := ncExternal(a.Dtype()) + if err != nil { + return 0 + } + return size +} + +// ncAppendPayload appends arr's elements [start, end) in the external +// format of its dtype: the slot is reserved once and each element is +// stored into it, so no element pays for the append's own capacity +// check. Only the dtypes the classic model stores are accepted: a +// narrower array carries no integer payload for the int branch to +// read, so any other dtype is a named error rather than a bounds +// failure or a silent cast. SaveNetCDF's validation refuses the same +// dtypes first, so the error is a defence in depth for direct callers. +func ncAppendPayload(buf []byte, arr *core.Array, start, end int) ([]byte, error) { + switch arr.Dtype() { + case core.Float: + vals := arr.RawFloats()[start:end] + at := len(buf) + buf = append(buf, make([]byte, 8*len(vals))...) + for i, x := range vals { + binary.BigEndian.PutUint64(buf[at+i*8:], math.Float64bits(x)) + } + case core.Float32: + vals := arr.RawFloat32s()[start:end] + at := len(buf) + buf = append(buf, make([]byte, 4*len(vals))...) + for i, x := range vals { + binary.BigEndian.PutUint32(buf[at+i*4:], math.Float32bits(x)) + } + case core.Int: + vals := arr.RawInts()[start:end] + at := len(buf) + buf = append(buf, make([]byte, 4*len(vals))...) + for i, x := range vals { + binary.BigEndian.PutUint32(buf[at+i*4:], uint32(int32(x))) + } + default: + return nil, base.Errf("SaveNetCDF: cannot store dtype %s in a NetCDF classic file; convert with Astype", arr.Dtype()) + } + return buf, nil +} + +// LoadNetCDF reads a NetCDF classic file (CDF-1 and CDF-2), returning +// its dimensions, its variables and its global attributes. Every +// variable lands the core dtype its classic type code carries: +// NC_BYTE as int8 (signed in the classic model), NC_CHAR as uint8 raw +// bytes (CHAR carries bytes at the array level, never text), NC_SHORT +// as int16, NC_INT as int32, and NC_FLOAT and NC_DOUBLE as float64, +// where every classic value is exact. A type code beyond the classic +// six is refused by name: the unsigned and 64-bit codes exist only in +// formats this module does not speak. A variable that lands a narrow +// dtype here is refused by SaveNetCDF, whose writer stores float64, +// float32 and int64 arrays only; convert with Astype first. Variable +// and dimension attributes are returned on the variables themselves; +// only the global attributes come back in the map. A record dimension +// comes back with length zero, and each record variable's first axis +// carries the record count the file declares. +func LoadNetCDF(path string) ([]NetCDFDim, []NetCDFVar, map[string]string, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, nil, nil, base.Errf("LoadNetCDF: %w", err) + } + return parseNetCDF(data) +} + +// maxIndex is the largest int this platform holds: a declared +// dimension length beyond it cannot appear in the shapes the core +// builds. +const maxIndex = int64(^uint(0) >> 1) + +// ncReader walks the file bytes with bounds checks, recording the +// first violation so the parse can read straight through a structure +// and report once at the end. +type ncReader struct { + data []byte + pos int + err error +} + +func (r *ncReader) fail(format string, args ...any) { + if r.err == nil { + r.err = base.Errf("LoadNetCDF: "+format, args...) + } +} + +func (r *ncReader) take(n int) []byte { return r.take64(int64(n)) } + +// take64 is take with the arithmetic in int64, so a length read from +// the header cannot wrap on a 32-bit platform before the bounds check +// has seen it. +func (r *ncReader) take64(n int64) []byte { + if r.err != nil { + return nil + } + if n < 0 || n > int64(len(r.data)-r.pos) { + r.fail("file ends inside the header at byte %d", r.pos) + return nil + } + b := r.data[r.pos : r.pos+int(n)] + r.pos += int(n) + return b +} + +// boundedCount caps a record count declared in the header by the bytes +// that could possibly hold it: every record of a list is at least +// minBytes long, so n of them cannot fit into what is left of the +// file. A header that lies about its sizes fails here instead of +// making the reader allocate the claimed amount, which for a uint32 +// count is tens of gigabytes from a file of a few bytes. +func (r *ncReader) boundedCount(n uint32, minBytes int64, what string) int { + if r.err != nil { + return 0 + } + left := int64(len(r.data) - r.pos) + if int64(n) > left/minBytes { + r.fail("the header declares %d %s, more than the %d remaining bytes can hold", n, what, left) + return 0 + } + return int(n) +} + +func (r *ncReader) u32() uint32 { + b := r.take(4) + if b == nil { + return 0 + } + return binary.BigEndian.Uint32(b) +} + +func (r *ncReader) u64() uint64 { + b := r.take(8) + if b == nil { + return 0 + } + return binary.BigEndian.Uint64(b) +} + +// name reads a length-prefixed, 4-byte-padded name. +func (r *ncReader) name() string { + n := r.u32() + b := r.take64(int64(n) + int64(ncPadLen(int(n)))) + if b == nil { + return "" + } + return string(b[:n]) +} + +func parseNetCDF(data []byte) ([]NetCDFDim, []NetCDFVar, map[string]string, error) { + r := &ncReader{data: data} + magic := r.take(4) + if r.err != nil { + return nil, nil, nil, r.err + } + if string(magic[:3]) != "CDF" { + return nil, nil, nil, base.Errf("LoadNetCDF: not a NetCDF file, the magic is %q", magic[:3]) + } + wide := false + switch magic[3] { + case 1: + case 2: + wide = true // 64-bit offsets; numrecs stays 32-bit + default: + return nil, nil, nil, base.Errf("LoadNetCDF: unsupported NetCDF version %d", magic[3]) + } + numrecs := int64(r.u32()) + + dims := r.dimList() + attrs := r.attList() + vars := r.varList(wide, dims) + if r.err != nil { + return nil, nil, nil, r.err + } + + // The record variables' slabs interleave, so a record variable's + // bytes are not contiguous: the span check and the gather both need + // the size of a whole record, which is the sum of every record + // variable's slab padded to a four-byte boundary. Each factor is + // bounded per variable here, exactly as the read below bounds it + // before it multiplies: without that, one hostile declaration + // inflates the record size for every other record variable, ones + // already validated included, and the wrapped span that follows + // walks the reader off the end of the file. + slabs := make([]int64, len(vars)) + recordSize := int64(0) + for i, v := range vars { + if !isRecordVar(dims, v.dimIDs) { + continue + } + width, ok := ncTypeWidth(v.ncType) + if !ok { + continue // the read below reports the unknown type + } + per := int64(1) + for _, did := range v.dimIDs[1:] { + l := int64(dims[did].Length) + if l <= 0 || per > int64(len(r.data))/l { + return nil, nil, nil, base.Errf("LoadNetCDF: variable %q declares more elements than the file holds", v.name) + } + per *= l + } + slabs[i] = ncPad4Size(per * width) + if slabs[i] > math.MaxInt64-recordSize { + return nil, nil, nil, base.Errf("LoadNetCDF: the record variables declare more bytes per record than can be addressed") + } + recordSize += slabs[i] + } + + out := make([]NetCDFVar, len(vars)) + for i, v := range vars { + isRecord := isRecordVar(dims, v.dimIDs) + // The element count is a product of header fields, so it can + // wrap: a variable with more elements than the file has bytes + // cannot be backed by it, and the division keeps the product + // itself inside int64 while the check runs. The record axis + // counts once here; its extent is numrecs, applied below. + count := int64(1) + first := 0 + if isRecord { + first = 1 + } + for _, did := range v.dimIDs[first:] { + l := int64(dims[did].Length) + if l <= 0 || count > int64(len(r.data))/l { + return nil, nil, nil, base.Errf("LoadNetCDF: variable %q declares more elements than the file holds", v.name) + } + count *= l + } + perRecord := count + if isRecord { + if numrecs > int64(len(r.data)) || (perRecord > 0 && numrecs > int64(len(r.data))/perRecord) { + return nil, nil, nil, base.Errf("LoadNetCDF: variable %q declares more elements than the file holds", v.name) + } + count = perRecord * numrecs + } + names := make([]string, len(v.dimIDs)) + shape := make([]int, len(v.dimIDs)) + for j, did := range v.dimIDs { + names[j] = dims[did].Name + shape[j] = dims[did].Length + if isRecord && j == 0 { + shape[j] = int(numrecs) + } + } + // A rank-0 (scalar) variable has no shape at all, and an Array + // always carries at least one dimension, so the scalar comes + // back as a one-element vector with no dimension names. + if len(shape) == 0 { + shape = []int{1} + } + // The payload is allocated once in the dtype the type code + // lands and the decode fills it in place, so no element is + // copied twice and no integer value is rounded by a widening. + landing, ok := ncLandingDtype(v.ncType) + if !ok { + return nil, nil, nil, base.Errf("LoadNetCDF: variable %q has unknown type %d", v.name, v.ncType) + } + arr := core.New(landing, shape...) + if arr == nil { + return nil, nil, nil, base.Errf("LoadNetCDF: variable %q: shape %v cannot build an array", v.name, shape) + } + if isRecord { + r.valuesRecord(arr, v.begin, v.ncType, perRecord, numrecs, recordSize) + } else { + r.values(arr, v.begin, v.ncType, count) + } + if r.err != nil { + return nil, nil, nil, r.err + } + out[i] = NetCDFVar{Name: v.name, Dims: names, Values: arr, Attrs: v.attrs} + } + return dims, out, attrs, nil +} + +// ncVar is the parsed header record of one variable. +type ncVar struct { + name string + dimIDs []int + attrs map[string]string + ncType uint32 + begin int64 +} + +func (r *ncReader) dimList() []NetCDFDim { + if r.err != nil { + return nil + } + if r.u32() == 0 { + r.u32() // ABSENT is two zero words + return nil + } + n := r.boundedCount(r.u32(), 8, "dimensions") + dims := make([]NetCDFDim, 0, n) + seen := map[string]bool{} + for range n { + name := r.name() + length := r.u32() + if r.err != nil { + return nil + } + // Names must be unique: a variable refers to dimensions by + // index, and with a duplicated name its own name list would + // resolve differently by position and by name, so any answer + // would be a silent guess. + if seen[name] { + r.fail("dimension %q is declared twice", name) + return nil + } + seen[name] = true + if length == 0 { + // The record dimension: exactly one, and only the first. + // Its extent is the number of records, which each record + // variable carries on its leading axis. + if len(dims) != 0 { + r.fail("dimension %q has length 0 but is not the first: only one record dimension exists, and it leads the list", name) + return nil + } + dims = append(dims, NetCDFDim{Name: name, Length: 0}) + continue + } + // A declared length must fit an int on every platform the + // library targets. Whether the file actually holds the data is + // decided per variable, where the elements are read. + if int64(length) > maxIndex { + r.fail("dimension %q is declared %d long, beyond this platform's index range", name, length) + return nil + } + dims = append(dims, NetCDFDim{Name: name, Length: int(length)}) + } + return dims +} + +func (r *ncReader) attList() map[string]string { + if r.err != nil { + return nil + } + if r.u32() == 0 { + r.u32() // ABSENT + return nil + } + n := r.boundedCount(r.u32(), 12, "global attributes") + attrs := make(map[string]string, n) + seen := map[string]bool{} + for range n { + name := r.name() + ncType := r.u32() + count := r.boundedCount(r.u32(), 1, "attribute elements") + if r.err != nil { + return nil + } + if seen[name] { + r.fail("attribute %q is declared twice", name) + return nil + } + seen[name] = true + attrs[name] = r.attValue(ncType, count) + if r.err != nil { + return nil + } + } + return attrs +} + +// attValue reads one attribute payload. Text comes back verbatim; +// numeric types, which the classic model allows here, are formatted as +// decimal so callers see the value the file carries either way. +func (r *ncReader) attValue(ncType uint32, count int) string { + if r.err != nil { + return "" + } + switch ncType { + case ncTypeChar: + b := r.take(count + ncPadLen(count)) + if b == nil { + return "" + } + return string(b[:count]) + case ncTypeByte: + b := r.take(count + ncPadLen(count)) + if b == nil { + return "" + } + parts := make([]string, count) + for i := range count { + parts[i] = strconv.FormatInt(int64(int8(b[i])), 10) + } + return strings.Join(parts, ", ") + case ncTypeShort: + b := r.take(count*2 + ncPadLen(count*2)) + if b == nil { + return "" + } + parts := make([]string, count) + for i := range count { + parts[i] = strconv.FormatInt(int64(int16(binary.BigEndian.Uint16(b[i*2:]))), 10) + } + return strings.Join(parts, ", ") + case ncTypeInt: + b := r.take(count * 4) + if b == nil { + return "" + } + parts := make([]string, count) + for i := range count { + parts[i] = strconv.FormatInt(int64(int32(binary.BigEndian.Uint32(b[i*4:]))), 10) + } + return strings.Join(parts, ", ") + case ncTypeFloat: + b := r.take(count * 4) + if b == nil { + return "" + } + parts := make([]string, count) + for i := range count { + parts[i] = strconv.FormatFloat(float64(math.Float32frombits(binary.BigEndian.Uint32(b[i*4:]))), 'g', -1, 32) + } + return strings.Join(parts, ", ") + case ncTypeDouble: + b := r.take(count * 8) + if b == nil { + return "" + } + parts := make([]string, count) + for i := range count { + parts[i] = strconv.FormatFloat(math.Float64frombits(binary.BigEndian.Uint64(b[i*8:])), 'g', -1, 64) + } + return strings.Join(parts, ", ") + default: + r.fail("unknown attribute type %d", ncType) + return "" + } +} + +func (r *ncReader) varList(wide bool, dims []NetCDFDim) []ncVar { + if r.err != nil { + return nil + } + if r.u32() == 0 { + r.u32() // ABSENT + return nil + } + n := r.boundedCount(r.u32(), 28, "variables") + vars := make([]ncVar, 0, n) + seenVar := map[string]bool{} + for range n { + v := ncVar{name: r.name()} + if seenVar[v.name] { + r.fail("variable %q is declared twice", v.name) + return nil + } + seenVar[v.name] = true + rank := int(r.u32()) + if rank < 0 || int64(rank) > int64(len(r.data)-r.pos)/4 { + r.fail("variable %q declares %d dimensions, more than the file can hold", v.name, rank) + return nil + } + v.dimIDs = make([]int, rank) + for j := range rank { + id := int(r.u32()) + if id < 0 || id >= len(dims) { + r.fail("variable %q refers to dimension id %d of %d", v.name, id, len(dims)) + return nil + } + v.dimIDs[j] = id + } + v.attrs = r.attList() + v.ncType = r.u32() + r.u32() // vsize: redundant, the dimensions carry the count + if wide { + v.begin = int64(r.u64()) + } else { + v.begin = int64(r.u32()) + } + if r.err != nil { + return nil + } + vars = append(vars, v) + } + return vars +} + +// isRecordVar reports whether the variable's leading dimension is the +// record dimension, which is what makes its slabs interleave. +func isRecordVar(dims []NetCDFDim, dimIDs []int) bool { + return len(dimIDs) > 0 && dims[dimIDs[0]].Length == 0 +} + +// ncPad4Size rounds a byte count up to a four-byte boundary, the +// padding the classic format gives every variable slab. +func ncPad4Size(n int64) int64 { return (n + 3) &^ 3 } + +// ncTypeWidth returns the stored size of a classic type code. CHAR +// carries one byte per element, which the array level lands as uint8 +// raw bytes rather than as text. +func ncTypeWidth(ncType uint32) (int64, bool) { + switch ncType { + case ncTypeByte, ncTypeChar: + return 1, true + case ncTypeShort: + return 2, true + case ncTypeInt, ncTypeFloat: + return 4, true + case ncTypeDouble: + return 8, true + } + return 0, false +} + +// ncLandingDtype maps a classic type code onto the core dtype its +// values land in. NC_BYTE is signed in the classic model; NC_CHAR +// lands uint8 raw bytes, which say nothing about text at the array +// level. NC_FLOAT and NC_DOUBLE keep their float64 landing, where +// every classic value is exact. Codes 7 to 11 exist only in the +// unsigned and 64-bit extensions of the classic formats this module +// speaks, and stay refused by name: no supported writer emits them. +func ncLandingDtype(ncType uint32) (core.Dtype, bool) { + switch ncType { + case ncTypeByte: + return core.Int8, true + case ncTypeChar: + return core.Uint8, true + case ncTypeShort: + return core.Int16, true + case ncTypeInt: + return core.Int32, true + case ncTypeFloat, ncTypeDouble: + return core.Float, true + } + return 0, false +} + +// valuesRecord reads a record variable into arr's payload: each of the +// numrecs records carries its own slab of perRecord elements at the +// variable's begin offset plus a whole multiple of recordSize, because +// the other record variables' slabs sit between them. The padding +// inside a slab is skipped, never decoded. +func (r *ncReader) valuesRecord(arr *core.Array, begin int64, ncType uint32, perRecord, numrecs, recordSize int64) { + if r.err != nil { + return + } + width, ok := ncTypeWidth(ncType) + if !ok { + r.fail("variable has unknown type %d", ncType) + return + } + if perRecord < 0 || numrecs < 0 || recordSize < perRecord*width { + r.fail("record variable declares %d elements per record over %d records", perRecord, numrecs) + return + } + // Every slab must lie inside the file: with records present the last + // record's slab ends at begin + (numrecs-1)*recordSize + + // perRecord*width, and with none there is nothing to read at all. + // The record count is divided out before it is multiplied: a hostile + // record size makes the product wrap, and a wrapped span passes a + // check that compares the product itself. + failSpan := func() { + r.fail("variable data at offset %d (%d records of %d bytes) runs past the end of the %d-byte file", + begin, numrecs, recordSize, len(r.data)) + } + span := int64(0) + if numrecs > 0 { + span = perRecord * width + fits := span <= int64(len(r.data)) + if numrecs > 1 { + fits = fits && recordSize <= (int64(len(r.data))-span)/(numrecs-1) + } + if !fits { + failSpan() + return + } + span += (numrecs - 1) * recordSize + } + if begin < 0 || begin > int64(len(r.data))-span { + failSpan() + return + } + for rec := range numrecs { + slab := begin + rec*recordSize + // The window, not the tail: decodeCells fills [lo, hi) cells. + lo := int(rec * perRecord) + if err := decodeCells(arr, r.data[slab:slab+perRecord*width], ncType, lo, lo+int(perRecord)); err != nil { + r.fail("%v", err) + return + } + } +} + +// values reads count elements of the given type from the absolute file +// offset into arr's payload, in the dtype the type code lands: every +// classic integer type fits its own narrow dtype exactly, and the +// float codes keep their float64 landing. +func (r *ncReader) values(arr *core.Array, begin int64, ncType uint32, count int64) { + if r.err != nil { + return + } + width, ok := ncTypeWidth(ncType) + if !ok { + r.fail("variable has unknown type %d", ncType) + return + } + // The declared count is a product of header fields, so it can be + // negative (a wrapped multiply) or larger than the address space. + // The comparison is done in int64 against the bytes actually left + // in the file, which rejects both without ever forming a size the + // allocator would have to refuse. + span := count * width + if count < 0 || span < 0 || begin < 0 || begin > int64(len(r.data))-span { + r.fail("variable data at offset %d (%d %d-byte elements) runs past the end of the %d-byte file", + begin, count, width, len(r.data)) + return + } + if err := decodeCells(arr, r.data[begin:begin+span], ncType, 0, int(count)); err != nil { + r.fail("%v", err) + } +} + +// decodeCells decodes one contiguous run of stored cells into the +// payload window [lo, hi) of arr, which the caller allocated in the +// dtype ncLandingDtype maps ncType onto; the raw run and the window +// must agree in element count for the type. A type code without a case +// is refused: the classic formats carry only the codes +// ncLandingDtype maps, and an unhandled one left the payload silently +// zero once, which is the worst answer a reader can give. +func decodeCells(arr *core.Array, raw []byte, ncType uint32, lo, hi int) error { + switch ncType { + case ncTypeByte: + dst := arr.RawInt8s()[lo:hi] + for i := range dst { + dst[i] = int8(raw[i]) + } + case ncTypeChar: + copy(arr.RawUint8s()[lo:hi], raw[:hi-lo]) + case ncTypeShort: + dst := arr.RawInt16s()[lo:hi] + for i := range dst { + dst[i] = int16(binary.BigEndian.Uint16(raw[i*2:])) + } + case ncTypeInt: + dst := arr.RawInt32s()[lo:hi] + for i := range dst { + dst[i] = int32(binary.BigEndian.Uint32(raw[i*4:])) + } + case ncTypeFloat: + dst := arr.RawFloats()[lo:hi] + for i := range dst { + dst[i] = float64(math.Float32frombits(binary.BigEndian.Uint32(raw[i*4:]))) + } + case ncTypeDouble: + dst := arr.RawFloats()[lo:hi] + for i := range dst { + dst[i] = math.Float64frombits(binary.BigEndian.Uint64(raw[i*8:])) + } + default: + return base.Errf("variable has unknown type %d", ncType) + } + return nil +} + +// ncExternal maps a dtype onto the NetCDF type code and the external +// size the file stores it in. +func ncExternal(dt core.Dtype) (code uint32, size int, err error) { + switch dt { + case core.Float: + return ncTypeDouble, 8, nil + case core.Float32: + return ncTypeFloat, 4, nil + case core.Int: + return ncTypeInt, 4, nil + default: + return 0, 0, base.Errf("the classic model stores float64, float32 and int arrays, got dtype %s", dt) + } +} + +func ncExternalType(dt core.Dtype) uint32 { + code, _, _ := ncExternal(dt) + return code +} + +// ncCheckName enforces the traditional NetCDF name grammar: the first +// character alphanumeric or '_', the rest alphanumeric or one of +// '_.@+-', which every classic reader accepts without escaping. +func ncCheckName(s string) error { + if s == "" { + return base.Errf("names must not be empty") + } + for i := range s { + c := s[i] + ok := c == '_' || c >= '0' && c <= '9' || c >= 'A' && c <= 'Z' || c >= 'a' && c <= 'z' + if i > 0 { + ok = ok || c == '.' || c == '@' || c == '+' || c == '-' + } + if !ok { + return base.Errf("name %q contains the character %q, which the traditional name grammar forbids", s, c) + } + } + return nil +} + +// ncAppendName writes a length-prefixed name padded to 4 bytes. +func ncAppendName(buf []byte, s string) []byte { + buf = binary.BigEndian.AppendUint32(buf, uint32(len(s))) + buf = append(buf, s...) + for range ncPadLen(len(s)) { + buf = append(buf, 0) + } + return buf +} + +// ncPadLen is the number of bytes needed to reach the next 4-byte +// boundary. +func ncPadLen(n int) int { return (4 - n%4) % 4 } + +// ncPad4 appends null bytes up to the next 4-byte boundary. The +// header pads with zeros; the payload types this module writes are all +// multiples of four bytes wide, so no fill value is ever needed. +func ncPad4(buf []byte) []byte { + for range ncPadLen(len(buf)) { + buf = append(buf, 0) + } + return buf +} diff --git a/io/netcdf_payload_test.go b/io/netcdf_payload_test.go new file mode 100644 index 0000000..a1572d2 --- /dev/null +++ b/io/netcdf_payload_test.go @@ -0,0 +1,102 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/hex" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestNetCDFAppendPayloadRefusesNarrowDtypes pins the defensive check +// in ncAppendPayload: only the dtypes the classic model stores encode, +// and a narrower array is a named error rather than a read past a +// payload it does not carry. SaveNetCDF's own validation refuses the +// same dtypes before the append runs, so the error is unreachable +// through the public API; the pin holds the direct contract. +func TestNetCDFAppendPayloadRefusesNarrowDtypes(t *testing.T) { + f16, err := core.FromFloat16s([]float64{1, 0}, 2) + if err != nil { + t.Fatalf("FromFloat16s: %v", err) + } + i8, err := core.FromInt8s([]int8{-1, 1}, 2) + if err != nil { + t.Fatalf("FromInt8s: %v", err) + } + b, err := core.FromBools([]bool{true, false}, 2) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + for _, tc := range []struct { + name string + arr *core.Array + }{ + {"float16", f16}, + {"int8", i8}, + {"bool", b}, + } { + t.Run(tc.name, func(t *testing.T) { + buf, err := ncAppendPayload(nil, tc.arr, 0, tc.arr.Len()) + if err == nil { + t.Fatalf("ncAppendPayload accepted a %s array", tc.arr.Dtype()) + } + if buf != nil { + t.Fatalf("ncAppendPayload returned %d bytes alongside the error", len(buf)) + } + want := "cannot store dtype " + tc.arr.Dtype().String() + if !strings.Contains(err.Error(), want) { + t.Fatalf("error %q does not name the dtype, want %q", err, want) + } + }) + } +} + +// TestNetCDFAppendPayloadBytePins pins the external encodings byte for +// byte: float64 as NC_DOUBLE, float32 as NC_FLOAT and int64 as NC_INT, +// all big-endian, exactly as the classic model stores them. +func TestNetCDFAppendPayloadBytePins(t *testing.T) { + f64, err := core.FromFloats([]float64{1.5, -2.25}, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + f32, err := core.FromFloat32s([]float32{1.5, -2.25}, 2) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + ints, err := core.FromInts([]int64{1, -2}, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + for _, tc := range []struct { + name string + arr *core.Array + want string + }{ + {"float64", f64, "3ff8000000000000c002000000000000"}, + {"float32", f32, "3fc00000c0100000"}, + {"int64", ints, "00000001fffffffe"}, + } { + t.Run(tc.name, func(t *testing.T) { + buf, err := ncAppendPayload(nil, tc.arr, 0, tc.arr.Len()) + if err != nil { + t.Fatalf("ncAppendPayload: %v", err) + } + if got := hex.EncodeToString(buf); got != tc.want { + t.Fatalf("payload = %s, want %s", got, tc.want) + } + }) + } + // The window form encodes exactly the elements [start, end): a + // sliced append carries the same bytes the whole array of that + // window would. + buf, err := ncAppendPayload(nil, f64, 1, 2) + if err != nil { + t.Fatalf("ncAppendPayload: %v", err) + } + if got, want := hex.EncodeToString(buf), "c002000000000000"; got != want { + t.Fatalf("window payload = %s, want %s", got, want) + } +} diff --git a/io/netcdf_test.go b/io/netcdf_test.go new file mode 100644 index 0000000..72f4398 --- /dev/null +++ b/io/netcdf_test.go @@ -0,0 +1,634 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/hex" + "math" + "os" + "path/filepath" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +func ncTempPath(t *testing.T) string { + t.Helper() + return filepath.Join(t.TempDir(), "data.nc") +} + +// TestNetCDFRoundTrip moves two variables of different dtypes through +// the file format and back, values, shapes, dimensions and attributes +// included. +func TestNetCDFRoundTrip(t *testing.T) { + dims := []NetCDFDim{{Name: "lat", Length: 3}, {Name: "lon", Length: 4}} + temp := mustFloats(t, []float64{ + 1.5, -2.25, 3.125, 4, + -5.5, 6.75, -7.875, 8, + 9.25, -10.5, 11.125, -12, + }, 3, 4) + // A scalar carries one value and no dimensions. + scalar, err := core.FromFloats([]float64{42}, 1) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + vars := []NetCDFVar{ + { + Name: "temp", + Dims: []string{"lat", "lon"}, + Values: temp, + Attrs: map[string]string{"units": "degC", "long_name": "sea surface"}, + }, + { + Name: "quality", + Values: scalar, + }, + } + + path := ncTempPath(t) + attrs := map[string]string{"title": "round trip", "source": "tensor"} + if err := SaveNetCDF(path, dims, vars, attrs); err != nil { + t.Fatalf("SaveNetCDF: %v", err) + } + gotDims, gotVars, gotAttrs, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if len(gotDims) != 2 || gotDims[0] != dims[0] || gotDims[1] != dims[1] { + t.Fatalf("dims = %v, want %v", gotDims, dims) + } + for k, want := range attrs { + if gotAttrs[k] != want { + t.Fatalf("global attr %q = %q, want %q", k, gotAttrs[k], want) + } + } + if len(gotVars) != 2 { + t.Fatalf("got %d variables, want 2", len(gotVars)) + } + tv := gotVars[0] + if tv.Name != "temp" || len(tv.Dims) != 2 || tv.Dims[0] != "lat" || tv.Dims[1] != "lon" { + t.Fatalf("temp dims = %v", tv.Dims) + } + if tv.Values.NDim() != 2 || tv.Values.Shape()[0] != 3 || tv.Values.Shape()[1] != 4 { + t.Fatalf("temp shape = %v", tv.Values.Shape()) + } + for i := range temp.Len() { + if tv.Values.FloatAt(i) != temp.FloatAt(i) { + t.Fatalf("temp[%d] = %g, want %g", i, tv.Values.FloatAt(i), temp.FloatAt(i)) + } + } + for k, want := range map[string]string{"units": "degC", "long_name": "sea surface"} { + if tv.Attrs[k] != want { + t.Fatalf("temp attr %q = %q, want %q", k, tv.Attrs[k], want) + } + } + qv := gotVars[1] + if qv.Name != "quality" || len(qv.Dims) != 0 || qv.Values.FloatAt(0) != 42 { + t.Fatalf("quality = %v %v %g", qv.Name, qv.Dims, qv.Values.FloatAt(0)) + } +} + +// TestNetCDFIntAndFloat32Ranges pins the integer and single-precision +// external types, extremes included: NC_INT is a signed 32-bit value +// that lands int32 on read, and the float32 round trip is exact +// through the float64 landing NC_FLOAT keeps. +func TestNetCDFIntAndFloat32Ranges(t *testing.T) { + ints, err := core.FromInts([]int64{-2147483648, -1, 0, 1, 2147483647}, 5) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + f32, err := core.FromFloat32s([]float32{1.5, -2.25, 0, math.MaxFloat32, math.SmallestNonzeroFloat32}, 5) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + path := ncTempPath(t) + vars := []NetCDFVar{ + {Name: "i", Dims: []string{"n"}, Values: ints}, + {Name: "f", Dims: []string{"n"}, Values: f32}, + } + if err := SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 5}}, vars, nil); err != nil { + t.Fatalf("SaveNetCDF: %v", err) + } + _, got, _, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if dt := got[0].Values.Dtype(); dt != core.Int32 { + t.Fatalf("int dtype = %s, want int32", dt) + } + wantInts := []int32{-2147483648, -1, 0, 1, 2147483647} + for i, want := range wantInts { + if got[0].Values.RawInt32s()[i] != want { + t.Fatalf("int[%d] = %d, want %d", i, got[0].Values.RawInt32s()[i], want) + } + } + if dt := got[1].Values.Dtype(); dt != core.Float { + t.Fatalf("float32 dtype = %s, want the float64 landing NC_FLOAT keeps", dt) + } + for i := range 5 { + if got[1].Values.FloatAt(i) != float64(f32.RawFloat32s()[i]) { + t.Fatalf("f32[%d] = %g, want %g", i, got[1].Values.FloatAt(i), f32.RawFloat32s()[i]) + } + } +} + +// TestNetCDFParsesKnownBytes pins the parser against the specification +// rather than against our own writer: this is the canonical tiny.nc +// from the NetCDF file format documentation, one SHORT variable "vx" +// of length 5 with values 3, 1, 4, 1, 5 and one fill-padded tail. +func TestNetCDFParsesKnownBytes(t *testing.T) { + const dump = "" + + "43444601" + // magic CDF\x01 + "00000000" + // numrecs = 0 + "0000000a" + // NC_DIMENSION + "00000001" + // one dimension + "0000000364696d00" + // name "dim" + pad + "00000005" + // length 5 + "0000000000000000" + // gatt_list ABSENT + "0000000b" + // NC_VARIABLE + "00000001" + // one variable + "0000000276780000" + // name "vx" + pad + "00000001" + // rank 1 + "00000000" + // dimid 0 + "0000000000000000" + // vatt_list ABSENT + "00000003" + // NC_SHORT + "0000000c" + // vsize 12 (5 shorts + 1 fill pad) + "00000050" + // begin 80 + "00030001000400010005" + // 3, 1, 4, 1, 5 + "8001" // fill pad + raw, err := hex.DecodeString(dump) + if err != nil { + t.Fatalf("hex: %v", err) + } + if len(raw) != 92 { + t.Fatalf("fixture is %d bytes, the specification says 92", len(raw)) + } + path := ncTempPath(t) + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + dims, vars, attrs, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if len(dims) != 1 || dims[0].Name != "dim" || dims[0].Length != 5 { + t.Fatalf("dims = %v", dims) + } + if len(vars) != 1 || vars[0].Name != "vx" || len(vars[0].Dims) != 1 || vars[0].Dims[0] != "dim" { + t.Fatalf("vars = %v", vars) + } + // NC_SHORT lands int16, the width the classic file stores. + if dt := vars[0].Values.Dtype(); dt != core.Int16 { + t.Fatalf("vx dtype = %s, want int16", dt) + } + if got := vars[0].Values.RawInt16s()[:5]; !slices.Equal(got, []int16{3, 1, 4, 1, 5}) { + t.Fatalf("vx = %v, want [3 1 4 1 5]", got) + } + want := []float64{3, 1, 4, 1, 5} + for i, w := range want { + if got := vars[0].Values.FloatAt(i); got != w { + t.Fatalf("vx[%d] = %g, want %g", i, got, w) + } + } + if len(attrs) != 0 { + t.Fatalf("attrs = %v, want none", attrs) + } +} + +// TestNetCDFErrors pins the refusal paths: a record dimension, an +// unknown version, a truncated file, bad names and a shape mismatch. +func TestNetCDFErrors(t *testing.T) { + path := ncTempPath(t) + a := mustFloats(t, []float64{1, 2, 3}, 3) + + // The record dimension is supported, but only where the classic + // model allows it: leading the dimension list, leading each record + // variable's dimensions, and holding a whole number of records. + err := SaveNetCDF(path, []NetCDFDim{{Name: "x", Length: 2}, {Name: "t", Length: 0}}, + []NetCDFVar{{Name: "v", Dims: []string{"x"}, Values: a}}, nil) + if err == nil || !strings.Contains(err.Error(), "leads the list") { + t.Fatalf("record dimension not first: %v", err) + } + err = SaveNetCDF(path, []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 3}}, + []NetCDFVar{{Name: "v", Dims: []string{"x", "t"}, Values: a}}, nil) + if err == nil || !strings.Contains(err.Error(), "record axis leads") { + t.Fatalf("record axis not leading: %v", err) + } + err = SaveNetCDF(path, []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 2}}, + []NetCDFVar{{Name: "v", Dims: []string{"t", "x"}, Values: a}}, nil) + if err == nil || !strings.Contains(err.Error(), "whole number of records") { + t.Fatalf("partial record: %v", err) + } + two, terr := core.FromFloats([]float64{1, 2}, 2) + if terr != nil { + t.Fatalf("FromFloats: %v", terr) + } + err = SaveNetCDF(path, []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 1}}, + []NetCDFVar{ + {Name: "u", Dims: []string{"t"}, Values: a}, + {Name: "w", Dims: []string{"t", "x"}, Values: two}, + }, nil) + if err == nil || !strings.Contains(err.Error(), "records, the others") { + t.Fatalf("record count disagreement: %v", err) + } + + // An int64 value outside int32 is refused, not truncated. + big, ferr := core.FromInts([]int64{1 << 40}, 1) + if ferr != nil { + t.Fatalf("FromInts: %v", ferr) + } + err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 1}}, + []NetCDFVar{{Name: "v", Dims: []string{"n"}, Values: big}}, nil) + if err == nil || !strings.Contains(err.Error(), "int32 range") { + t.Fatalf("int64 out of range: %v", err) + } + + // A landed narrow variable is refused by the writer, whose classic + // model stores float64, float32 and int64 arrays only. + i8, ierr := core.FromInt8s([]int8{1, 2, 3}, 3) + if ierr != nil { + t.Fatalf("FromInt8s: %v", ierr) + } + err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 3}}, + []NetCDFVar{{Name: "v", Dims: []string{"n"}, Values: i8}}, nil) + if err == nil || !strings.Contains(err.Error(), "SaveNetCDF") || + !strings.Contains(err.Error(), "int8") || + !strings.Contains(err.Error(), "the classic model stores float64, float32 and int arrays") { + t.Fatalf("int8 variable: %v", err) + } + + // The value count must match the dimensions. + err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 4}}, + []NetCDFVar{{Name: "v", Dims: []string{"n"}, Values: a}}, nil) + if err == nil || !strings.Contains(err.Error(), "its dimensions hold 4") { + t.Fatalf("shape mismatch: %v", err) + } + + // Names follow the traditional grammar. + err = SaveNetCDF(path, []NetCDFDim{{Name: "bad name", Length: 1}}, + []NetCDFVar{{Name: "v", Dims: []string{"bad name"}, Values: mustFloats(t, []float64{1}, 1)}}, nil) + if err == nil || !strings.Contains(err.Error(), "forbids") { + t.Fatalf("bad name: %v", err) + } + + // Duplicate variable names are refused. + err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 1}}, + []NetCDFVar{ + {Name: "v", Dims: []string{"n"}, Values: mustFloats(t, []float64{1}, 1)}, + {Name: "v", Dims: []string{"n"}, Values: mustFloats(t, []float64{2}, 1)}, + }, nil) + if err == nil || !strings.Contains(err.Error(), "declared twice") { + t.Fatalf("duplicate variable: %v", err) + } + + // Not a NetCDF file at all. + if err := os.WriteFile(path, []byte("not a netcdf file at all, sorry"), 0o644); err != nil { + t.Fatalf("write: %v", err) + } + if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "not a NetCDF file") { + t.Fatalf("bad magic: %v", err) + } + + // A version this module does not speak. + if err := os.WriteFile(path, []byte{'C', 'D', 'F', 9}, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "unsupported NetCDF version") { + t.Fatalf("bad version: %v", err) + } + + // CDF-5, the 64-bit-count extension: refused by the same version + // gate rather than parsed as classic. + if err := os.WriteFile(path, []byte{'C', 'D', 'F', 5}, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "unsupported NetCDF version") { + t.Fatalf("CDF-5 version: %v", err) + } + + // A truncated header. + if err := os.WriteFile(path, []byte{'C', 'D', 'F', 1, 0, 0}, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "ends inside the header") { + t.Fatalf("truncated: %v", err) + } +} + +// TestNetCDFRecordDimensionRead pins the record dimension at load +// time: a length of 0 in the leading dimension is the unlimited +// dimension, so the header parses and the dimension comes back with +// length 0 (its extent lives on each record variable's first axis). +func TestNetCDFRecordDimensionRead(t *testing.T) { + const dump = "" + + "43444601" + // magic + "00000001" + // numrecs = 1 + "0000000a00000001" + // NC_DIMENSION, one dimension + "0000000174000000" + // name "t" + pad + "00000000" + // length 0: the record dimension + "0000000000000000" + // gatt ABSENT + "0000000b00000000" // NC_VARIABLE, zero variables + raw, err := hex.DecodeString(dump) + if err != nil { + t.Fatalf("hex: %v", err) + } + path := ncTempPath(t) + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + dims, vars, _, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("record dimension: %v", err) + } + if len(dims) != 1 || dims[0].Name != "t" || dims[0].Length != 0 { + t.Fatalf("dims = %+v, want one record dimension named t with length 0", dims) + } + if len(vars) != 0 { + t.Fatalf("vars = %+v, want none", vars) + } + // A second zero-length dimension is refused: only one unlimited + // dimension exists, and it leads the list. + const twoRecords = "" + + "43444601" + // magic + "00000001" + // numrecs = 1 + "0000000a00000002" + // NC_DIMENSION, two dimensions + "0000000174000000" + // name "t" + pad + "00000000" + // length 0: the record dimension + "0000000178000000" + // name "x" + pad + "00000000" + // length 0 again: refused + "0000000000000000" + // gatt ABSENT + "0000000b00000000" // NC_VARIABLE, zero variables + raw, err = hex.DecodeString(twoRecords) + if err != nil { + t.Fatalf("hex: %v", err) + } + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + if _, _, _, err := LoadNetCDF(path); err == nil { + t.Fatal("expected an error for a second zero-length dimension") + } +} + +// TestNetCDFClassicNativeLandings pins the classic type codes onto the +// core dtypes they land, against a hand-built CDF-1 file rather than +// the package's own writer: NC_BYTE as int8, NC_SHORT as int16, +// NC_INT as int32, NC_CHAR as uint8 raw bytes, and NC_FLOAT and +// NC_DOUBLE keeping their float64 landing. The values carry the +// extremes of every width, so a widened or unsigned reading of any of +// them fails the pin. +func TestNetCDFClassicNativeLandings(t *testing.T) { + const dump = "" + + "43444601" + // magic CDF\x01 + "00000000" + // numrecs = 0 + "0000000a00000001" + // NC_DIMENSION, one dimension + "000000016e000000" + // name "n" + pad + "00000004" + // length 4 + "0000000000000000" + // gatt_list ABSENT + "0000000b00000006" + // NC_VARIABLE, six variables + "000000016200000000000001000000000000000000000000" + // "b": rank 1, dim 0, no attrs + "00000001" + // NC_BYTE + "00000004" + // vsize 4 + "00000104" + // begin 260 + "000000017300000000000001000000000000000000000000" + // "s" + "00000003" + // NC_SHORT + "00000008" + + "00000108" + // begin 264 + "000000016900000000000001000000000000000000000000" + // "i" + "00000004" + // NC_INT + "00000010" + + "00000110" + // begin 272 + "000000016300000000000001000000000000000000000000" + // "c" + "00000002" + // NC_CHAR + "00000004" + + "00000120" + // begin 288 + "000000016600000000000001000000000000000000000000" + // "f" + "00000005" + // NC_FLOAT + "00000010" + + "00000124" + // begin 292 + "000000016400000000000001000000000000000000000000" + // "d" + "00000006" + // NC_DOUBLE + "00000020" + + "00000134" + // begin 308 + "80ff007f" + // b: -128, -1, 0, 127 + "8000ffff00017fff" + // s: -32768, -1, 1, 32767 + "80000000ffffffff000000017fffffff" + // i: int32 extremes + "0041c8ff" + // c: raw bytes 0, 65, 200, 255 + "3fc00000c02000000000000040500000" + // f: 1.5, -2.5, 0, 3.25 + "3ff8000000000000c004000000000000" + // d: 1.5, -2.5, + "00000000000000004010000000000000" // 0, 4 + raw, err := hex.DecodeString(dump) + if err != nil { + t.Fatalf("hex: %v", err) + } + path := ncTempPath(t) + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + dims, vars, attrs, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if len(dims) != 1 || dims[0].Name != "n" || dims[0].Length != 4 { + t.Fatalf("dims = %v", dims) + } + if len(attrs) != 0 { + t.Fatalf("attrs = %v, want none", attrs) + } + if len(vars) != 6 { + t.Fatalf("vars = %d, want 6", len(vars)) + } + byName := map[string]NetCDFVar{} + for _, v := range vars { + byName[v.Name] = v + if s := v.Values.Shape(); len(s) != 1 || s[0] != 4 { + t.Fatalf("%s shape = %v, want [4]", v.Name, s) + } + } + if dt := byName["b"].Values.Dtype(); dt != core.Int8 { + t.Fatalf("NC_BYTE dtype = %s, want int8", dt) + } + if got, want := byName["b"].Values.RawInt8s()[:4], []int8{-128, -1, 0, 127}; !slices.Equal(got, want) { + t.Fatalf("NC_BYTE values = %v, want %v", got, want) + } + if dt := byName["s"].Values.Dtype(); dt != core.Int16 { + t.Fatalf("NC_SHORT dtype = %s, want int16", dt) + } + if got, want := byName["s"].Values.RawInt16s()[:4], []int16{-32768, -1, 1, 32767}; !slices.Equal(got, want) { + t.Fatalf("NC_SHORT values = %v, want %v", got, want) + } + if dt := byName["i"].Values.Dtype(); dt != core.Int32 { + t.Fatalf("NC_INT dtype = %s, want int32", dt) + } + if got, want := byName["i"].Values.RawInt32s()[:4], []int32{-2147483648, -1, 1, 2147483647}; !slices.Equal(got, want) { + t.Fatalf("NC_INT values = %v, want %v", got, want) + } + // CHAR carries bytes, not text: the array level lands uint8 and + // what the bytes spell is the caller's question. + if dt := byName["c"].Values.Dtype(); dt != core.Uint8 { + t.Fatalf("NC_CHAR dtype = %s, want uint8", dt) + } + if got, want := byName["c"].Values.RawUint8s()[:4], []uint8{0, 65, 200, 255}; !slices.Equal(got, want) { + t.Fatalf("NC_CHAR values = %v, want %v", got, want) + } + if dt := byName["f"].Values.Dtype(); dt != core.Float { + t.Fatalf("NC_FLOAT dtype = %s, want float", dt) + } + for i, want := range []float64{1.5, -2.5, 0, 3.25} { + if got := byName["f"].Values.FloatAt(i); got != want { + t.Fatalf("f[%d] = %v, want %v", i, got, want) + } + } + if dt := byName["d"].Values.Dtype(); dt != core.Float { + t.Fatalf("NC_DOUBLE dtype = %s, want float", dt) + } + for i, want := range []float64{1.5, -2.5, 0, 4} { + if got := byName["d"].Values.FloatAt(i); got != want { + t.Fatalf("d[%d] = %v, want %v", i, got, want) + } + } +} + +// TestNetCDFRecordNarrowLandings pins the record path on a hand-built +// CDF-1 file: an NC_BYTE record variable's slabs land int8 in place +// and an NC_SHORT record variable's slabs land int16, including the +// slab padding the classic layout writes between records, and a fixed +// NC_SHORT variable lands int16 beside them. +func TestNetCDFRecordNarrowLandings(t *testing.T) { + const dump = "" + + "43444601" + // magic + "00000002" + // numrecs = 2 + "0000000a00000002" + // NC_DIMENSION, two dimensions + "0000000372656300" + // name "rec" + pad + "00000000" + // length 0: the record dimension + "0000000178000000" + // name "x" + pad + "00000003" + // length 3 + "0000000000000000" + // gatt_list ABSENT + "0000000b00000003" + // NC_VARIABLE, three variables + "0000000272620000" + // name "rb" + pad + "00000001" + // rank 1 + "00000000" + // dimid 0: the record dimension + "0000000000000000" + // vatt ABSENT + "00000001" + // NC_BYTE + "00000004" + // vsize 4: one byte padded to four + "000000ac" + // begin 172: its slab inside the first record + "0000000273730000" + // name "ss" + pad + "00000001" + // rank 1 + "00000001" + // dimid 1: x, fixed + "0000000000000000" + // vatt ABSENT + "00000003" + // NC_SHORT + "00000008" + // vsize 8: six bytes padded to eight + "000000a4" + // begin 164 + "0000000272730000" + // name "rs" + pad + "00000001" + // rank 1 + "00000000" + // dimid 0: the record dimension + "0000000000000000" + // vatt ABSENT + "00000003" + // NC_SHORT + "00000004" + // vsize 4: one short padded to four + "000000b0" + // begin 176: its slab inside the first record + "0005fffa0007" + // ss: 5, -6, 7 + "0000" + // slab padding + "07000000" + // rb record 0: 7 + "12340000" + // rs record 0: 0x1234 + "c8000000" + // rb record 1: 200 as int8 is -56 + "fff00000" // rs record 1: -16 + raw, err := hex.DecodeString(dump) + if err != nil { + t.Fatalf("hex: %v", err) + } + path := ncTempPath(t) + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + dims, vars, _, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if len(dims) != 2 || dims[0].Name != "rec" || dims[0].Length != 0 || dims[1].Length != 3 { + t.Fatalf("dims = %v", dims) + } + byName := map[string]NetCDFVar{} + for _, v := range vars { + byName[v.Name] = v + } + rb := byName["rb"].Values + if rb.Dtype() != core.Int8 { + t.Fatalf("rb dtype = %s, want int8", rb.Dtype()) + } + if s := rb.Shape(); len(s) != 1 || s[0] != 2 { + t.Fatalf("rb shape = %v, want [2]: the record axis carries the record count", s) + } + if got, want := rb.RawInt8s()[:2], []int8{7, -56}; !slices.Equal(got, want) { + t.Fatalf("rb values = %v, want %v", got, want) + } + rs := byName["rs"].Values + if rs.Dtype() != core.Int16 { + t.Fatalf("rs dtype = %s, want int16", rs.Dtype()) + } + if s := rs.Shape(); len(s) != 1 || s[0] != 2 { + t.Fatalf("rs shape = %v, want [2]: the record axis carries the record count", s) + } + if got, want := rs.RawInt16s()[:2], []int16{0x1234, -16}; !slices.Equal(got, want) { + t.Fatalf("rs values = %v, want %v", got, want) + } + ss := byName["ss"].Values + if ss.Dtype() != core.Int16 { + t.Fatalf("ss dtype = %s, want int16", ss.Dtype()) + } + if got, want := ss.RawInt16s()[:3], []int16{5, -6, 7}; !slices.Equal(got, want) { + t.Fatalf("ss values = %v, want %v", got, want) + } +} + +// TestNetCDFUnknownTypeRefused pins the refusal of every type code +// beyond the classic six: the unsigned and 64-bit codes exist only in +// formats this module does not speak, and a header carrying one is +// refused by name rather than decoded into some other dtype. +func TestNetCDFUnknownTypeRefused(t *testing.T) { + const template = "" + + "43444601" + // magic + "00000000" + // numrecs = 0 + "0000000000000000" + // dim_list ABSENT + "0000000000000000" + // gatt_list ABSENT + "0000000b" + // NC_VARIABLE + "00000001" + // one variable + "0000000176000000" + // name "v" + pad + "00000000" + // rank 0 + "0000000000000000" + // vatt ABSENT + "00000007" + // type code, replaced per case + "00000000" + // vsize + "00000000" // begin + for _, tc := range []struct{ code, word string }{ + {"7", "00000007"}, {"8", "00000008"}, {"9", "00000009"}, + {"10", "0000000a"}, {"11", "0000000b"}, + } { + raw, err := hex.DecodeString(strings.Replace(template, "00000007", tc.word, 1)) + if err != nil { + t.Fatalf("hex: %v", err) + } + path := ncTempPath(t) + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + _, _, _, err = LoadNetCDF(path) + if err == nil || !strings.Contains(err.Error(), "unknown type "+tc.code) { + t.Fatalf("type code %s: err = %v, want the unknown-type refusal", tc.code, err) + } + } +} + +// TestNetCDFDecodeCellsRefused pins the explicit refusal inside +// decodeCells itself: an ncType without a case is a named error, never +// a buffer left silently untouched. +func TestNetCDFDecodeCellsRefused(t *testing.T) { + arr := core.New(core.Float, 2) + for _, code := range []uint32{0, 7, 8, 9, 10, 11, 12, 99} { + err := decodeCells(arr, make([]byte, 16), code, 0, 2) + if err == nil || !strings.Contains(err.Error(), "unknown type") { + t.Fatalf("decodeCells with type %d: err = %v, want the unknown-type refusal", code, err) + } + } +} diff --git a/io/netcdfrecord_test.go b/io/netcdfrecord_test.go new file mode 100644 index 0000000..f79104b --- /dev/null +++ b/io/netcdfrecord_test.go @@ -0,0 +1,158 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/hex" + "os" + "path/filepath" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The record (unlimited) dimension: written and read back, and read +// from a file another implementation wrote. + +// TestNetCDFRecordRoundTrip writes a file with a record dimension and a +// fixed variable, then reads it back: the record count, the shapes and +// the values must survive, and the record axis must stay first. +func TestNetCDFRecordRoundTrip(t *testing.T) { + path := filepath.Join(t.TempDir(), "rec.nc") + tp := mustFloats(t, []float64{1.5, 2.5, 3.5}, 3) + sp := mustInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3) + vp := mustFloats32(t, []float32{7, 8, 9}, 3) + dims := []NetCDFDim{{Name: "time", Length: 0}, {Name: "x", Length: 3}} + vars := []NetCDFVar{ + {Name: "t", Dims: []string{"time"}, Values: tp}, + {Name: "s", Dims: []string{"time", "x"}, Values: sp}, + {Name: "v", Dims: []string{"x"}, Values: vp}, + } + if err := SaveNetCDF(path, dims, vars, map[string]string{"title": "record round trip"}); err != nil { + t.Fatalf("SaveNetCDF: %v", err) + } + gotDims, gotVars, attrs, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if len(gotDims) != 2 || gotDims[0].Name != "time" || gotDims[0].Length != 0 || gotDims[1].Length != 3 { + t.Fatalf("dims = %+v, want the record dimension first with length 0", gotDims) + } + if attrs["title"] != "record round trip" { + t.Fatalf("attrs = %v", attrs) + } + if len(gotVars) != 3 { + t.Fatalf("vars = %d, want 3", len(gotVars)) + } + byName := map[string]NetCDFVar{} + for _, v := range gotVars { + byName[v.Name] = v + } + if s := byName["t"].Values.Shape(); s[0] != 3 { + t.Fatalf("t shape = %v, want [3]", s) + } + if got := byName["s"].Values.Shape(); got[0] != 3 || got[1] != 3 { + t.Fatalf("s shape = %v, want [3 3]", got) + } + for i, want := range []float64{1.5, 2.5, 3.5} { + if got := byName["t"].Values.FloatAt(i); got != want { + t.Fatalf("t[%d] = %v, want %v", i, got, want) + } + } + for i, want := range []float64{1, 2, 3, 4, 5, 6, 7, 8, 9} { + if got := byName["s"].Values.FloatAt(i); got != want { + t.Fatalf("s[%d] = %v, want %v", i, got, want) + } + } + for i, want := range []float64{7, 8, 9} { + if got := byName["v"].Values.FloatAt(i); got != want { + t.Fatalf("v[%d] = %v, want %v", i, got, want) + } + } +} + +// TestNetCDFRecordFromNetCDF4 reads a file the netCDF reference +// implementation wrote: an unlimited time axis, a 1-D record variable +// of NC_DOUBLE, a 2-D record variable of NC_SHORT whose six-byte slab +// is padded to eight on disk, and a fixed NC_FLOAT variable. It pins +// the layout a foreign writer produces, padding included. +func TestNetCDFRecordFromNetCDF4(t *testing.T) { + const dump = "" + + "43444601000000020000000a000000020000000474696d6500000000000000017800000000000003" + + "00000000000000000000000b00000003000000017400000000000001000000000000000000000000" + + "0000000600000008000000b400000001730000000000000200000000000000010000000000000000" + + "0000000300000008000000bc00000001760000000000000100000001000000000000000000000005" + + "0000000c000000a840e0000041000000411000003ff8000000000000000100020003800140040000" + + "000000000004000500068001" + raw, err := hex.DecodeString(dump) + if err != nil { + t.Fatalf("hex: %v", err) + } + path := filepath.Join(t.TempDir(), "foreign.nc") + if err := os.WriteFile(path, raw, 0o644); err != nil { + t.Fatalf("write: %v", err) + } + dims, vars, _, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if len(dims) != 2 || dims[0].Name != "time" || dims[0].Length != 0 || dims[1].Length != 3 { + t.Fatalf("dims = %+v", dims) + } + if len(vars) != 3 { + t.Fatalf("vars = %d, want 3", len(vars)) + } + byName := map[string]NetCDFVar{} + for _, v := range vars { + byName[v.Name] = v + } + // Two records of one double. + if got := byName["t"].Values.RawFloats(); len(got) != 2 || got[0] != 1.5 || got[1] != 2.5 { + t.Fatalf("t = %v, want [1.5 2.5]", got) + } + // Two records of three shorts, each slab padded on disk. + if got := byName["s"].Values.Shape(); got[0] != 2 || got[1] != 3 { + t.Fatalf("s shape = %v, want [2 3]", got) + } + for i, want := range []float64{1, 2, 3, 4, 5, 6} { + if got := byName["s"].Values.FloatAt(i); got != want { + t.Fatalf("s[%d] = %v, want %v", i, got, want) + } + } + // The fixed variable, written before the records. Every variable + // comes back widened to float64, whatever the file stores. + for i, want := range []float64{7, 8, 9} { + if got := byName["v"].Values.FloatAt(i); got != want { + t.Fatalf("v[%d] = %v, want %v", i, got, want) + } + } +} + +// TestNetCDFRecordZeroRecords pins the empty unlimited dimension: a +// file may declare the record dimension and hold no records at all. +func TestNetCDFRecordZeroRecords(t *testing.T) { + path := filepath.Join(t.TempDir(), "empty.nc") + empty, err := core.FromFloats([]float64{}, 0) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + err = SaveNetCDF(path, []NetCDFDim{{Name: "time", Length: 0}}, + []NetCDFVar{{Name: "t", Dims: []string{"time"}, Values: empty}}, nil) + if err != nil { + t.Fatalf("SaveNetCDF: %v", err) + } + dims, vars, _, err := LoadNetCDF(path) + if err != nil { + t.Fatalf("LoadNetCDF: %v", err) + } + if len(dims) != 1 || dims[0].Length != 0 { + t.Fatalf("dims = %+v", dims) + } + if len(vars) != 1 || vars[0].Values.Len() != 0 { + t.Fatalf("vars = %+v, want one empty variable", vars) + } + if s := vars[0].Values.Shape(); s[0] != 0 { + t.Fatalf("shape = %v, want [0]", s) + } +} diff --git a/io/testdata/fuzz/FuzzLoadHDF5/2fe5cd680e6d8b63 b/io/testdata/fuzz/FuzzLoadHDF5/2fe5cd680e6d8b63 new file mode 100644 index 0000000..e6b3b40 --- /dev/null +++ b/io/testdata/fuzz/FuzzLoadHDF5/2fe5cd680e6d8b63 @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x89HDF\r\n\x1a\n\x000000\b\b0000000000000000000000000000000000000000000000000`\x00\x00\x00\x00\x00\x00\x00000000000000000000000000\x010\x06\x000000\x00\x01\x00\x000000\x01\x00\x18\x000000\x01\x0100000000000\x00\x00\x00\x04\x00\x00\x00\x00\x00\x00\x00\x03\x00\x18\x00\x01\x00\x00\x00\x11 ?\x00\b\x00\x00\x00\x00\x00@\x004\v\x004\xff\x03\x00\x00\x00\x00\x00\x00\x05\x00\b\x00\x01\x00\x00\x00\x02\x03\x02\x01\x00\x00\x00\x00\v\x008\x00\x01\x00\x00\x00\x01\x02\x00\x00\x00\x00\x00\x00\x02\x00\b\x00\x01\x00\x01\x00shuffle\x00\b\x00\x00\x00\x00\x00\x00\x00\x01\x00\b\x00\x01\x00\x01\x00deflate\x00\x04\x00\x00\x00\x00\x00\x00\x00\b\x00\x18\x00\x00\x00\x00\x00\x03\x02\x02\x00 \x00\x00\x00\x00\x7f\x00\x02\x00\x00\x00\b\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00H\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00") diff --git a/io/testdata/fuzz/FuzzLoadHDF5/a7ab50e52191807a b/io/testdata/fuzz/FuzzLoadHDF5/a7ab50e52191807a new file mode 100644 index 0000000..bb55e9c --- /dev/null +++ b/io/testdata/fuzz/FuzzLoadHDF5/a7ab50e52191807a @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x89HDF\r\n\x1a\n\x000000\b\b0000000000000000000000000000000000000000000000000`\x00\x00\x00\x00\x00\x00\x00000000000000000000000000\x010\x01\x000000\x18\x00\x00\x000000\x06\x00\a\x000000\x01700000000000000") diff --git a/io/testdata/fuzz/FuzzLoadHDF5/ae4657f5205db80f b/io/testdata/fuzz/FuzzLoadHDF5/ae4657f5205db80f new file mode 100644 index 0000000..e7f6688 --- /dev/null +++ b/io/testdata/fuzz/FuzzLoadHDF5/ae4657f5205db80f @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x89HDF\r\n\x1a\n\x000000\b\b0000000000000000000000000000000000000000000000000`\x00\x00\x00\x00\x00\x00\x00000000000000000000000000\x010\x03\x000000\x18\x00\x00\x000000\x06\x00\x00\x00000000\x00\x00000000\x00\x000000") diff --git a/io/testdata/fuzz/FuzzLoadNetCDF/14c890c96cda7d9c b/io/testdata/fuzz/FuzzLoadNetCDF/14c890c96cda7d9c new file mode 100644 index 0000000..ba7f325 --- /dev/null +++ b/io/testdata/fuzz/FuzzLoadNetCDF/14c890c96cda7d9c @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("CDF\x0100000000\x00\x00\x00\x02\x00\x00\x00\x030000\x00\x00\x00\x03\x00\x00\x00\x030000\x00\x00\x00\x040000\x00\x00\x00\x01\x00\x00\x00\x0500000000\x00\x00\x00\x02\x00\x00\x00\x0400000000\x00\x00\x00\x02\x00\x00\x00\x040000\x00\x00\x00\x02\x00\x00\x00\x00\x00\x00\x00\x010000\x00\x00\x00\x00\x00\x00\x00\x060000\x00\x00\x000\x00\x00\x00\x040000\x00\x00\x00\x01\x00\x00\x00\x000000\x00\x00\x00\x00\x00\x00\x00\x040000\x00\x00\x000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000") diff --git a/io/testdata/fuzz/FuzzLoadNetCDF/c4c632324ec4a450 b/io/testdata/fuzz/FuzzLoadNetCDF/c4c632324ec4a450 new file mode 100644 index 0000000..e360b83 --- /dev/null +++ b/io/testdata/fuzz/FuzzLoadNetCDF/c4c632324ec4a450 @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("CDF\x0100000000\x00\x00\x00\x02\x00\x00\x00\x000000\x00\x00\x00\x01000000000000\x00\x00\x00\x000000\x00\x00\x00\x01\x00\x00\x00\x00\x00\x00\x00\x000000\x00\x00\x00\x00\x00\x00\x00\x010000\x00\x00\x000") diff --git a/io/testdata/h5/bool_enum.h5 b/io/testdata/h5/bool_enum.h5 new file mode 100644 index 0000000000000000000000000000000000000000..d582dab663c9127b2acb18f93a4e9072b6a3f9b8 GIT binary patch literal 712 zcmeD5aB<`1lHy_j0S*oZ76t(@6Gr@pf)h*-5f~pPp8#brLg@}Dy#p@J$N-X)fbs>Q z=A)|%337F10IGzU52K;l7+ydb98lWB)iD6Xgt-=G{|%@-j7rN%OfLpsjT1TZ{7;-L6|#03X~o1;%KFj5&pft<*YF4oLN} zz3~QI`v{09;1PHPo`7qa*_knsQGX;r!D3d4XFNOJpU<^4**@uMe|L+tcaov;xvGvAl+zNg<+#2CEnk9z3AhFk zWWB84May;P{1r5cxK55(t=|tc31YH2sI#9qIJm>*25ivo$#8-4d5E=)0l;2LTrp$p{Rc4n-Z2-W?CQDr1!|o^7ntoT0PO*ZSGb(O!I;U_ZeK3ycw(4Ymt% z$Q9(-#a2F-;?lDaAbQ&!#wR72P74`gSWwP%yoQlOsRc_@uyi84ig*e+o~B^c`#TB9 zz{i_|{p~@T;z&^5PV8^HX@4v07WEbJ7z7ae`g~~XWoYTJ^$#6S`yD!z_(-Il@;sw) zxFCMooOJQNR|W5Hw^rwwa@>~ZnZZuaPjKhWGrnCyQ|AnA`>1^UaO0`>m5yqEctHTL zy&R{#VK&&K$qkd1rG8J?TpT}LEmTvCfDtePM!*Od0V7}pjDQg^0!F|H7=eF;0R4}r z|Mc{jhn@-1<2HJfHhldiY2|bC->*RJvwQDP;pVMI1KbLH1jyIWsf8c>=J>T_1dM%p6Eti+xDH>}y+*m8b65 Z(asH38H?ooY_toC)8^f=G^6sB{s3SHfQA47 literal 0 HcmV?d00001 diff --git a/io/testdata/h5/fletcher.h5 b/io/testdata/h5/fletcher.h5 new file mode 100644 index 0000000000000000000000000000000000000000..f80c835a1497563017ebba8f0396c46f4e625371 GIT binary patch literal 3612 zcmeD5aB<`1lHy_j0S*oZ76t(@6Gr@p0vSGt2#gPtPk=HQp>zk7Ucm%mFfxE31A_!q zTo7tLy1I}cS62q0N|^aD8mf)KfCa+hfC-G!BPs+uTpa^I9*%(e8kR~=K+_p4FjAll zSbFq;Nsvi1GO&TuFN6T4P)JHB43M>_hQi2TKYtgH`(PF*K>Y=iAEk*40Z@6z2#j4=IR~SqaA;q_3z3k6 z%0uHCuKEy~JfuK}OEWw`ljlcBfm#2c@-ShTJS?7J;-mCv2#kinXb6mkz-S1Jh5-FS zpdv0g!GTTw%|ae_j-!lAS!Ebh?=x|1nRx3hNQQw?klo;@0^`!u1zVV=gcUI`$o^qG zUGj%d6e24MlwH8K)O5iXu_@OI85opqFma|x-dhQgwFSy1Ff8R=xJ7V^Ss4R^ngNsP IzJ%sA07CE7#SIvq2|D85e7y<1$zdF>m3+O zxEW0T10`93u4Crn0J@x+5h}?b0dvC!7-Q(rTVsTJiI0v{Cnl)ZJgDczncrVcMvo3y zEW>DKBBDbAn$n;Y1H%xE4j+XB($ovtQIkePU^E0qLx4UZ@WCExj{}6}fYJgE0Dvq- AX8-^I literal 0 HcmV?d00001 diff --git a/io/walk_cycle_pins_test.go b/io/walk_cycle_pins_test.go new file mode 100644 index 0000000..d2ff9a6 --- /dev/null +++ b/io/walk_cycle_pins_test.go @@ -0,0 +1,204 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "encoding/binary" + "fmt" + "math" + "runtime" + "strings" + "testing" + "time" +) + +// Regression pins for the HDF5 group walk: deep chains that must stay +// linear, the depth cap, self-links and hard-link diamonds that must +// be refused as cycles. + +// groupChain builds a well-formed HDF5 file holding a chain of +// depth groups, each carrying one named attribute and a hard link to the +// next. Every group is legal, the attributes are legal scalar attributes +// and the chain terminates, so the reader has no grounds to refuse it: +// the walk has to stay cheap on its own. +func groupChain(depth int) []byte { + const ( + firstGroup = 128 + stride = 88 + ) + n := firstGroup + depth*stride + 16 + f := make([]byte, n) + copy(f, hdf5Magic) + f[8] = 0 // superblock version 0 + f[13] = 8 + f[14] = 8 + binary.LittleEndian.PutUint64(f[32:], math.MaxUint64) // free space undefined + binary.LittleEndian.PutUint64(f[40:], uint64(n)) // end of file + binary.LittleEndian.PutUint64(f[48:], math.MaxUint64) // driver undefined + binary.LittleEndian.PutUint64(f[64:], firstGroup) // root object header + + for k := range depth { + a := firstGroup + k*stride + nmsg := 2 + if k == depth-1 { + nmsg = 1 // the tail group only carries its attribute + } + f[a] = 1 // object header version 1 + binary.LittleEndian.PutUint16(f[a+2:], uint16(nmsg)) + binary.LittleEndian.PutUint32(f[a+4:], 1) // reference count + binary.LittleEndian.PutUint32(f[a+8:], 88) // message data size + + // Attribute message (type 12), size 40, body at a+24. + binary.LittleEndian.PutUint16(f[a+16:], hdf5MsgAttribute) + binary.LittleEndian.PutUint16(f[a+18:], 40) + f[a+24] = 1 // version + binary.LittleEndian.PutUint16(f[a+26:], 8) // name length + binary.LittleEndian.PutUint16(f[a+28:], 8) // datatype length + binary.LittleEndian.PutUint16(f[a+30:], 8) // dataspace length + copy(f[a+32:], fmt.Sprintf("a%07d", k)) // 8-byte attribute name + f[a+40] = 0x10 // fixed-point datatype + binary.LittleEndian.PutUint32(f[a+44:], 1) // one byte + f[a+48] = 1 // scalar dataspace + f[a+56] = byte(k) // one byte of value + + if k == depth-1 { + continue + } + // Link message (type 6), size 12, body at a+72. + binary.LittleEndian.PutUint16(f[a+64:], hdf5MsgLink) + binary.LittleEndian.PutUint16(f[a+66:], 12) + f[a+72] = 1 // version + f[a+73] = 0 // flags: hard link, 1-byte name length + f[a+74] = 1 // name length + f[a+75] = 'a' + binary.LittleEndian.PutUint64(f[a+76:], uint64(a+stride)) + } + return f +} + +// selfLink builds the smallest HDF5 file whose root group links to +// itself: 136 bytes, one object header, one link message. +func selfLink() []byte { + const N = 136 + f := make([]byte, N) + copy(f, hdf5Magic) + f[8] = 0 // superblock version 0 + f[13] = 8 + f[14] = 8 + binary.LittleEndian.PutUint64(f[32:], math.MaxUint64) + binary.LittleEndian.PutUint64(f[40:], N) + binary.LittleEndian.PutUint64(f[48:], math.MaxUint64) + binary.LittleEndian.PutUint64(f[64:], 96) // root object header + + f[96] = 1 // object header version 1 + binary.LittleEndian.PutUint16(f[98:], 1) + binary.LittleEndian.PutUint32(f[100:], 1) + binary.LittleEndian.PutUint32(f[104:], 24) + + binary.LittleEndian.PutUint16(f[112:], hdf5MsgLink) + binary.LittleEndian.PutUint16(f[114:], 12) + f[120] = 1 // link message version + f[121] = 0 // hard link, 1-byte name length + f[122] = 1 // name length + f[123] = 'a' + binary.LittleEndian.PutUint64(f[124:], 96) // the link target: itself + return f +} + +// TestWalkRefusesHardLinkCycle pins the crash that took the editor down: +// walk had no visited set, so a group linked into itself recursed for +// ever while the path string grew, and the heap grew with it at hundreds +// of megabytes per second until the host ran out of memory. The reader +// must refuse the file with an error. +func TestWalkRefusesHardLinkCycle(t *testing.T) { + guard := time.AfterFunc(20*time.Second, func() { panic("LoadHDF5 on a cyclic file did not return") }) + defer guard.Stop() + + path := writeHostile(t, "selflink.h5", selfLink()) + stop := warnHeap(t, 256<<20) + _, err := LoadHDF5(path) + stop() + if err == nil { + t.Fatal("LoadHDF5 accepted a group that hard-links into itself") + } + if !strings.Contains(err.Error(), "hard-link cycle") { + t.Fatalf("LoadHDF5 = %v, want the cycle refusal", err) + } +} + +// TestWalkDeepChainStaysLinear pins the other half of the same crash: a +// well-formed chain of groups used to cost memory quadratic in its +// depth, because every level copied the inherited attribute map and +// built a longer path string. The live memory must grow with the depth, +// not with its square, and the walk carries one shared path buffer and +// one shared attribute map to make that so. +func TestWalkDeepChainStaysLinear(t *testing.T) { + guard := time.AfterFunc(60*time.Second, func() { panic("deep chain walk did not return") }) + defer guard.Stop() + + measure := func(depth int) (uint64, int) { + data := groupChain(depth) + path := writeHostile(t, fmt.Sprintf("chain%d.h5", depth), data) + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + if _, err := LoadHDF5(path); err != nil { + t.Fatalf("depth %d: LoadHDF5 refused a well-formed file: %v", depth, err) + } + runtime.ReadMemStats(&after) + return after.TotalAlloc - before.TotalAlloc, len(data) + } + small, smallFile := measure(150) + big, bigFile := measure(450) + // Linear: three times the depth is about three times the bytes. The + // per-level copies this replaced measured 16.5x for the same step. + if ratio := float64(big) / float64(small); ratio > 6 { + t.Errorf("allocation grows with the square of the depth: %.1fx for 3x depth (%d B for %d B, %d B for %d B)", + ratio, big, bigFile, small, smallFile) + } +} + +// TestWalkRefusesUnboundedNesting pins the depth cap: past it the reader +// refuses the file loudly and cheaply instead of walking every level +// first. +func TestWalkRefusesUnboundedNesting(t *testing.T) { + guard := time.AfterFunc(60*time.Second, func() { panic("deep chain walk did not return") }) + defer guard.Stop() + + deep := writeHostile(t, "deep.h5", groupChain(2*hdf5MaxGroupDepth)) + stop := warnHeap(t, 256<<20) + _, err := LoadHDF5(deep) + stop() + if err == nil { + t.Fatalf("LoadHDF5 accepted a chain %d groups deep", 2*hdf5MaxGroupDepth) + } + if !strings.Contains(err.Error(), "nest deeper than") { + t.Fatalf("LoadHDF5 = %v, want the nesting refusal", err) + } +} + +// warnHeap fails the test as soon as the live heap passes limit, so a +// regression cannot consume the machine: the watchdog panics, which +// unwinds the offending walk instead of letting it allocate on. +func warnHeap(t *testing.T, limit uint64) func() { + t.Helper() + done := make(chan struct{}) + go func() { + ticker := time.NewTicker(20 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-done: + return + case <-ticker.C: + var m runtime.MemStats + runtime.ReadMemStats(&m) + if m.HeapAlloc > limit { + panic("walk heap above its cap") + } + } + } + }() + return func() { close(done) } +} diff --git a/io/wave_bench_test.go b/io/wave_bench_test.go new file mode 100644 index 0000000..5574070 --- /dev/null +++ b/io/wave_bench_test.go @@ -0,0 +1,61 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +// Benchmarks for the paths the coordinator's battery does not isolate. +// Every fixture is deterministic and built before the measured loop. + +import ( + "os" + "path/filepath" + "testing" +) + +// BenchmarkWaveFITSBinaryTextCells measures the character-column decode +// alone: one 8A column of four thousand rows, the per-cell work the +// binary-table benchmark mixes with seven numeric columns. +func BenchmarkWaveFITSBinaryTextCells(b *testing.B) { + const rows = 4000 + path := filepath.Join(b.TempDir(), "textcells.fits") + body := make([]byte, rows*8) + for r := range rows { + copy(body[r*8:], "star ") + } + cards := []string{ + fitsStringCardRaw("XTENSION", "BINTABLE"), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 2), + fitsIntCard("NAXIS1", 8), + fitsIntCard("NAXIS2", rows), + fitsIntCard("PCOUNT", 0), + fitsIntCard("GCOUNT", 1), + fitsIntCard("TFIELDS", 1), + fitsStringCardRaw("TTYPE1", "STAR"), + fitsStringCardRaw("TFORM1", "8A"), + fitsEndCard(), + } + out := fitsAppendCards(nil, []string{ + fitsBoolCard("SIMPLE", true), + fitsIntCard("BITPIX", 8), + fitsIntCard("NAXIS", 0), + fitsBoolCard("EXTEND", true), + fitsEndCard(), + }) + out = fitsAppendCards(out, cards) + out = append(out, body...) + out = fitsAppendZeroPad(out) + if err := os.WriteFile(path, out, 0o644); err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + table, err := LoadFITSTable(path) + if err != nil { + b.Fatal(err) + } + if table.Rows != rows || table.Text[0] == nil || table.Text[0][0] != "star" { + b.Fatalf("table came back as %s with %d rows", table.Kind, table.Rows) + } + } +} diff --git a/io/zz_guard_test.go b/io/zz_guard_test.go new file mode 100644 index 0000000..58e257b --- /dev/null +++ b/io/zz_guard_test.go @@ -0,0 +1,61 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package io + +import ( + "flag" + "fmt" + "os" + "runtime" + "testing" + "time" +) + +// TestMain installs a heap watchdog for the whole package run. The io +// tests are hostile-input tests by design: they feed crafted headers to +// the readers and check that nothing panics and nothing grows without +// bound. A regression therefore cannot fail politely, it allocates: a +// hard-link cycle and a deep group chain both took the host down once. +// The watchdog turns any such runaway into an immediate panic that +// unwinds the offending code path, so the worst case is a failed test +// rather than an OOM kill of the editor that started it. +func TestMain(m *testing.M) { + // The aggregate read budget bounds what one input may allocate + // legally; the watchdog sits an order above it. Fuzzing runs many + // workers holding inputs at once and the engine carries its own + // corpus and coverage bookkeeping, so the fuzz run watches a wider + // ceiling while still catching the unbounded growth the guard + // exists for: a runaway reaches any finite cap in seconds. + // + // The flags are parsed here, before the ceiling is decided: the + // test flags are only registered at TestMain and m.Run parses them + // later, so an earlier read of test.fuzz sees its empty default and + // the fuzz ceiling never fires. + flag.Parse() + const plainCap = 3 << 30 + ceiling := int64(plainCap) + if f := flag.Lookup("test.fuzz"); f != nil && f.Value.String() != "" { + ceiling = 32 << 30 + } + stop := make(chan struct{}) + go func() { + ticker := time.NewTicker(20 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-stop: + return + case <-ticker.C: + var m runtime.MemStats + runtime.ReadMemStats(&m) + if int64(m.HeapAlloc) > ceiling { + panic(fmt.Sprintf("io test binary heap above %d GiB: a reader is allocating without bound", ceiling>>30)) + } + } + } + }() + code := m.Run() + close(stop) + os.Exit(code) +} diff --git a/justfile b/justfile new file mode 100644 index 0000000..d4c9901 --- /dev/null +++ b/justfile @@ -0,0 +1,295 @@ +# tensor: a scientific computing library in pure Go. +# +# A library: nothing to install, nothing to run. Below the variable +# block is the standard recipe contract; docs-check and fuzz-all are +# the project extensions. The portable build is the only build. + +# What the test, race, unit and bench recipes sweep: every logic +# package, the examples aside. The examples are main programs with no +# tests; the build compiles them, and the coverage floor is a property +# of the library packages alone. +packages := ". ./internal/... ./grad/... ./integrate/... ./io/... ./linalg/... ./optim/... ./plot/... ./signal/... ./spmd/... ./stats/..." + +# Where the fuzz targets live; the fuzz-all sweep discovers them by name. +fuzz_packages := "./io ./grad" + +# The memory fence for the test recipes: a cgroup ceiling with swap off, so a +# runaway run dies as a failed run and never eats the machine. 4G is the +# default; raise it only with a reason recorded here. +memlimit := "4G" + +# The representative set the bench-report measures, one benchmark per +# kernel family across the packages. The recipe measures the names this +# release tag and the head tree share, so an old tag never blocks the +# report. +bench_set := "Add1M Exp100k Dot SumAxis2D SumAxis3DDim1 MeanAxis2D Norm2D Prod2D CumSum2D ArgSort1M MatMul1024 MatMulSmall64 MatMulOdd130 MatMulTall512x64x2048 MatVec KernelTransposeTiled256 EinsumBatchedMatMul Cholesky256 Solve256 SVD128 FFT4096 FFT2_512x512 Conv1D CWTMorlet SavitzkyGolay WelchPSD IntegrateRK4 IntegrateBackwardEuler IntegrateHeat2D MinimiseLBFGSBounded MinimiseDifferentialEvolution BackwardTwoLayer LinearRegression LogisticRegression KernelDensity Median CovarianceMatrix" + +default: + @just --list + +# Compile. Zero errors, zero warnings. +build: + go build ./... + +# The test gate: the suite, no cache, the coverage floor, under the memory fence. +test: + #!/usr/bin/env perl + my @fence = (q{systemd-run}, q{--user}, q{--scope}, + q{-p}, q{MemoryMax={{memlimit}}}, q{-p}, q{MemorySwapMax=0}); + system(@fence, q{go}, q{test}, q{-count=1}, q{-timeout}, q{30m}, + q{-coverprofile}, q{coverage.out}, qw({{packages}})) == 0 + or die qq{the test suite failed\n}; + 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); + +# The same suite under the race detector. The expensive one, still fenced. +race: + systemd-run --user --scope -p MemoryMax={{memlimit}} -p MemorySwapMax=0 go test -race -count=1 -timeout 30m {{packages}} + +# Fast scoped run for iterating. This is the one that runs after every edit. +unit pkgs=packages run=".*": + systemd-run --user --scope -p MemoryMax={{memlimit}} -p MemorySwapMax=0 go test {{pkgs}} -run '{{run}}' + +# Time-boxed fuzz of one target in one package. The package is required; never a gate. +fuzz target pkg fuzztime="60s": + systemd-run --user --scope -p MemoryMax={{memlimit}} -p MemorySwapMax=0 go test -run '^$' -fuzz '{{target}}' -fuzztime={{fuzztime}} {{pkg}} + +# Benchmarks. On an idle machine only, deliberately unfenced: a ceiling would distort the measurement. +bench pkgs=packages: + go test -run '^$' -bench=. -benchmem -count=5 {{pkgs}} + +# Format in place. +fmt: + gofmt -w . + +# Zero diff. Prints nothing when everything is formatted. +fmt-check: + #!/usr/bin/env perl + open(my $g, q{-|}, q{gofmt}, q{-l}, q{.}) or die qq{gofmt: $!}; + my @bad = <$g>; + close($g); + print @bad; + exit(@bad ? 1 : 0); + +# Both static gates: go vet and go fix -diff. +vet: + go vet ./... + go fix -diff ./... + +# The definition of done, in one command. Once per task, never per edit. +gates: build fmt-check vet test race + +# Remove the build artefacts. +clean: + rm -rf bin/ coverage.out + +# Run every Go program in README.md. An extension: no standard gate covers prose. +docs-check: + #!/usr/bin/env perl + my $root = `pwd`; chomp $root; + my $tmp = $ENV{TMPDIR} || q{/tmp}; + my $base = qq{$tmp/tensor-docs-check}; + system(q{rm}, q{-rf}, $base) == 0 or die qq{rm $base failed\n}; + mkdir $base or die qq{mkdir $base: $!}; + open(my $md, q{<}, q{README.md}) or die qq{README.md: $!}; + my (@blocks, $cur); + while (my $l = <$md>) { + if ($l =~ m{^```go\s*$}) { $cur = q{}; next } + if (defined $cur && $l =~ m{^```\s*$}) { push @blocks, $cur; undef $cur; next } + $cur .= $l if defined $cur; + } + close($md) or die qq{README.md close failed\n}; + my ($i, $failed) = (0, 0); + for my $b (@blocks) { + $i++; + my $dir = sprintf qq{%s/%02d}, $base, $i; + mkdir $dir or die qq{mkdir $dir: $!}; + open(my $g, q{>}, qq{$dir/go.mod}) or die qq{$dir/go.mod: $!}; + print $g qq{module readme$i\n\ngo 1.27.1\n\n} + . qq{require sourcedock.dev/petrbalvin/tensor v0.0.0\n\n} + . qq{replace sourcedock.dev/petrbalvin/tensor => $root\n}; + close($g) or die qq{go.mod close failed\n}; + open(my $m, q{>}, qq{$dir/main.go}) or die qq{$dir/main.go: $!}; + print $m $b; + close($m) or die qq{main.go close failed\n}; + chdir $dir or die qq{chdir $dir: $!}; + local $ENV{GOFLAGS} = q{-mod=mod}; + open(my $r, q{-|}, q{go}, q{run}, q{.}) or die qq{go run $dir: $!}; + my @out = <$r>; + # close reports true on a clean exit; $? carries the status of one that failed. + my $ok = close($r); + my $status = $? >> 8; + unless ($ok) { + $failed++; + print qq{FAILED: README program $i in $dir (exit $status)\n}, @out; + } + chdir $root or die qq{chdir $root: $!}; + } + print qq{README programs: $i, failures: $failed\n}; + exit($failed ? 1 : 0); + +# Fuzz every target for FUZZTIME each, under the same fence. Exploration, never a gate. +fuzz-all fuzztime="5s": + #!/usr/bin/env perl + my @fence = (q{systemd-run}, q{--user}, q{--scope}, + q{-p}, q{MemoryMax={{memlimit}}}, q{-p}, q{MemorySwapMax=0}); + my $bad = 0; + for my $pkg (split q{ }, q({{fuzz_packages}})) { + open(my $l, q{-|}, q{go}, q{test}, q{-list}, q{^Fuzz}, $pkg) + or die qq{go test -list $pkg: $!}; + my @targets = map { chomp; $_ } grep { m{^Fuzz} } <$l>; + close($l) or die qq{go test -list $pkg failed\n}; + for my $t (@targets) { + print qq{=== fuzz $t in $pkg for {{fuzztime}} ===\n}; + system(@fence, q{go}, q{test}, q{-run}, q{^$}, q{-fuzz}, qq{^$t\$}, + q{-fuzztime}, q({{fuzztime}}), $pkg) == 0 + or do { $bad = 1; print qq{FAILED: $t\n} }; + } + } + exit($bad ? 1 : 0); + +# Benchmark the working tree against the latest release tag, rewriting docs/benchmarks/release-vs-head.md. +bench-report benchtime="0.5s": + #!/usr/bin/env perl + use v5.40; + my $repo = `git rev-parse --show-toplevel`; chomp $repo; + my @tags = grep { chomp; $_ } `git tag --list v* --sort=-v:refname`; + @tags or die qq{no release tag in the repository\n}; + my ($tag) = @tags; + my ($head, $dirty) = (`git rev-parse --short HEAD` =~ s/\s+\z//r); + $dirty = `git status --porcelain` eq q{} ? q{} : q{, dirty working tree}; + chomp $head; + my $rel_dir = $repo =~ s{/[^/]+\z}{}r . qq{/tensor-bench-release}; + END { system(q{git}, q{worktree}, q{remove}, q{--force}, $rel_dir) } + if (-d $rel_dir) { system(q{git}, q{worktree}, q{remove}, q{--force}, $rel_dir) } + system(q{git}, q{worktree}, q{add}, q{--detach}, $rel_dir, $tag) == 0 + or die qq{worktree for $tag failed\n}; + print qq{=== $tag ($rel_dir) against head $head$dirty ===\n}; + my %wanted = map { $_ => 1 } split q{ }, q({{bench_set}}); + my %map; # side -> name -> package + for my $side ([release => $rel_dir], [head => $repo]) { + my @pkgs = split /\n/, `go -C $side->[1] list ./...`; + for my $p (@pkgs) { + open(my $l, q{-|}, q{go}, q{-C}, $side->[1], q{test}, $p, q{-list}, q{Benchmark}) + or die qq{go -C $side->[1] test -list $p: $!}; + while (my $n = <$l>) { + next unless $n =~ m{^Benchmark(\w+)\s*\z}; + chomp $n; + $map{$side->[0]}{$1} = $p if $wanted{$1}; + } + close($l); + } + } + my %pkgs; # side -> package -> [names both sides know] + for my $n (sort keys %{$map{release}}) { + next unless $map{head}{$n}; + push @{$pkgs{release}{$map{release}{$n}}}, $n; + push @{$pkgs{head}{$map{head}{$n}}}, $n; + } + my (%ns, %bytes, %allocs); + for my $round (1 .. 4) { + my @sides = $round % 2 ? ([release => $rel_dir], [head => $repo]) + : ([head => $repo], [release => $rel_dir]); + for my $s (@sides) { + my ($side, $dir) = @$s; + for my $p (sort keys %{$pkgs{$side}}) { + my @names = @{$pkgs{$side}{$p}}; + my $sel = q{^(} . join(q{|}, map { qq{Benchmark$_} } @names) . q{)$}; + open(my $r, q{-|}, q{go}, q{-C}, $dir, q{test}, $p, q{-run}, q{^$}, + q{-bench}, $sel, q{-benchmem}, qq{-benchtime={{benchtime}}}, q{-count=2}) + or die qq{go -C $dir bench $p: $!}; + while (my $l = <$r>) { + next unless $l =~ m{^Benchmark(\w+)\S*\s+\d+\s+([\d.]+) ns/op(?:\s+(\d+) B/op(?:\s+(\d+) allocs/op)?)?}; + my $n = $1; + next unless $wanted{$n} && $map{release}{$n} && $map{head}{$n}; + push @{$ns{$n}{$side}}, $2; + push @{$bytes{$n}{$side}}, $3 if defined $3; + push @{$allocs{$n}{$side}}, $4 if defined $4; + } + close($r) or die qq{the bench run failed: $p ($side)\n}; + } + } + print qq{=== round $round done ===\n}; + } + sub median { my @v = sort { $a <=> $b } @_; my $m = int(@v / 2); @v % 2 ? $v[$m] : ($v[$m - 1] + $v[$m]) / 2 } + sub unit { + my ($ns) = @_; + return sprintf qq{%.2f ms}, $ns / 1e6 if $ns >= 1e6; + return sprintf qq{%.2f µs}, $ns / 1e3 if $ns >= 1e3; + return sprintf qq{%.0f ns}, $ns; + } + my (@rows, $faster, $slower, $same); + ($faster, $slower, $same) = (0, 0, 0); + for my $n (sort keys %ns) { + next unless @{$ns{$n}{release}} && @{$ns{$n}{head}}; + my ($mr, $mh) = (median(@{$ns{$n}{release}}), median(@{$ns{$n}{head}})); + my $f = $mr / $mh; + $f >= 1.05 ? $faster++ : $f <= 0.95 ? $slower++ : $same++; + my $bs = defined $bytes{$n}{release} && defined $bytes{$n}{head} + ? sprintf qq{%d → %d}, median(@{$bytes{$n}{release}}), median(@{$bytes{$n}{head}}) : q{—}; + my $as = defined $allocs{$n}{release} && defined $allocs{$n}{head} + ? sprintf qq{%d → %d}, median(@{$allocs{$n}{release}}), median(@{$allocs{$n}{head}}) : q{—}; + push @rows, [$n, unit($mr), unit($mh), sprintf(qq{%.2f}, $f), $bs, $as]; + } + @rows or die qq{no benchmark measured on both sides\n}; + @rows = sort { $b->[3] <=> $a->[3] } @rows; + open(my $ci, q{<}, q{/proc/cpuinfo}) or die qq{/proc/cpuinfo: $!}; + my ($cpu_name) = grep { m{^model name} } <$ci>; + close($ci); + ($cpu_name = $cpu_name // q{unknown CPU}) =~ s{^model name\s*:\s*}{}; + chomp $cpu_name; + my $cores = do { open(my $c3, q{<}, q{/proc/cpuinfo}); grep { m{^processor} } <$c3> }; + my $mem = do { open(my $m, q{<}, q{/proc/meminfo}); (grep { m{^MemTotal} } <$m>)[0] }; + $mem =~ s{^MemTotal:\s+([0-9]+) kB.*}{sprintf qq{%.0f GiB}, $1 / 1048576}e; + chomp(my $go_v = `go version`); chomp(my $os = `uname -sr`); chomp($mem); + my $report = qq{$repo/docs/benchmarks/release-vs-head.md}; + mkdir qq{$repo/docs/benchmarks}; + open(my $o, q{>}, $report) or die qq{$report: $!}; + print $o <<"EOF"; + # Benchmark report: $tag against head $head + + One live comparison, regenerated by `just bench-report` and never + accumulated: the newest release tag against the working tree, taken in + one interleaved session on an idle machine. Re-run it before a release + and commit the file the run writes. + + ## Environment + + - CPU: $cpu_name, $cores logical cores + - Memory: $mem + - OS: $os + - Go: $go_v + - Build: the portable build, no pinned `GOAMD64` level, no `GOEXPERIMENT` + + ## Revisions + + - Release: `$tag` + - Head: `$head`$dirty + - The set: the benchmarks `bench_set` names that both revisions carry; + a name either side lacks is left out, never counted as a result. + + ## Method + + Four interleaved rounds, the side order alternating between rounds, two + counts per side per round, `-benchtime={{benchtime}}` with `-benchmem`, + medians over the eight samples a side collects. A factor at or above + 1.05 is faster, at or below 0.95 slower, anything between noise. Time + of the release divides time of the head, so above one means the head is + faster. + + ## Results + + | Benchmark | $tag | head $head | factor | B/op release → head | allocs/op release → head | + |---|---|---|---|---|---| + EOF + for my $r (@rows) { print $o qq{| `$r->[0]` | $r->[1] | $r->[2] | $r->[3]x | $r->[4] | $r->[5] |\n} } + print $o qq{\n## Summary\n\n$faster of }, scalar @rows, + qq{ benchmarks sit above the noise band (faster), $slower below it (slower), $same inside it.\n}; + close($o) or die qq{$report close failed\n}; + print qq{=== $report written: $faster faster, $slower slower, $same unchanged ===\n}; + diff --git a/leak_test.go b/leak_test.go new file mode 100644 index 0000000..5f87dbc --- /dev/null +++ b/leak_test.go @@ -0,0 +1,176 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package tensor + +import ( + "os" + "path/filepath" + "runtime" + "testing" +) + +// The leak harness. A kernel that returns must leave nothing behind: +// every parallel worker joined, every buffer handed back or +// unreachable, no reference parked in a package global. Go garbage +// collects, so a leak here means unbounded growth, not a lost block, +// and the check is the same shape either way: run the heavy paths, +// force collection, and require the goroutine count and the live heap +// to come back to where they started. The buffers sync.Pool hoards are +// released by the collections the check forces first. + +// leakOps is the representative load: parallel kernels across every +// domain, the ones that spawn workers and the ones that allocate +// scratch proportional to the input, plus the untrusted-input readers. +// The data is sized so that a single retained buffer per round shows up +// far above the tolerance below: the vectors are 2 MiB and the matrix +// 32 KiB, while the tolerance is a fraction of a megabyte. +func leakOps(t *testing.T) []func() { + t.Helper() + const bigN = 1 << 18 // 2 MiB of float64 + big := randA(t, 31, bigN) + med := randA(t, 32, 1<<16) + sq := mustA(t, med.RawFloats()[:64*64], 64, 64) + // Symmetric, diagonally dominant: both factorisations stay stable. + mm := sq.RawFloats() + for i := range 64 { + for j := range 64 { + mm[i*64+j] = (mm[i*64+j] + mm[j*64+i]) / 2 + } + mm[i*64+i] += 20 + } + sig := randA(t, 33, 1<<14) + // Files the readers must refuse: a header that lies about its sizes + // and a table header with no data behind it. Reading them allocates + // nothing beyond the few dozen bytes they contain, and nothing may + // stay behind either. + dir := t.TempDir() + hostileNC := filepath.Join(dir, "hostile.nc") + if err := os.WriteFile(hostileNC, hostileOracleHeader(), 0o644); err != nil { + t.Fatal(err) + } + hostileFITS := filepath.Join(dir, "hostile.fits") + if err := os.WriteFile(hostileFITS, hostileOracleTable(), 0o644); err != nil { + t.Fatal(err) + } + return []func(){ + func() { _, _ = Add(big, big) }, + func() { _, _ = Mul(big, big) }, + func() { _, _ = MatMul2D(sq, sq) }, + func() { _, _ = Einsum("ij,jk->ik", sq, sq) }, + func() { _, _ = Sort(big) }, + func() { _, _ = ArgSort(big) }, + func() { _, _ = FFT(sig) }, + func() { _, _, _, _ = SVD(sq) }, + func() { _, _, _ = Eigen(sq) }, + func() { _, _ = Inv(sq) }, + func() { _, _ = Cholesky(sq) }, + func() { _, _, _ = QR(sq) }, + func() { _, _, _ = WelchPSD(sig, 1000, 256, 128, "hann") }, + func() { _, _ = SavitzkyGolay(sig, 21, 3) }, + func() { _, _ = DWT(sig, 4) }, + func() { _, _ = CWT(sig, Morlet, []float64{1, 2, 4, 8, 16}, 1) }, + func() { _, _ = SobolPoints(1<<12, 8, 0) }, + func() { _, _ = HaltonPoints(1<<12, 8, 0) }, + func() { _, _ = Gradient1D(sig, 1) }, + func() { _, _ = Laplacian(sq, 1, 1) }, + func() { + _, _ = IntegrateODE(func(t float64, y *Array) (*Array, error) { + return MulF(y, -1), nil + }, 0, 5, mustA(t, []float64{1}, 1), ODEOptions{MaxSteps: 100000}) + }, + func() { _, _, _, _ = LoadNetCDF(hostileNC) }, + func() { _, _ = LoadFITSTable(hostileFITS) }, + func() { _, _ = LoadHDF5(filepath.Join("io", "testdata", "h5", "fixture.h5")) }, + func() { + _, _ = LinearRegression(mustA(t, randA(t, 34, 4096).RawFloats(), 4096, 1), + randA(t, 35, 4096)) + }, + } +} + +// liveHeap returns the live heap after forcing every reclaimable byte +// to be collected twice: sync.Pool keeps one victim generation, so two +// rounds are what it takes for pooled buffers to be unreachable. +func liveHeap() (heap uint64, goroutines int) { + runtime.GC() + runtime.GC() + var m runtime.MemStats + runtime.ReadMemStats(&m) + return m.HeapAlloc, runtime.NumGoroutine() +} + +// leakBlock runs one measurement block over the given operations and +// returns the live heap and goroutine count after it, so the leak +// detector itself can be measured by the same instrument it uses. +func leakBlock(ops []func(), rounds int) (heap uint64, goroutines int) { + for range rounds { + for _, op := range ops { + op() + } + } + return liveHeap() +} + +// TestLeakHarnessDetectsLeaks is the harness's self-check: the same +// block measurement run over an operation that deliberately retains +// half a megabyte per round must fail, so a regression that empties +// the comparison or flips its direction cannot pass silently. +func TestLeakHarnessDetectsLeaks(t *testing.T) { + var retained [][]float64 + leaky := []func(){ + func() { retained = append(retained, make([]float64, 64<<10)) }, + } + leakBlock(leaky, 3) // settle, as the real harness does + first, _ := leakBlock(leaky, 10) + second, _ := leakBlock(leaky, 10) + // The closure must stay reachable through the second block's + // collections: without a later use, the caller's slot is dead + // during the call and the collector would free exactly the data + // the self-check means to catch. + runtime.KeepAlive(leaky) + if second <= first+256<<10 { + t.Fatalf("the leak harness measured %d then %d bytes: a 5 MiB retention over 10 rounds went undetected", first, second) + } +} + +func TestNoResourceLeaks(t *testing.T) { + ops := leakOps(t) + if len(ops) == 0 { + t.Fatal("the leak harness has no operations: a refactor emptied leakOps and the measurement below would pass vacuously") + } + run := func(rounds int) { + for range rounds { + for _, op := range ops { + op() + } + } + } + // Caches, pooled buffers and lazy tables settle in the first block. + // The measurement then compares successive blocks rather than the + // whole run against a baseline: a leak retains data per round and + // grows steadily block over block, while the runtime's own + // structures settle. Comparing endpoint to start once is what let + // a leak of a few hundred kilobytes per block hide under the + // settling curve. + const ( + block = 10 + tolerance = 256 << 10 + ) + run(3) + run(block) + prev, g0 := liveHeap() + for range 2 { + run(block) + heap, g := liveHeap() + if heap > prev+tolerance { + t.Errorf("live heap grew from %d to %d bytes over %d rounds of the facade load (tolerance %d): something retains data per round", + prev, heap, block, tolerance) + } + if g > g0 { + t.Errorf("goroutines grew from %d to %d over %d rounds of the facade load: a worker did not return", + g0, g, block) + } + prev = heap + } +} diff --git a/linalg/bench_decomp_test.go b/linalg/bench_decomp_test.go new file mode 100644 index 0000000..e13db82 --- /dev/null +++ b/linalg/bench_decomp_test.go @@ -0,0 +1,512 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "fmt" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Decomposition benchmarks for the shapes the solver benchmarks in +// perf_bench_test.go do not cover: the tall least-squares system, the +// pivoted QR sweep, the complex and matrix-function paths, and the SVD +// shapes whose cost sits in the reconstruction kernels rather than in +// the bidiagonalisation. Every input is a fixed literal formula, so the +// timings compare like with like across runs. + +// benchDecompFloats fills an m×n row-major matrix with a deterministic +// literal pattern: a diagonally weighted band plus a bounded pseudo-random +// ripple. +func benchDecompFloats(m, n int, diagonal float64) []float64 { + v := make([]float64, m*n) + for i := range m { + for j := range n { + v[i*n+j] = 0.5*float64((i*11+j*7)%17) - 4 + 0.25*float64((i*j*5)%13) + } + v[i*n+i%n] += diagonal + } + return v +} + +// benchDecompSPD builds a symmetric strictly diagonally dominant +// matrix, which is positive definite however the ripple lands. +func benchDecompSPD(n int) []float64 { + v := make([]float64, n*n) + for i := range n { + for j := range n { + v[i*n+j] = float64((i*5+j*3)%9) - 4 + } + } + for i := range n { + for j := i + 1; j < n; j++ { + avg := (v[i*n+j] + v[j*n+i]) / 2 + v[i*n+j], v[j*n+i] = avg, avg + } + row := 0.0 + for j := range n { + if j != i { + row += v[i*n+j] + if v[i*n+j] < 0 { + row += 2 * -v[i*n+j] + } + } + } + v[i*n+i] += row + float64(n) + } + return v +} + +// benchDecompHermitian builds a Hermitian matrix with a real diagonal +// whose entries are well separated, so the Jacobi sweep deflates +// quickly. +func benchDecompHermitian(n int) []complex128 { + v := make([]complex128, n*n) + for i := range n { + v[i*n+i] = complex(4*float64(i)+8, 0) + for j := i + 1; j < n; j++ { + re := 0.5*float64((i*3+j*5)%7) - 1.5 + im := 0.25*float64((i+j)%5) - 0.5 + v[i*n+j] = complex(re, im) + v[j*n+i] = complex(re, -im) + } + } + return v +} + +func benchDecompArray(b *testing.B, v []float64, m, n int) *core.Array { + b.Helper() + a, err := core.FromFloats(v, m, n) + if err != nil { + b.Fatal(err) + } + return a +} + +func benchDecompComplexArray(b *testing.B, v []complex128, m, n int) *core.Array { + b.Helper() + a, err := core.FromComplexes(v, m, n) + if err != nil { + b.Fatal(err) + } + return a +} + +// BenchmarkCholesky512 measures the blocked Cholesky sweep above the +// size where the solver benchmark's 256 stops. +func BenchmarkCholesky512(b *testing.B) { + a := benchDecompArray(b, benchDecompSPD(512), 512, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := Cholesky(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkLeastSquares512x32 measures the tall least-squares route, +// where the QR factor's orthogonal accumulation dominates. +func BenchmarkLeastSquares512x32(b *testing.B) { + const m, n = 512, 32 + a := benchDecompArray(b, benchDecompFloats(m, n, 3), m, n) + rhs, err := core.FromFloats(benchDecompFloats(m, 1, 1), m) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := LeastSquares(a, rhs); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkRRQR192x48 measures the pivoted QR sweep, whose cost is the +// per-step column-norm scan and the reflector application. +func BenchmarkRRQR192x48(b *testing.B) { + const m, n = 192, 48 + a := benchDecompArray(b, benchDecompFloats(m, n, 3), m, n) + b.ReportAllocs() + for b.Loop() { + if _, _, _, _, err := RRQR(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkRRQR128x32 measures the pivoted QR sweep at the size where a +// single reflector's work is small enough that crew sizing decides it. +func BenchmarkRRQR128x32(b *testing.B) { + const m, n = 128, 32 + a := benchDecompArray(b, benchDecompFloats(m, n, 3), m, n) + b.ReportAllocs() + for b.Loop() { + if _, _, _, _, err := RRQR(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSVDReconstruct1024x16 measures the SVD shapes whose cost sits +// in the reconstruction kernels: the full QR factor is far larger than +// the bidiagonalisation. +func BenchmarkSVDReconstruct1024x16(b *testing.B) { + const m, n = 1024, 16 + a := benchDecompArray(b, benchDecompFloats(m, n, 2), m, n) + b.ReportAllocs() + for b.Loop() { + if _, _, _, err := SVD(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkEigenComplex8 measures the complex Hermitian Jacobi sweep at +// the largest size its convergence floor admits: the sweep's off-mass +// threshold sits below the rounding floor the rotations leave behind +// above 8, so a bigger input errors out before it measures anything. +func BenchmarkEigenComplex8(b *testing.B) { + const n = 8 + a := benchDecompComplexArray(b, benchDecompHermitian(n), n, n) + b.ReportAllocs() + for b.Loop() { + if _, _, err := EigenComplex(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSVDComplex96x48 measures the direct complex bidiagonalisation +// and its Golub-Reinsch sweep. +func BenchmarkSVDComplex96x48(b *testing.B) { + const m, n = 96, 48 + flat := benchDecompFloats(m, n, 2) + v := make([]complex128, m*n) + for i := range m * n { + v[i] = complex(flat[i], 0.5*float64((i*3)%7)-1.5) + } + a := benchDecompComplexArray(b, v, m, n) + b.ReportAllocs() + for b.Loop() { + if _, _, _, err := SVDComplex(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkEigenGeneral32 measures the Hessenberg reduction, the shifted +// QR sweep and the eigenvector back-substitution on a general matrix. +func BenchmarkEigenGeneral32(b *testing.B) { + const n = 32 + v := benchDecompFloats(n, n, 4) + a := benchDecompArray(b, v, n, n) + b.ReportAllocs() + for b.Loop() { + if _, _, err := EigenGeneral(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSchurComplex48 measures the complex Schur decomposition the +// matrix functions are built on. +func BenchmarkSchurComplex48(b *testing.B) { + const n = 48 + flat := benchDecompFloats(n, n, 3) + v := make([]complex128, n*n) + for i := range n * n { + v[i] = complex(flat[i], 0.25*float64((i*5)%9)-1) + } + a := benchDecompComplexArray(b, v, n, n) + b.ReportAllocs() + for b.Loop() { + if _, _, err := SchurComplex(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMatrixExp32 measures the Padé scaling-and-squaring kernel. +func BenchmarkMatrixExp32(b *testing.B) { + const n = 32 + v := benchDecompFloats(n, n, 0) + v[0] = -0.5 // keep the 1-norm inside the unscaled degrees + a := benchDecompArray(b, v, n, n) + b.ReportAllocs() + for b.Loop() { + if _, err := MatrixExp(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMatrixSqrt64 measures the symmetric eigen route of a matrix +// function, whose cost is the eigendecomposition plus two triple +// products. +func BenchmarkMatrixSqrt64(b *testing.B) { + const n = 64 + a := benchDecompArray(b, benchDecompSPD(n), n, n) + b.ReportAllocs() + for b.Loop() { + if _, err := MatrixSqrt(a); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMatrixLogSchur32 measures the Schur-Parlett logarithm, the +// square-root walk and the Mercator series included. +func BenchmarkMatrixLogSchur32(b *testing.B) { + const n = 32 + flat := benchDecompSPD(n) + v := benchDecompComplexSPD(flat, n) + a := benchDecompComplexArray(b, v, n, n) + b.ReportAllocs() + for b.Loop() { + if _, err := MatrixLog(a); err != nil { + b.Fatal(err) + } + } +} + +// benchDecompComplexSPD lifts a symmetric positive definite real matrix +// to complex with a small positive imaginary part on the strict upper +// triangle and its conjugate below, keeping the spectrum off the +// non-positive real axis. +func benchDecompComplexSPD(flat []float64, n int) []complex128 { + v := make([]complex128, n*n) + for i := range n { + for j := range n { + v[i*n+j] = complex(flat[i*n+j], 0) + } + } + for i := range n { + for j := i + 1; j < n; j++ { + im := 0.05 * float64((i+j)%4+1) + v[i*n+j] = complex(flat[i*n+j], im) + v[j*n+i] = complex(flat[j*n+i], -im) + } + } + return v +} + +// BenchmarkTikhonov256x64 measures the SVD solve route, whose cost is a +// full SVD of the system matrix plus the Uᵀb and V·d products. +func BenchmarkTikhonov256x64(b *testing.B) { + const m, n = 256, 64 + a := benchDecompArray(b, benchDecompFloats(m, n, 2), m, n) + rhs, err := core.FromFloats(benchDecompFloats(m, 1, 1), m) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := SolveTikhonov(a, rhs, 0.5); err != nil { + b.Fatal(err) + } + } +} + +// TestDecompWorkerCountBitIdentity pins the parallel splits: every +// kernel this file benchmarks must answer bit for bit the same whether +// the crew is one goroutine or the machine's full width. A split that +// moved an addend or a boundary shows up here before the oracle's +// smaller pinned shapes see it. +func TestDecompWorkerCountBitIdentity(t *testing.T) { + prev := engine.SetNumWorkers(1) + defer engine.SetNumWorkers(prev) + + const ( + cholN = 320 + lsM = 384 + lsN = 24 + rrM = 160 + rrN = 40 + svdM = 384 + svdN = 24 + eigN = 192 + cplxN = 8 + ) + spd := benchDecompSPD(cholN) + tall := benchDecompFloats(lsM, lsN, 3) + rr := benchDecompFloats(rrM, rrN, 3) + wide := benchDecompFloats(svdM, svdN, 2) + herm := benchDecompHermitian(cplxN) + square := benchDecompSPD(eigN) + complexSquare := make([]complex128, cplxN*cplxN) + flatComplex := benchDecompFloats(cplxN, cplxN, 3) + for i := range cplxN * cplxN { + complexSquare[i] = complex(flatComplex[i], 0.25*float64((i*5)%9)-1) + } + + spdA := mustFromFloats(t, spd, cholN, cholN) + tallA := mustFromFloats(t, tall, lsM, lsN) + rrA := mustFromFloats(t, rr, rrM, rrN) + wideA := mustFromFloats(t, wide, svdM, svdN) + hermA := mustFromComplexes(t, herm, cplxN, cplxN) + squareA := mustFromFloats(t, square, eigN, eigN) + complexA := mustFromComplexes(t, complexSquare, cplxN, cplxN) + rhs := mustFromFloats(t, benchDecompFloats(lsM, 1, 1), lsM) + rrRHS := mustFromFloats(t, benchDecompFloats(rrM, 1, 1), rrM) + // A tall system wide enough that the blocked Qᵀb dispatches a crew, + // so the split is compared against the serial order too. + const lsWideN = 48 + wideLS := mustFromFloats(t, benchDecompFloats(lsM, lsWideN, 3), lsM, lsWideN) + + type snap struct { + name string + run func() []*core.Array + } + cases := []snap{ + {"Cholesky", func() []*core.Array { + l, err := Cholesky(spdA) + if err != nil { + t.Fatal(err) + } + return []*core.Array{l} + }}, + {"LeastSquares", func() []*core.Array { + x, err := LeastSquares(tallA, rhs) + if err != nil { + t.Fatal(err) + } + return []*core.Array{x} + }}, + {"LeastSquaresBlockedQtb", func() []*core.Array { + x, err := LeastSquares(wideLS, rhs) + if err != nil { + t.Fatal(err) + } + return []*core.Array{x} + }}, + {"RRQR", func() []*core.Array { + q, r, perm, rank, err := RRQR(rrA) + if err != nil { + t.Fatal(err) + } + pf := make([]float64, len(perm)+1) + for i, p := range perm { + pf[i] = float64(p) + } + pf[len(perm)] = float64(rank) + return []*core.Array{q, r, mustFromFloats(t, pf, len(pf), 1)} + }}, + {"SolveRRQR", func() []*core.Array { + x, err := SolveRRQR(rrA, rrRHS) + if err != nil { + t.Fatal(err) + } + return []*core.Array{x} + }}, + {"SVD", func() []*core.Array { + u, s, vt, err := SVD(wideA) + if err != nil { + t.Fatal(err) + } + return []*core.Array{u, s, vt} + }}, + {"Eigen", func() []*core.Array { + v, q, err := Eigen(squareA) + if err != nil { + t.Fatal(err) + } + return []*core.Array{v, q} + }}, + {"EigenComplex", func() []*core.Array { + v, q, err := EigenComplex(hermA) + if err != nil { + t.Fatal(err) + } + return []*core.Array{v, q} + }}, + {"SVDComplex", func() []*core.Array { + u, s, vh, err := SVDComplex(complexA) + if err != nil { + t.Fatal(err) + } + return []*core.Array{u, s, vh} + }}, + } + + serial := make([][]*core.Array, len(cases)) + for i, c := range cases { + serial[i] = c.run() + } + engine.SetNumWorkers(0) // the machine's full width + for i, c := range cases { + got := c.run() + for k := range got { + if got[k].Len() != serial[i][k].Len() { + t.Fatalf("%s: result %d length %d under the full crew, %d serial", + c.name, k, got[k].Len(), serial[i][k].Len()) + } + if !rawBitsEqual(got[k], serial[i][k]) { + t.Fatalf("%s: result %d differs bitwise between the serial and parallel crew", c.name, k) + } + } + } +} + +// rawBitsEqual compares two arrays' payloads bit for bit, complex +// payloads included. +func rawBitsEqual(a, b *core.Array) bool { + if a.Dtype() == core.Complex || b.Dtype() == core.Complex { + ac, bc := a.RawComplexes(), b.RawComplexes() + if len(ac) != len(bc) { + return false + } + for i := range ac { + if ac[i] != bc[i] { + return false + } + } + return true + } + af, bf := a.RawFloats(), b.RawFloats() + if len(af) != len(bf) { + return false + } + for i := range af { + if af[i] != bf[i] { + return false + } + } + return true +} + +// BenchmarkLeastSquaresTall measures the solve at the shapes the +// reflector route exists for: many more rows than columns, where forming +// Q would dominate everything else. +func BenchmarkLeastSquaresTall(b *testing.B) { + for _, c := range []struct{ m, n int }{{512, 32}, {2048, 64}, {8192, 16}} { + a := make([]float64, c.m*c.n) + s := uint64(20260920) + for i := range a { + s = s*6364136223846793005 + 1442695040888963407 + a[i] = float64((s>>40)%9+1) * 0.5 + } + bm := make([]float64, c.m) + for i := range bm { + bm[i] = float64(i%11) - 5 + } + am, err := core.FromFloats(a, c.m, c.n) + if err != nil { + b.Fatal(err) + } + bv, err := core.FromFloats(bm, c.m, 1) + if err != nil { + b.Fatal(err) + } + b.Run(fmt.Sprintf("%dx%d", c.m, c.n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := LeastSquares(am, bv); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/linalg/bench_factor_test.go b/linalg/bench_factor_test.go new file mode 100644 index 0000000..b8e22b8 --- /dev/null +++ b/linalg/bench_factor_test.go @@ -0,0 +1,421 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "fmt" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the dense LU kernel, the sparse factorisations and the +// sparse rank-one sweep. Every input is built from fixed literals, so +// the patterns, the pivots and the orderings are the same on every run. + +// denseLUInput builds an n×n row-major matrix whose rows view one flat +// buffer, with a pristine copy the benchmark restores from: the +// factorisation consumes its argument, and rebuilding the matrix inside +// the timed region would measure the rebuild. +func denseLUInput(n int) (rows [][]float64, pristine []float64) { + flat := make([]float64, n*n) + for i := range n { + for j := range n { + flat[i*n+j] = float64((i*5+j*11)%13) - 6 + } + // Diagonal dominance keeps the elimination on the + // well-conditioned side, so the timing measures the kernel. + flat[i*n+i] += float64(2 * n) + } + rows = make([][]float64, n) + for i := range n { + rows[i] = flat[i*n : (i+1)*n] + } + pristine = make([]float64, len(flat)) + copy(pristine, flat) + return rows, pristine +} + +func restoreDenseRows(rows [][]float64, pristine []float64) { + off := 0 + for _, row := range rows { + copy(row, pristine[off:off+len(row)]) + off += len(row) + } +} + +// BenchmarkDenseLUFactor measures the LU elimination alone at the sizes +// the solvers reach. +func BenchmarkDenseLUFactor(b *testing.B) { + for _, n := range []int{128, 256, 512} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + rows, pristine := denseLUInput(n) + b.ReportAllocs() + for b.Loop() { + restoreDenseRows(rows, pristine) + base.Factor(rows) + } + }) + } +} + +// denseLUSolveInput builds the same matrix as a core array, for the +// end-to-end solve path. +func denseLUSolveInput(b *testing.B, n int) (*core.Array, *core.Array) { + b.Helper() + flat := make([]float64, n*n) + for i := range n { + for j := range n { + flat[i*n+j] = float64((i*5+j*11)%13) - 6 + } + flat[i*n+i] += float64(2 * n) + } + a, err := core.FromFloats(flat, n, n) + if err != nil { + b.Fatal(err) + } + rhs := make([]float64, n) + for i := range rhs { + rhs[i] = float64(i%7) - 3 + } + x, err := core.FromFloats(rhs, n) + if err != nil { + b.Fatal(err) + } + return a, x +} + +// BenchmarkDenseLUSolve measures Solve end to end: the copy of the +// matrix, the factorisation, the permutation and the substitution. +func BenchmarkDenseLUSolve(b *testing.B) { + for _, n := range []int{128, 256, 512} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + a, x := denseLUSolveInput(b, n) + b.ReportAllocs() + for b.Loop() { + if _, err := Solve(a, x); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// sparseCOOFrom assembles a COO matrix from the triplets a builder +// appended, refusing nothing: the builders below emit canonical, +// duplicate-free triplets. +func sparseCOOFrom(b *testing.B, n int, idx []int64, vals []float64) *core.SparseCOO { + b.Helper() + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + b.Fatal(err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + b.Fatal(err) + } + return coo +} + +// gridLaplacian builds the 5-point Laplacian on a w×h grid in row-major +// order: symmetric positive definite, banded, and the pattern every +// sparse direct solver is measured on. +func gridLaplacian(b *testing.B, w, h int) *core.SparseCOO { + b.Helper() + n := w * h + idx := make([]int64, 0, 5*n) + vals := make([]float64, 0, 5*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + at := func(x, y int) int { return y*w + x } + for y := range h { + for x := range w { + add(at(x, y), at(x, y), 4) + if x+1 < w { + add(at(x, y), at(x+1, y), -1) + add(at(x+1, y), at(x, y), -1) + } + if y+1 < h { + add(at(x, y), at(x, y+1), -1) + add(at(x, y+1), at(x, y), -1) + } + } + } + return sparseCOOFrom(b, n, idx, vals) +} + +// arrowHead builds the arrowhead matrix of order n: a diagonal plus a +// dense first row and column. The diagonal dominates the first arrow +// (n+1 against n−1), so the matrix is positive definite, and the +// ordering has a genuine choice to make on the arrow. +func arrowHead(b *testing.B, n int) *core.SparseCOO { + b.Helper() + idx := make([]int64, 0, 3*n) + vals := make([]float64, 0, 3*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + add(0, 0, float64(n+1)) + for i := 1; i < n; i++ { + add(i, i, float64(i+2)) + add(0, i, 1) + add(i, 0, 1) + } + return sparseCOOFrom(b, n, idx, vals) +} + +// bandedNonsymmetric builds a banded, diagonally dominant matrix with a +// periodic spike below the diagonal: every tenth column has a +// subdiagonal entry larger than its diagonal, so the partial pivoting +// of the LU factorisation swaps rows and its label bookkeeping is +// exercised rather than measured cold. +func bandedNonsymmetric(b *testing.B, n int) *core.SparseCOO { + b.Helper() + idx := make([]int64, 0, 6*n) + vals := make([]float64, 0, 6*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + d := 6.0 + if i%10 == 3 { + d = 0.5 // the spike below takes this column's pivot + } + add(i, i, d) + if i+1 < n { + add(i, i+1, 1+0.25*float64(i%3)) + add(i+1, i, 1.5+0.5*float64(i%5)) + } + if i+3 < n { + add(i, i+3, 0.5) + } + } + return sparseCOOFrom(b, n, idx, vals) +} + +// swapHeavyBanded builds a tridiagonal matrix that pivots on every +// column: a small diagonal against a large subdiagonal, so the largest +// entry at or below the diagonal is always the row below. The factor +// stays banded, so the measurement is the pivot bookkeeping rather +// than the elimination arithmetic. +func swapHeavyBanded(b *testing.B, n int) *core.SparseCOO { + b.Helper() + idx := make([]int64, 0, 3*n) + vals := make([]float64, 0, 3*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 0.125) + if i+1 < n { + add(i, i+1, 1) + add(i+1, i, 16) + } + } + return sparseCOOFrom(b, n, idx, vals) +} + +// BenchmarkSparseLUFactorSwapped measures the elimination of a system +// that pivots at every column, so the cost of relabelling L's stored +// rows is part of the number. +func BenchmarkSparseLUFactorSwapped(b *testing.B) { + for _, n := range []int{100, 200, 400} { + coo := swapHeavyBanded(b, n) + b.Run(fmt.Sprintf("banded-%d", n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := NewSparseLU(coo); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// denseRHS builds a deterministic right hand side of length n. +func denseRHS(b *testing.B, n int) *core.Array { + b.Helper() + v := make([]float64, n) + for i := range v { + v[i] = float64(i%11) - 5 + 0.5*float64(i%3) + } + x, err := core.FromFloats(v, n) + if err != nil { + b.Fatal(err) + } + return x +} + +// BenchmarkSparseCholeskyFactor measures the symbolic and numeric +// elimination of a few hundred unknowns: the Laplacian grid with the +// natural order, then the same grid through the reverse Cuthill-McKee +// order, which meets a much smaller factor. +func BenchmarkSparseCholeskyFactor(b *testing.B) { + coo := gridLaplacian(b, 19, 19) + b.Run("grid-natural", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := NewSparseCholesky(coo, SparseOrderingNatural); err != nil { + b.Fatal(err) + } + } + }) + b.Run("grid-rcm", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := NewSparseCholesky(coo, SparseOrderingReverseCuthillMcKee); err != nil { + b.Fatal(err) + } + } + }) + coo = arrowHead(b, 300) + b.Run("arrow-natural", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := NewSparseCholesky(coo, SparseOrderingNatural); err != nil { + b.Fatal(err) + } + } + }) +} + +// BenchmarkSparseCholeskySolve measures the solve on a factor built +// once outside the loop: the gather, the two substitutions and the +// scatter. +func BenchmarkSparseCholeskySolve(b *testing.B) { + for _, size := range []struct { + name string + w, h int + ordering SparseOrdering + }{ + {"grid-natural", 19, 19, SparseOrderingNatural}, + {"grid-rcm", 19, 19, SparseOrderingReverseCuthillMcKee}, + } { + coo := gridLaplacian(b, size.w, size.h) + f, err := NewSparseCholesky(coo, size.ordering) + if err != nil { + b.Fatal(err) + } + rhs := denseRHS(b, size.w*size.h) + b.Run(size.name, func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := f.Solve(rhs); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// BenchmarkSparseLUFactor measures the left-looking elimination of a +// banded nonsymmetric system with pivoting. +func BenchmarkSparseLUFactor(b *testing.B) { + for _, n := range []int{200, 400} { + coo := bandedNonsymmetric(b, n) + b.Run(fmt.Sprintf("banded-%d", n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := NewSparseLU(coo); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// BenchmarkSparseLUSolve measures the forward and backward substitution +// on a factor built once outside the loop. +func BenchmarkSparseLUSolve(b *testing.B) { + const n = 400 + coo := bandedNonsymmetric(b, n) + f, err := NewSparseLU(coo) + if err != nil { + b.Fatal(err) + } + rhs := denseRHS(b, n) + b.Run("banded-400", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := f.Solve(rhs); err != nil { + b.Fatal(err) + } + } + }) +} + +// sparseCholRankOneInput factors the tridiagonal Laplacian and returns +// the factor with a two-entry vector whose support sits on the stored +// pattern, so both the update and the downdate are accepted. +func sparseCholRankOneInput(b *testing.B, n int) (*SparseCholesky, *core.Array) { + b.Helper() + idx := make([]int64, 0, 3*n) + vals := make([]float64, 0, 3*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 4) + if i+1 < n { + add(i, i+1, -1) + add(i+1, i, -1) + } + } + coo := sparseCOOFrom(b, n, idx, vals) + f, err := NewSparseCholesky(coo, SparseOrderingNatural) + if err != nil { + b.Fatal(err) + } + x := core.New(core.Float, []int{n}...) + x.RawFloats()[0] = 0.5 + x.RawFloats()[1] = 0.25 + return f, x +} + +// BenchmarkSparseCholRankOnePair measures one rank-one update followed +// by the matching downdate: each iteration leaves the factor where it +// found it, so the sweep is timed rather than the refactorisation the +// modification replaces. +func BenchmarkSparseCholRankOnePair(b *testing.B) { + for _, n := range []int{150, 300} { + f, x := sparseCholRankOneInput(b, n) + b.Run(fmt.Sprintf("tridiag-%d", n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if err := f.Update(x); err != nil { + b.Fatal(err) + } + if err := f.Downdate(x); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// BenchmarkSparseCholUpdateRefusal measures the refusal path: a vector +// whose support reaches outside the stored pattern, so every call +// returns the pattern error after the fill check has walked the support. +// It is the cost the check pays before it can refuse, and the error +// construction is deliberately part of it. +func BenchmarkSparseCholUpdateRefusal(b *testing.B) { + const n = 300 + f, _ := sparseCholRankOneInput(b, n) + x := core.New(core.Float, []int{n}...) + x.RawFloats()[0] = 0.5 + x.RawFloats()[n-1] = 0.25 + b.ReportAllocs() + for b.Loop() { + if err := f.Update(x); err == nil { + b.Fatal("expected the update to need fill the pattern does not hold") + } + } +} diff --git a/linalg/bench_iterative_test.go b/linalg/bench_iterative_test.go new file mode 100644 index 0000000..bfb5a18 --- /dev/null +++ b/linalg/bench_iterative_test.go @@ -0,0 +1,251 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Iterative solver benchmarks: the Krylov solves and the least-squares +// recursions, with and without the ILU(0) preconditioner. Every system +// is built from fixed literals so a run is reproducible, and each one +// is sized to finish well inside a second, which keeps the timings on +// the solvers' own per-iteration work rather than on input +// construction. + +// benchGridCOO builds the five-point stencil of an nx by ny grid as an +// n x n COO matrix (n = nx·ny) in natural row ordering. The diagonal is +// 4 and the vertical coupling −1 on both sides. A nonzero skew makes +// the horizontal coupling −1∓skew, which breaks the symmetry while +// keeping the matrix diagonally dominant; skew 0 leaves it symmetric +// positive-definite. +func benchGridCOO(b *testing.B, nx, ny int, skew float64) *core.SparseCOO { + b.Helper() + n := nx * ny + idx := make([]int64, 0, 5*n*2) + vals := make([]float64, 0, 5*n) + add := func(i, j int, v float64) { + idx = append(idx, int64(i), int64(j)) + vals = append(vals, v) + } + for r := range ny { + for c := range nx { + i := r*nx + c + add(i, i, 4) + if c+1 < nx { + add(i, i+1, -1-skew) + } + if c > 0 { + add(i, i-1, -1+skew) + } + if r+1 < ny { + add(i, i+nx, -1) + } + if r > 0 { + add(i, i-nx, -1) + } + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + b.Fatal(err) + } + values, err := core.FromFloats(vals, len(vals)) + if err != nil { + b.Fatal(err) + } + coo, err := core.NewSparseCOO(indices, values, []int{n, n}) + if err != nil { + b.Fatal(err) + } + return coo +} + +// benchTallCOO builds the overdetermined (2n − 2) x n matrix whose +// first n rows are the identity and whose remaining rows are the +// second-difference stencil over three adjacent columns. It is full +// rank and banded, so LSQR and LSMR converge in a bounded number of +// steps. +func benchTallCOO(b *testing.B, n int) *core.SparseCOO { + b.Helper() + m := 2*n - 2 + idx := make([]int64, 0, (4*n)*2) + vals := make([]float64, 0, 4*n) + add := func(i, j int, v float64) { + idx = append(idx, int64(i), int64(j)) + vals = append(vals, v) + } + for j := range n { + add(j, j, 1) + } + for r := range n - 2 { + i := n + r + add(i, r, -1) + add(i, r+1, 2) + add(i, r+2, -1) + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + b.Fatal(err) + } + values, err := core.FromFloats(vals, len(vals)) + if err != nil { + b.Fatal(err) + } + coo, err := core.NewSparseCOO(indices, values, []int{m, n}) + if err != nil { + b.Fatal(err) + } + return coo +} + +// benchPatternRHS builds the deterministic right-hand side of length n +// from a fixed integer pattern. +func benchPatternRHS(b *testing.B, n int) *core.Array { + b.Helper() + v := make([]float64, n) + for i := range v { + v[i] = float64(i%7) - 3 + } + a, err := core.FromFloats(v, n) + if err != nil { + b.Fatal(err) + } + return a +} + +// BenchmarkGMRESSparse measures a restarted GMRES solve of a +// nonsymmetric five-point system. +func BenchmarkGMRESSparse(b *testing.B) { + const n = 400 + coo := benchGridCOO(b, 20, 20, 0.5) + c, err := cooToCSR(coo, "BenchmarkGMRESSparse") + if err != nil { + b.Fatal(err) + } + rowStart, colIdx, vals := c.rowStart, c.colIdx, c.vals + op := func(v *core.Array) (*core.Array, error) { + out := core.New(core.Float, n) + dst := out.RawFloats() + for i := range n { + sum := 0.0 + for p := rowStart[i]; p < rowStart[i+1]; p++ { + sum += vals[p] * v.FloatAt(colIdx[p]) + } + dst[i] = sum + } + return out, nil + } + bv := benchPatternRHS(b, n) + b.ReportAllocs() + for b.Loop() { + if _, err := GMRES(op, bv, 30, 200, 1e-10); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSpSolveCG measures a conjugate-gradient solve of a +// symmetric five-point system under the Jacobi preconditioner. +func BenchmarkSpSolveCG(b *testing.B) { + coo := benchGridCOO(b, 20, 20, 0) + bv := benchPatternRHS(b, 400) + b.ReportAllocs() + for b.Loop() { + if _, err := SpSolve(coo, bv, 1e-10, 0); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSpSolveCGILU measures the same system under the ILU(0) +// preconditioner. +func BenchmarkSpSolveCGILU(b *testing.B) { + coo := benchGridCOO(b, 20, 20, 0) + bv := benchPatternRHS(b, 400) + ilu, err := NewSparseILU(coo) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := SpSolve(coo, bv, 1e-10, 0, ilu); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSpSolveBiCGSTAB measures a BiCGSTAB solve of a nonsymmetric +// five-point system under the Jacobi preconditioner. +func BenchmarkSpSolveBiCGSTAB(b *testing.B) { + coo := benchGridCOO(b, 20, 20, 0.5) + bv := benchPatternRHS(b, 400) + b.ReportAllocs() + for b.Loop() { + if _, err := SpSolveBiCGSTAB(coo, bv, 1e-10, 0); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSpSolveBiCGSTABILU measures the same system under the +// ILU(0) preconditioner. +func BenchmarkSpSolveBiCGSTABILU(b *testing.B) { + coo := benchGridCOO(b, 20, 20, 0.5) + bv := benchPatternRHS(b, 400) + ilu, err := NewSparseILU(coo) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := SpSolveBiCGSTAB(coo, bv, 1e-10, 0, ilu); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSpLSQR measures the Golub-Kahan least-squares recursion of +// SpLSQR on an overdetermined banded system. +func BenchmarkSpLSQR(b *testing.B) { + const n = 400 + coo := benchTallCOO(b, n) + bv := benchPatternRHS(b, 2*n-2) + _, info, err := SpLSQR(coo, bv, 1e-10, 0, 0) + if err != nil { + b.Fatal(err) + } + steps := float64(info.Iterations) + b.ReportAllocs() + for b.Loop() { + if _, _, err := SpLSQR(coo, bv, 1e-10, 0, 0); err != nil { + b.Fatal(err) + } + } + // The system and the tolerance are fixed, so the step count is too: + // reporting it shows the size of the work the timing covers. It goes + // after the loop, which clears any metric reported before it. + b.ReportMetric(steps, "steps") +} + +// BenchmarkSpLSMR measures the LSMR recursion on the same system. +func BenchmarkSpLSMR(b *testing.B) { + const n = 400 + coo := benchTallCOO(b, n) + bv := benchPatternRHS(b, 2*n-2) + _, info, err := SpLSMR(coo, bv, 1e-10, 0, 0) + if err != nil { + b.Fatal(err) + } + steps := float64(info.Iterations) + b.ReportAllocs() + for b.Loop() { + if _, _, err := SpLSMR(coo, bv, 1e-10, 0, 0); err != nil { + b.Fatal(err) + } + } + b.ReportMetric(steps, "steps") +} diff --git a/linalg/bench_sparse_test.go b/linalg/bench_sparse_test.go new file mode 100644 index 0000000..641e50e --- /dev/null +++ b/linalg/bench_sparse_test.go @@ -0,0 +1,318 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "fmt" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the sparse surface: the matrix-vector and matrix-matrix +// products, the coordinate-to-compressed construction, and the Krylov +// solvers that drive them. Every input is a fixed literal pattern, so a +// run is reproducible and two builds compare like with like. + +// benchSink keeps a kernel's result reachable so the compiler cannot +// elide the call it came from. +var ( + benchSinkCSR *SparseCSR + benchSinkCSC *SparseCSC + benchSinkArr *core.Array +) + +// laplacianTriples returns the flat (row, col, value) triples of the +// n×n symmetric banded Laplacian of the given half-bandwidth: 2·width +// on the diagonal and -1 at every offset up to width. It is symmetric +// positive definite by strict diagonal dominance, which is what the +// Lanczos and conjugate-gradient benchmarks need, and its non-zeros per +// row are 2·width+1 at every row. +func laplacianTriples(n, width int) []float64 { + tri := make([]float64, 0, n*(2*width+1)*3) + for i := range n { + tri = append(tri, float64(i), float64(i), float64(2*width)) + for d := 1; d <= width; d++ { + j := i + d + if j >= n { + break + } + tri = append(tri, float64(i), float64(j), -1) + tri = append(tri, float64(j), float64(i), -1) + } + } + return tri +} + +// duplicateTriples returns the triples of a matrix whose coordinates +// repeat: entry k carries row (k·13) mod n and column (k·37) mod n, so +// each coordinate appears about len/n times and the construction path +// has to merge duplicates rather than only sort distinct pairs. The +// values step along a literal modular sequence, so no merge cancels by +// accident. +func duplicateTriples(n, len_ int) []float64 { + tri := make([]float64, 0, len_*3) + for k := range len_ { + row := (k * 13) % n + col := (k * 37) % n + v := float64((k*7)%23) - 11 + if v == 0 { + v = 1.5 + } + tri = append(tri, float64(row), float64(col), v) + } + return tri +} + +// benchCOO builds a sparse matrix from flat (row, col, value) triples. +func benchCOO(b *testing.B, rows, cols int, tri []float64) *core.SparseCOO { + b.Helper() + nnz := len(tri) / 3 + idx := make([]int64, 0, nnz*2) + vals := make([]float64, 0, nnz) + for i := range nnz { + idx = append(idx, int64(tri[i*3]), int64(tri[i*3+1])) + vals = append(vals, tri[i*3+2]) + } + indices, err := core.FromInts(idx, nnz, 2) + if err != nil { + b.Fatalf("FromInts: %v", err) + } + values, err := core.FromFloats(vals, nnz) + if err != nil { + b.Fatalf("FromFloats: %v", err) + } + coo, err := core.NewSparseCOO(indices, values, []int{rows, cols}) + if err != nil { + b.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// benchVector returns a dense vector of length n whose entries come +// from a literal modular sequence, deliberately without zeros so a +// scaled product cannot skip work. +func benchVector(b *testing.B, n int) *core.Array { + b.Helper() + v := make([]float64, n) + for i := range n { + v[i] = float64((i*11)%17)/8 - 1 + if v[i] == 0 { + v[i] = 0.25 + } + } + arr, err := core.FromFloats(v, n) + if err != nil { + b.Fatalf("FromFloats: %v", err) + } + return arr +} + +func BenchmarkSparseCSRMatVec(b *testing.B) { + for _, n := range []int{256, 131072} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + coo := benchCOO(b, n, n, laplacianTriples(n, 1)) + csr, err := CSRFromCOO(coo) + if err != nil { + b.Fatalf("CSRFromCOO: %v", err) + } + x := benchVector(b, n) + for b.Loop() { + benchSinkArr, _ = csr.MatVec(x) + } + }) + } +} + +func BenchmarkSparseCSCMatVec(b *testing.B) { + for _, n := range []int{256, 131072} { + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + coo := benchCOO(b, n, n, laplacianTriples(n, 1)) + csc, err := CSCFromCOO(coo) + if err != nil { + b.Fatalf("CSCFromCOO: %v", err) + } + x := benchVector(b, n) + for b.Loop() { + benchSinkArr, _ = csc.MatVec(x) + } + }) + } +} + +func BenchmarkSparseCSRMatMulSparse(b *testing.B) { + const n = 256 + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + coo := benchCOO(b, n, n, laplacianTriples(n, 4)) + csr, err := CSRFromCOO(coo) + if err != nil { + b.Fatalf("CSRFromCOO: %v", err) + } + for b.Loop() { + benchSinkCSR, _ = csr.MatMulSparse(csr) + } + }) +} + +func BenchmarkSparseCSRMatMulDense(b *testing.B) { + const n = 512 + const cols = 4 + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + coo := benchCOO(b, n, n, laplacianTriples(n, 2)) + csr, err := CSRFromCOO(coo) + if err != nil { + b.Fatalf("CSRFromCOO: %v", err) + } + v := make([]float64, n*cols) + for i := range v { + v[i] = float64((i*11)%17)/8 - 1 + } + x, err := core.FromFloats(v, n, cols) + if err != nil { + b.Fatalf("FromFloats: %v", err) + } + for b.Loop() { + benchSinkArr, _ = csr.MatMulDense(x) + } + }) +} + +func BenchmarkCSRFromCOODuplicates(b *testing.B) { + const n = 1500 + tri := duplicateTriples(n, 6000) + b.Run("nnz=6000", func(b *testing.B) { + coo := benchCOO(b, n, n, tri) + for b.Loop() { + benchSinkCSR, _ = CSRFromCOO(coo) + } + }) +} + +func BenchmarkCSCFromCOODuplicates(b *testing.B) { + const n = 1500 + tri := duplicateTriples(n, 6000) + b.Run("nnz=6000", func(b *testing.B) { + coo := benchCOO(b, n, n, tri) + for b.Loop() { + benchSinkCSC, _ = CSCFromCOO(coo) + } + }) +} + +// BenchmarkSpEigenBanded measures the real symmetric Krylov solve end +// to end: the coordinate construction, the banded matrix-vector +// products and the full reorthogonalisation of every Lanczos step. +func BenchmarkSpEigenBanded(b *testing.B) { + const n = 800 + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + coo := benchCOO(b, n, n, laplacianTriples(n, 1)) + for b.Loop() { + benchSinkArr, _, _ = SpEigen(coo, 6, core.NewGenerator(7)) + } + }) +} + +// BenchmarkSpEigenGeneralBanded measures the Arnoldi solve on the same +// matrix: one explicit orthogonalisation per basis column instead of +// the three-term recurrence. +func BenchmarkSpEigenGeneralBanded(b *testing.B) { + const n = 400 + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + coo := benchCOO(b, n, n, laplacianTriples(n, 1)) + for b.Loop() { + benchSinkArr, _, _ = SpEigenGeneral(coo, 4, core.NewGenerator(7)) + } + }) +} + +func BenchmarkSpEigenComplexBanded(b *testing.B) { + const n = 400 + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + idx := make([]int64, 0, n*6) + vals := make([]complex128, 0, n*3) + for i := range n { + idx = append(idx, int64(i), int64(i)) + vals = append(vals, 2+0i) + if i+1 < n { + idx = append(idx, int64(i), int64(i+1)) + vals = append(vals, 1i) + idx = append(idx, int64(i+1), int64(i)) + vals = append(vals, -1i) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + b.Fatalf("FromInts: %v", err) + } + values, err := core.FromComplexes(vals, len(vals)) + if err != nil { + b.Fatalf("FromComplexes: %v", err) + } + coo, err := core.NewSparseCOO(indices, values, []int{n, n}) + if err != nil { + b.Fatalf("NewSparseCOO: %v", err) + } + for b.Loop() { + benchSinkArr, _, _ = SpEigenComplex(coo, 4, core.NewGenerator(7)) + } + }) +} + +// BenchmarkSpExpApplyBanded measures the Krylov projection of the +// matrix exponential action, which shares the Lanczos kernel with the +// symmetric eigensolver. +func BenchmarkSpExpApplyBanded(b *testing.B) { + const n = 400 + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + coo := benchCOO(b, n, n, laplacianTriples(n, 1)) + v := benchVector(b, n) + for b.Loop() { + benchSinkArr, _ = SpExpApply(coo, v, 0) + } + }) +} + +// BenchmarkSpSolveComplexCGBanded measures the conjugate-gradient solve +// whose every step is a complex sparse matrix-vector product with a +// Jacobi preconditioner. +func BenchmarkSpSolveComplexCGBanded(b *testing.B) { + const n = 256 + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + idx := make([]int64, 0, n*6) + vals := make([]complex128, 0, n*3) + for i := range n { + idx = append(idx, int64(i), int64(i)) + vals = append(vals, 4+0i) + if i+1 < n { + idx = append(idx, int64(i), int64(i+1)) + vals = append(vals, 1i) + idx = append(idx, int64(i+1), int64(i)) + vals = append(vals, -1i) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + b.Fatalf("FromInts: %v", err) + } + values, err := core.FromComplexes(vals, len(vals)) + if err != nil { + b.Fatalf("FromComplexes: %v", err) + } + coo, err := core.NewSparseCOO(indices, values, []int{n, n}) + if err != nil { + b.Fatalf("NewSparseCOO: %v", err) + } + ones := make([]complex128, n) + for i := range ones { + ones[i] = 1 + } + rhs, err := core.FromComplexes(ones, n) + if err != nil { + b.Fatalf("FromComplexes: %v", err) + } + for b.Loop() { + benchSinkArr, _ = SpSolveComplexCG(coo, rhs, 1e-12, 400) + } + }) +} diff --git a/linalg/cholupdate.go b/linalg/cholupdate.go new file mode 100644 index 0000000..9808250 --- /dev/null +++ b/linalg/cholupdate.go @@ -0,0 +1,120 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// Rank-one modification of a Cholesky factorisation. Refactorising +// from scratch costs a full O(n³) pass; a rank-one change of the +// matrix only needs one sweep of rotations over the factor, O(n²), +// which is what repeated posterior re-evaluation, sliding-window +// covariances and active-set methods live on. + +// CholeskyUpdate returns the lower Cholesky factor of A + x·xᵀ, given +// l, the lower factor of A. The factor is rebuilt by orthogonal +// rotations applied to l with x appended as an extra row, so the +// result is exact to rounding without a fresh factorisation. A rank-1 +// factor, a mismatched vector or a complex input is an error. +func CholeskyUpdate(l, x *core.Array) (*core.Array, error) { + return cholRankOne(l, x, "CholeskyUpdate", +1) +} + +// CholeskyDowndate returns the lower Cholesky factor of A − x·xᵀ, +// given l, the lower factor of A. The sweep runs hyperbolic rotations +// down the factor, each one squaring the remaining diagonal against +// the vector: when a diagonal can no longer dominate, the modified +// matrix has left the positive definite cone and the downdate is an +// error rather than a silently broken factor. +func CholeskyDowndate(l, x *core.Array) (*core.Array, error) { + return cholRankOne(l, x, "CholeskyDowndate", -1) +} + +// cholRankOne carries the shared sweep. The augmented matrix +// [l; xᵀ] holds the updated Gram matrix for +1; for −1 the vector row +// enters with a minus sign and only the hyperbolic rotations that keep +// l·lᵀ − x·xᵀ invariant may remove it, which is where the positive +// definiteness test lives. +func cholRankOne(l, x *core.Array, name string, sign int) (*core.Array, error) { + if l.Dtype() == core.Complex || x.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + if l.NDim() != 2 || l.Shape()[0] != l.Shape()[1] { + return nil, base.Errf("%s: the factor must be a square 2-D matrix, got shape %s", + name, base.ShapeText(l.Shape())) + } + n := l.Shape()[0] + if x.NDim() != 1 || x.Len() != n { + return nil, base.Errf("%s: the vector must have length %d, got shape %s", + name, n, base.ShapeText(x.Shape())) + } + out := core.New(core.Float, []int{n, n}...) + if l.Dtype() == core.Float && !l.Strided() && len(l.RawFloats()) == n*n { + copy(out.RawFloats(), l.RawFloats()) + } else { + for i := range n * n { + out.RawFloats()[i] = l.FloatAt(i) + } + } + v := make([]float64, n) + if x.Dtype() == core.Float && !x.Strided() && len(x.RawFloats()) == n { + copy(v, x.RawFloats()) + } else { + for i := range v { + v[i] = x.FloatAt(i) + } + } + // The rotations below assume a lower triangular factor; a nonzero + // strict upper triangle would silently corrupt the sweep, so it is + // refused up front. + for i := range n { + for j := i + 1; j < n; j++ { + if out.RawFloats()[i*n+j] != 0 { + return nil, base.Errf("%s: the factor must be lower triangular (nonzero entry at row %d, column %d)", name, i, j) + } + } + } + // The modified matrix reads [L, x]·J·[L, x]ᵀ with J = diag(I, ±1): + // one sweep of (hyperbolic for −1, orthogonal for +1) rotations + // eliminates the vector column, and the surviving columns are the + // new lower factor. Each rotation squares the running diagonal + // against the vector, which is where a downdate leaves the + // positive definite cone. + for k := range n { + d := out.RawFloats()[k*n+k] + if d <= 0 { + return nil, base.Errf("%s: the factor must be lower triangular with a positive diagonal", name) + } + var r, c, s float64 + if sign > 0 { + r = math.Hypot(d, v[k]) + c, s = d/r, v[k]/r + } else { + r2 := d*d - v[k]*v[k] + if math.IsNaN(r2) || math.IsInf(r2, 0) { + return nil, base.Errf("%s: the downdate at row %d overflows the float64 range", name, k) + } + if r2 <= 0 { + return nil, base.Errf("%s: the modified matrix is not positive definite at row %d", name, k) + } + r = math.Sqrt(r2) + c, s = d/r, v[k]/r + } + out.RawFloats()[k*n+k] = r + for i := k + 1; i < n; i++ { + lik, xi := out.RawFloats()[i*n+k], v[i] + if sign > 0 { + out.RawFloats()[i*n+k] = c*lik + s*xi + } else { + out.RawFloats()[i*n+k] = c*lik - s*xi + } + v[i] = c*xi - s*lik + } + } + return out, nil +} diff --git a/linalg/cholupdate_test.go b/linalg/cholupdate_test.go new file mode 100644 index 0000000..762e6a4 --- /dev/null +++ b/linalg/cholupdate_test.go @@ -0,0 +1,158 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// spdSample builds a deterministic symmetric positive definite matrix: +// B + Bᵀ + n·I from a fixed entry pattern, so no test randomness leaks. +func spdSample(n int) *core.Array { + raw := make([]float64, n*n) + for i := range n { + for j := range n { + raw[i*n+j] = math.Sin(float64(3*i+j+1)) + math.Cos(float64(i-j)) + } + } + vals := make([]float64, n*n) + for i := range n { + for j := range n { + vals[i*n+j] = raw[i*n+j] + raw[j*n+i] + } + vals[i*n+i] += float64(n) // shift onto the PD cone + } + a, _ := core.FromFloats(vals, n, n) + return a +} + +// gramReconstruct multiplies a lower triangle out: returns L·Lᵀ. +func gramReconstruct(t *testing.T, l *core.Array) *core.Array { + t.Helper() + n := l.Shape()[0] + out, err := zeros(core.Float, []int{n, n}) + if err != nil { + t.Fatalf("gramReconstruct: %v", err) + } + for i := range n { + for j := range i + 1 { + s := 0.0 + for k := range n { + s += l.FloatAt(i*n+k) * l.FloatAt(j*n+k) + } + out.SetFloatAt(i*n+j, s) + out.SetFloatAt(j*n+i, s) + } + } + return out +} + +func matDiff(a, b *core.Array) float64 { + worst := 0.0 + for i := range a.Len() { + d := math.Abs(a.FloatAt(i) - b.FloatAt(i)) + if d > worst { + worst = d + } + } + return worst +} + +// TestCholeskyUpdateGram pins the defining property: the rebuilt +// factor's Gram matrix equals A + x·xᵀ, checked against a fresh +// factorisation of the modified matrix. +func TestCholeskyUpdateGram(t *testing.T) { + n := 5 + a := spdSample(n) + l, err := Cholesky(a) + if err != nil { + t.Fatalf("Cholesky: %v", err) + } + x := mustFloats(t, []float64{1, -2, 0.5, 3, -1}, n) + updated, err := CholeskyUpdate(l, x) + if err != nil { + t.Fatalf("CholeskyUpdate: %v", err) + } + want, err := zeros(core.Float, []int{n, n}) + if err != nil { + t.Fatalf("zeros: %v", err) + } + for i := range n { + for j := range n { + want.SetFloatAt(i*n+j, a.FloatAt(i*n+j)+x.FloatAt(i)*x.FloatAt(j)) + } + } + got := gramReconstruct(t, updated) + if matDiff(got, want) > 1e-10 { + t.Fatalf("Gram mismatch %.3g", matDiff(got, want)) + } + reference, err := Cholesky(want) + if err != nil { + t.Fatalf("Cholesky of the update: %v", err) + } + if matDiff(updated, reference) > 1e-8 { + t.Fatalf("factors disagree %.3g", matDiff(updated, reference)) + } +} + +// TestCholeskyUpdateDowndateRoundTrip updates and downdates the same +// vector: the original factor must come back. +func TestCholeskyUpdateDowndateRoundTrip(t *testing.T) { + n := 6 + a := spdSample(n) + l, err := Cholesky(a) + if err != nil { + t.Fatalf("Cholesky: %v", err) + } + x := mustFloats(t, []float64{0.3, 1, -0.7, 2, -1, 0.2}, n) + up, err := CholeskyUpdate(l, x) + if err != nil { + t.Fatalf("CholeskyUpdate: %v", err) + } + back, err := CholeskyDowndate(up, x) + if err != nil { + t.Fatalf("CholeskyDowndate: %v", err) + } + if matDiff(back, l) > 1e-8 { + t.Fatalf("round trip lost %.3g", matDiff(back, l)) + } +} + +// TestCholeskyDowndateOutsideCone pins the honest failure: removing +// more than the matrix carries leaves the positive definite cone and +// the downdate must refuse. +func TestCholeskyDowndateOutsideCone(t *testing.T) { + a := spdSample(3) + l, err := Cholesky(a) + if err != nil { + t.Fatalf("Cholesky: %v", err) + } + big := mustFloats(t, []float64{10, 10, 10}, 3) + if _, err := CholeskyDowndate(l, big); err == nil { + t.Fatal("expected an error for a downdate outside the cone") + } + // A vector the matrix can absorb must succeed. + small := mustFloats(t, []float64{0.1, 0.1, 0.1}, 3) + if _, err := CholeskyDowndate(l, small); err != nil { + t.Fatalf("small downdate: %v", err) + } +} + +// TestCholeskyRankOneErrors pins shape and dtype validation. +func TestCholeskyRankOneErrors(t *testing.T) { + l, _ := Cholesky(spdSample(3)) + bad := mustFloats(t, []float64{1, 2}, 2) + if _, err := CholeskyUpdate(l, bad); err == nil { + t.Fatal("expected an error for a mismatched vector") + } + if _, err := CholeskyUpdate(bad, bad); err == nil { + t.Fatal("expected an error for a non-square factor") + } + cx, _ := core.FromComplexes([]complex128{1}, 1) + if _, err := CholeskyUpdate(l, cx); err == nil { + t.Fatal("expected an error for a complex vector") + } +} diff --git a/linalg/complex_refusal_pins_test.go b/linalg/complex_refusal_pins_test.go new file mode 100644 index 0000000..a18ea1c --- /dev/null +++ b/linalg/complex_refusal_pins_test.go @@ -0,0 +1,170 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression tests: entry points that reached a +// complex array's float accessor (a panic, not an error), a +// preconditioner applied to a system of another dimension, and Krylov +// callbacks returning the wrong length. + +// mustComplex builds a complex array. +func mustComplex(t *testing.T, vals []complex128, shape ...int) *core.Array { + t.Helper() + a, err := core.FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + return a +} + +// TestComplexInputsAreRefused pins the dtype gate of the four entry +// points that read their arguments as real numbers. +func TestComplexInputsAreRefused(t *testing.T) { + c2 := mustComplex(t, []complex128{1 + 1i, 1 + 1i}, 2) + c3 := mustComplex(t, []complex128{1 + 1i, 2 + 2i, 3 + 3i}, 3) + t.Run("SolveTridiagonal", func(t *testing.T) { + if _, err := SolveTridiagonal(c2, c3, c2, c3); err == nil { + t.Fatal("expected an error for complex diagonals") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-dtype refusal", err) + } + }) + t.Run("SolveCyclicTridiagonal", func(t *testing.T) { + if _, err := SolveCyclicTridiagonal(c3, c3, c3, c3); err == nil { + t.Fatal("expected an error for complex diagonals") + } + }) + t.Run("FitPolynomial", func(t *testing.T) { + if _, err := FitPolynomial(c2, c2, 1); err == nil { + t.Fatal("expected an error for complex samples") + } + }) + t.Run("NewCubicSpline", func(t *testing.T) { + if _, err := NewCubicSpline(c3, c3); err == nil { + t.Fatal("expected an error for complex knots") + } + }) +} + +// TestSparseILUDimensionMismatch pins the preconditioner check: an ILU +// built for another system must be refused rather than indexed. +func TestSparseILUDimensionMismatch(t *testing.T) { + mk := func(n int, vals []float64) *core.SparseCOO { + t.Helper() + idx := make([]int64, 0, n*2) + for i := range n { + idx = append(idx, int64(i), int64(i)) + } + ia, err := core.FromInts(idx, n, 2) + if err != nil { + t.Fatal(err) + } + va, err := core.FromFloats(vals, n) + if err != nil { + t.Fatal(err) + } + sp, err := core.NewSparseCOO(ia, va, []int{n, n}) + if err != nil { + t.Fatal(err) + } + return sp + } + big := mk(4, []float64{4, 4, 4, 4}) + small := mk(3, []float64{3, 3, 3}) + ilu, err := NewSparseILU(small) + if err != nil { + t.Fatalf("NewSparseILU: %v", err) + } + b := mustFloats(t, []float64{1, 1, 1, 1}, 4) + if _, err := SpSolve(big, b, 0, 50, ilu); err == nil { + t.Fatal("expected an error for a preconditioner of the wrong dimension") + } else if !strings.Contains(err.Error(), "dimension") { + t.Fatalf("error = %v, want a dimension refusal", err) + } + if _, err := SpSolveBiCGSTAB(big, b, 0, 50, ilu); err == nil { + t.Fatal("expected an error from BiCGSTAB for a preconditioner of the wrong dimension") + } +} + +// TestGMRESCallbackLength pins the op contract: a callback that +// returns fewer elements than the problem has must be an error. +func TestGMRESCallbackLength(t *testing.T) { + b := mustFloats(t, []float64{1, 2}, 2) + short := mustFloats(t, []float64{1}, 1) + if _, err := GMRES(func(*core.Array) (*core.Array, error) { return short, nil }, b, 2, 5, 1e-12); err == nil { + t.Fatal("expected an error for an op returning the wrong length") + } else if !strings.Contains(err.Error(), "elements") { + t.Fatalf("error = %v, want a length refusal", err) + } +} + +// TestQRLargeScale pins the Householder reflector at scales where the +// raw squares of the old norm overflowed: every intermediate of the +// reflector is now O(1) in scaled units, so the reconstruction holds at +// any magnitude. +func TestQRLargeScale(t *testing.T) { + for _, scale := range []float64{1, 1e10, 1e100, 1.2589e154, 1e200, 1e300} { + vals := []float64{2, -1, 0.5, 1, 3, -2, 4, 1, 1} + for i := range vals { + vals[i] *= scale + } + a, err := core.FromFloats(vals, 3, 3) + if err != nil { + t.Fatal(err) + } + q, r, err := QR(a) + if err != nil { + t.Fatalf("scale %g: %v", scale, err) + } + maxA := 0.0 + for _, v := range vals { + maxA = math.Max(maxA, math.Abs(v)) + } + worst := 0.0 + for i := range 3 { + for j := range 3 { + acc := 0.0 + for k := range 3 { + acc += q.FloatAt(i*3+k) * r.FloatAt(k*3+j) + } + worst = math.Max(worst, math.Abs(acc-vals[i*3+j])/maxA) + } + } + if !(worst < 1e-12) { + t.Errorf("scale %g: |QR−A|/|A| = %g, want below 1e-12", scale, worst) + } + } +} + +// TestSparseILUNil pins the nil guard: an explicitly nil preconditioner +// is an error, not a dereference. +func TestSparseILUNil(t *testing.T) { + a, err := core.FromFloats([]float64{2, 3}, 2) + if err != nil { + t.Fatal(err) + } + idx, err := core.FromInts([]int64{0, 0, 1, 1}, 2, 2) + if err != nil { + t.Fatal(err) + } + sp, err := core.NewSparseCOO(idx, a, []int{2, 2}) + if err != nil { + t.Fatal(err) + } + b, err := core.FromFloats([]float64{1, 1}, 2) + if err != nil { + t.Fatal(err) + } + if _, err := SpSolve(sp, b, 0, 10, nil); err == nil { + t.Fatal("expected an error for a nil preconditioner") + } +} diff --git a/linalg/csvd.go b/linalg/csvd.go new file mode 100644 index 0000000..1679aa4 --- /dev/null +++ b/linalg/csvd.go @@ -0,0 +1,315 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// Direct complex singular value decomposition by Golub-Kahan +// bidiagonalisation. Reducing A through Aᴴ·A squares the condition +// number, so a tiny singular value of A carries the accuracy of the +// square of a tiny eigenvalue; the direct route never forms the +// Gramian. core.Complex Householder reflectors chosen with real beta leave +// a real bidiagonal matrix, and the Golub-Reinsch shifted QR iteration +// then works on real numbers only, folding its rotations back into the +// complex unitary factors. + +// svdBidiagonalise reduces the complex m×n matrix a (m ≥ n) to upper +// bidiagonal form with real diagonal d and superdiagonal e, +// accumulating the left unitary into u (m×m) and the right one into v +// (n×n), both seeded with the identity. Real beta is the choice that +// keeps both bidiagonal bands real: the reflector maps its head to a +// real multiple of the first coordinate whatever the head's phase. +func svdBidiagonalise(a []complex128, m, n int) (d, e []float64, u, v []complex128) { + u = make([]complex128, m*m) + v = make([]complex128, n*n) + for i := range m { + u[i*m+i] = 1 + } + for i := range n { + v[i*n+i] = 1 + } + d = make([]float64, n) + e = make([]float64, n) + // Reflector scratch reused across the sweep; each step uses the + // prefix it needs. + vecBuf := make([]complex128, m) + vecRBuf := make([]complex128, n) + for k := range n { + // Left reflector on column k, rows k..m-1. + norm := 0.0 + for i := k; i < m; i++ { + norm += base.Real2(a[i*n+k]) + } + norm = math.Sqrt(norm) + if norm > 0 { + beta := -norm + if real(a[k*n+k]) < 0 { + beta = norm + } + vh := 0.0 + vec := vecBuf[:m-k] + for i := k; i < m; i++ { + vec[i-k] = a[i*n+k] + if i == k { + vec[0] -= complex(beta, 0) + } + vh += base.Real2(vec[i-k]) + } + if vh > 0 { + vhx := complex(0, 0) + for i := k; i < m; i++ { + vhx += cmplx.Conj(vec[i-k]) * a[i*n+k] + } + tau := complex(1, 0) / vhx + for j := k; j < n; j++ { + s := complex(0, 0) + for i := k; i < m; i++ { + s += cmplx.Conj(vec[i-k]) * a[i*n+j] + } + s *= tau + for i := k; i < m; i++ { + a[i*n+j] -= s * vec[i-k] + } + } + for r := range m { + s := complex(0, 0) + for j := k; j < m; j++ { + s += u[r*m+j] * vec[j-k] + } + s *= cmplx.Conj(tau) + for j := k; j < m; j++ { + u[r*m+j] -= s * cmplx.Conj(vec[j-k]) + } + } + } + } + d[k] = real(a[k*n+k]) + a[k*n+k] = complex(d[k], 0) + for i := k + 1; i < m; i++ { + a[i*n+k] = 0 + } + if k == n-1 { + break + } + // Right reflector on row k, columns k+1..n-1. + norm = 0.0 + for j := k + 1; j < n; j++ { + norm += base.Real2(a[k*n+j]) + } + norm = math.Sqrt(norm) + if norm > 0 { + beta := -norm + if real(a[k*n+k+1]) < 0 { + beta = norm + } + vh := 0.0 + vecR := vecRBuf[:n-k-1] + for j := k + 1; j < n; j++ { + vecR[j-k-1] = a[k*n+j] + if j == k+1 { + vecR[0] -= complex(beta, 0) + } + vh += base.Real2(vecR[j-k-1]) + } + if vh > 0 { + vhx := complex(0, 0) + for j := k + 1; j < n; j++ { + vhx += cmplx.Conj(vecR[j-k-1]) * a[k*n+j] + } + tau := complex(1, 0) / vhx + for i := range m { + s := complex(0, 0) + for j := k + 1; j < n; j++ { + s += a[i*n+j] * cmplx.Conj(vecR[j-k-1]) + } + s *= tau + for j := k + 1; j < n; j++ { + a[i*n+j] -= s * vecR[j-k-1] + } + } + for r := range n { + s := complex(0, 0) + for j := k + 1; j < n; j++ { + s += v[r*n+j] * cmplx.Conj(vecR[j-k-1]) + } + s *= tau + for j := k + 1; j < n; j++ { + v[r*n+j] -= s * vecR[j-k-1] + } + } + } + e[k] = real(a[k*n+k+1]) + a[k*n+k+1] = complex(e[k], 0) + for j := k + 2; j < n; j++ { + a[k*n+j] = 0 + } + } + } + return d, e, u, v +} + +// svdGolubReinsch diagonalises the real bidiagonal pair (d, e) in +// place, folding every rotation into the columns of u (m×m) and v +// (n×n). Each sweep is one implicit shifted QR step on BᵀB: a +// Wilkinson shift from the trailing two-by-two opens the sweep, and +// the bulge it raises is chased off the block's end by alternating +// left and right rotations. A zero leading diagonal deflates through a +// rotation chase that empties its column, which is how rank deficiency +// leaves the iteration honestly. +func svdGolubReinsch(d, e []float64, u, v []complex128, m, n int) error { + const name = "SVDComplex" + // Deflation floor relative to the bidiagonal norm, the same + // backward-stability floor symmetricQr uses: a purely + // neighbour-relative threshold never triggers when both diagonals + // sit at rounding-zero, which the FMA contraction of the vector + // build leaves the chase on, and the iteration then exhausts + // itself over an already negligible block. + scale := 0.0 + for i := range n { + scale += d[i] * d[i] + if i+1 < n { + scale += 2 * e[i] * e[i] + } + } + tolAbs := base.EpsF * math.Sqrt(scale) + negligible := func(i int) bool { + a := math.Abs(e[i]) + return a == 0 || + a <= base.EpsF*(math.Abs(d[i])+math.Abs(d[i+1])) || + a < tolAbs + } + for iter := 0; iter < 60*n+100; iter++ { + // Retire negligible superdiagonal entries. + for i := range n - 1 { + if negligible(i) { + e[i] = 0 + } + } + // Locate the trailing unreduced block [l..r]. + r := -1 + for i := n - 2; i >= 0; i-- { + if e[i] != 0 { + r = i + 1 + break + } + } + if r < 0 { + return nil // fully diagonal + } + l := r - 1 + for l > 0 && e[l-1] != 0 { + l-- + } + if d[l] == 0 || math.Abs(d[l]) < tolAbs { + // The block's leading column is zero, or sits at the + // rounding floor of the whole bidiagonal: the FMA + // contraction of the vector build leaves a denormal there + // where the portable build hits the exact zero, and a + // shifted chase over such a block cycles without + // deflating. Left rotations between row l and the rows + // below push the superdiagonal e[l] through the block and + // off it, deflating the negligible singular value without + // touching the rest. + w := e[l] + for k := l + 1; k <= r; k++ { + h := math.Hypot(w, d[k]) + c, s := d[k]/h, -w/h + d[k] = h + if k < r { + w = s * e[k] + e[k] *= c + } + for row := range m { + ul, uk := u[row*m+l], u[row*m+k] + u[row*m+l] = complex(c, 0)*ul + complex(s, 0)*uk + u[row*m+k] = complex(-s, 0)*ul + complex(c, 0)*uk + } + } + e[l] = 0 + continue + } + // Wilkinson shift from the trailing two-by-two of BᵀB. + t11 := d[r-1] * d[r-1] + if r-2 >= 0 { + t11 += e[r-2] * e[r-2] + } + t22 := d[r]*d[r] + e[r-1]*e[r-1] + t21 := d[r-1] * e[r-1] + delta := (t11 - t22) / 2 + denom := math.Abs(delta) + math.Sqrt(delta*delta+t21*t21) + mu := t22 + if denom > 0 { + sign := 1.0 + if delta < 0 { + sign = -1.0 + } + mu = t22 - sign*t21*t21/denom + } + // Opening right rotation on columns (l, l+1), taken from the + // first column of T − μI. + g0 := d[l]*d[l] - mu + h0 := math.Hypot(g0, d[l]*e[l]) + c0, s0 := g0/h0, d[l]*e[l]/h0 + dl, el, dl1 := d[l], e[l], d[l+1] + d[l] = c0*dl + s0*el + e[l] = -s0*dl + c0*el + d[l+1] = c0 * dl1 + bulge := s0 * dl1 // at (l+1, l) + for row := range n { + vl, vl1 := v[row*n+l], v[row*n+l+1] + v[row*n+l] = complex(c0, 0)*vl + complex(s0, 0)*vl1 + v[row*n+l+1] = complex(-s0, 0)*vl + complex(c0, 0)*vl1 + } + // Chase: the left rotation kills the bulge and raises a + // super-bulge; the right rotation kills that and re-raises the + // bulge one step down, until it falls off the block's end. + for k := l; k < r; k++ { + // Left rotation on rows (k, k+1). A deflated diagonal meets + // a deflated bulge with h = 0: the rotation is the identity + // there, where the division would raise 0/0 NaNs that spill + // through the factor. + h := math.Hypot(d[k], bulge) + c1, s1 := 1.0, 0.0 + if h > 0 { + c1, s1 = d[k]/h, bulge/h + } + ekOld, dk1Old := e[k], d[k+1] + d[k] = h + e[k] = c1*ekOld + s1*dk1Old + d[k+1] = -s1*ekOld + c1*dk1Old + var sb float64 + if k < r-1 { + sb = s1 * e[k+1] + e[k+1] *= c1 + } + for row := range m { + uk, uk1 := u[row*m+k], u[row*m+k+1] + u[row*m+k] = complex(c1, 0)*uk + complex(s1, 0)*uk1 + u[row*m+k+1] = complex(-s1, 0)*uk + complex(c1, 0)*uk1 + } + if k >= r-1 { + break + } + // Right rotation on columns (k+1, k+2). + h2 := math.Hypot(e[k], sb) + c2, s2 := e[k]/h2, sb/h2 + ek1Old, dk1Old, dk2Old := e[k+1], d[k+1], d[k+2] + e[k] = h2 + d[k+1] = c2*dk1Old + s2*ek1Old + e[k+1] = -s2*dk1Old + c2*ek1Old + bulge = s2 * dk2Old + d[k+2] = c2 * dk2Old + for row := range n { + vk1, vk2 := v[row*n+k+1], v[row*n+k+2] + v[row*n+k+1] = complex(c2, 0)*vk1 + complex(s2, 0)*vk2 + v[row*n+k+2] = complex(-s2, 0)*vk1 + complex(c2, 0)*vk2 + } + } + } + return base.Errf("SVDComplex: the bidiagonal iteration did not converge") +} diff --git a/linalg/csvd_test.go b/linalg/csvd_test.go new file mode 100644 index 0000000..8fe92ef --- /dev/null +++ b/linalg/csvd_test.go @@ -0,0 +1,191 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// csvdSample builds a deterministic complex m×n matrix. +func csvdSample(m, n int) *core.Array { + vals := make([]complex128, m*n) + for i := range m * n { + vals[i] = complex(math.Sin(float64(3*i+1)), math.Cos(float64(2*i+1))) + } + a, _ := core.FromComplexes(vals, m, n) + return a +} + +// csvdPlane returns the m×m unitary that rotates coordinates i, j by +// angle theta with phase phi, the building block for test unitaries +// with known spectra. +func csvdPlane(m, i, j int, theta, phi float64) []complex128 { + q := make([]complex128, m*m) + for k := range m { + q[k*m+k] = 1 + } + q[i*m+i] = complex(math.Cos(theta), 0) + q[j*m+j] = complex(math.Cos(theta), 0) + q[i*m+j] = complex(math.Sin(theta)*math.Cos(phi), math.Sin(theta)*math.Sin(phi)) + q[j*m+i] = complex(-math.Sin(theta)*math.Cos(phi), math.Sin(theta)*math.Sin(phi)) + return q +} + +// TestSVDComplexContracts pins the decomposition contract on a spread +// of deterministic shapes: reconstruction, unitarity of both factors, +// descending non-negative singular values. +func TestSVDComplexContracts(t *testing.T) { + cases := []struct{ m, n int }{{5, 3}, {4, 4}, {3, 1}, {2, 2}, {1, 1}, {6, 2}} + for _, tc := range cases { + a := csvdSample(tc.m, tc.n) + u, sigma, vh, err := SVDComplex(a) + if err != nil { + t.Fatalf("%dx%d: SVDComplex: %v", tc.m, tc.n, err) + } + r := min(tc.m, tc.n) + scale := 0.0 + for i := range a.Len() { + scale = math.Max(scale, cmplx.Abs(a.ComplexAt(i))) + } + // Reconstruction. + recon := 0.0 + for i := range tc.m { + for j := range tc.n { + s := complex(0, 0) + for k := range r { + s += u.ComplexAt(i*r+k) * complex(sigma.FloatAt(k), 0) * vh.ComplexAt(k*tc.n+j) + } + recon = math.Max(recon, cmplx.Abs(s-a.ComplexAt(i*tc.n+j))) + } + } + if recon > 1e-10*math.Max(1, scale) { + t.Fatalf("%dx%d: reconstruction error %.3g", tc.m, tc.n, recon) + } + // Unitarity of U's columns and of V. + gram := func(get func(i, j int) complex128, rows, cols int) float64 { + worst := 0.0 + for i := range cols { + for j := range cols { + s := complex(0, 0) + for l := range rows { + s += cmplx.Conj(get(l, i)) * get(l, j) + } + want := 0.0 + if i == j { + want = 1 + } + worst = math.Max(worst, math.Abs(cmplx.Abs(s)-want)) + } + } + return worst + } + if g := gram(func(i, j int) complex128 { return u.ComplexAt(i*r + j) }, tc.m, r); g > 1e-10 { + t.Fatalf("%dx%d: U not orthonormal, error %.3g", tc.m, tc.n, g) + } + if g := gram(func(i, j int) complex128 { return vh.ComplexAt(i*tc.n + j) }, tc.n, tc.n); g > 1e-10 { + t.Fatalf("%dx%d: Vᴴ not unitary, error %.3g", tc.m, tc.n, g) + } + for k := range sigma.Len() { + if sigma.FloatAt(k) < 0 { + t.Fatalf("%dx%d: negative singular value %g", tc.m, tc.n, sigma.FloatAt(k)) + } + if k > 0 && sigma.FloatAt(k) > sigma.FloatAt(k-1)+1e-12 { + t.Fatalf("%dx%d: singular values not descending", tc.m, tc.n) + } + } + } +} + +// TestSVDComplexIllConditioned is the reason the direct route exists: +// with a known spectrum 1, 1e-8, 1e-16 the small singular values keep +// their relative accuracy, where the squared-condition AᴴA route would +// lose half the digits. +func TestSVDComplexIllConditioned(t *testing.T) { + const m, n = 3, 3 + sigma := []float64{1, 1e-8, 1e-16} + // A = P·diag(σ)·Qᴴ for two deterministic complex unitaries P, Q. + p := csvdPlane(m, 0, 1, 0.7, 1.1) + q := csvdPlane(n, 1, 2, 1.3, 0.4) + _ = q + pq := csvdPlane(m, 1, 2, 0.5, 2.2) + // Compose P = p·pq. + pMat := make([]complex128, m*m) + for i := range m { + for j := range m { + s := complex(0, 0) + for k := range m { + s += p[i*m+k] * pq[k*m+j] + } + pMat[i*m+j] = s + } + } + vals := make([]complex128, m*n) + for i := range m { + for j := range n { + s := complex(0, 0) + for k := range m { + s += pMat[i*m+k] * complex(sigma[k], 0) * cmplx.Conj(q[j*n+k]) + } + vals[i*n+j] = s + } + } + a, err := core.FromComplexes(vals, m, n) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + _, sigmaOut, _, err := SVDComplex(a) + if err != nil { + t.Fatalf("SVDComplex: %v", err) + } + for k := range n { + want := sigma[k] + got := sigmaOut.FloatAt(k) + if k < 2 { + if rel := math.Abs(got-want) / want; rel > 1e-9 { + t.Fatalf("σ%d = %.17g, want %.17g (relative error %.3g)", k, got, want, rel) + } + } else { + // At the round-off floor the honest guarantee is absolute: + // the direct route pins σ to eps·σ_max, the squared route + // could not. + if math.Abs(got-want) > 1e-15*sigma[0] { + t.Fatalf("σ%d = %.17g, want %.17g (absolute error %.3g)", + k, got, want, math.Abs(got-want)) + } + } + } +} + +// TestSVDComplexMatchesReal cross-checks the complex solver against +// the independent real SVD on a real matrix embedded in complex. +func TestSVDComplexMatchesReal(t *testing.T) { + realPart := mustFloats(t, []float64{ + 3, 0, 1, + 1, 2, 1, + 1, 1, 2, + 0, 1, 4, + }, 4, 3) + uR, sigmaR, _, err := SVD(realPart) + if err != nil { + t.Fatalf("SVD: %v", err) + } + _ = uR + vals := make([]complex128, 12) + for i := range 12 { + vals[i] = complex(realPart.FloatAt(i), 0) + } + a, _ := core.FromComplexes(vals, 4, 3) + _, sigmaC, _, err := SVDComplex(a) + if err != nil { + t.Fatalf("SVDComplex: %v", err) + } + for k := range 3 { + if math.Abs(sigmaC.FloatAt(k)-sigmaR.FloatAt(k)) > 1e-10*math.Max(1, sigmaR.FloatAt(k)) { + t.Fatalf("σ%d: complex %.12g, real %.12g", k, sigmaC.FloatAt(k), sigmaR.FloatAt(k)) + } + } +} diff --git a/linalg/decomp.go b/linalg/decomp.go new file mode 100644 index 0000000..d84c61c --- /dev/null +++ b/linalg/decomp.go @@ -0,0 +1,596 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Matrix decompositions. All routines operate on float64 in +// and out; complex matrices are rejected. The QR routine uses +// Householder reflections, numerically stable and standard. Cholesky +// uses the Banachiewicz algorithm. LeastSquares builds on QR. + +// QR returns the QR decomposition a = Q * R of an m×n matrix a, with +// m ≥ n. Q is m×m orthogonal and R is m×n upper triangular. +func QR(a *core.Array) (q, r *core.Array, err error) { + if a.Dtype() == core.Complex { + return nil, nil, base.Errf("QR: complex matrices are not supported") + } + if a.NDim() != 2 { + return nil, nil, base.Errf("QR: needs a 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m < n { + return nil, nil, base.Errf("QR: needs m ≥ n, got shape %s", base.ShapeText(a.Shape())) + } + aMat := denseFloats(a, m, n) + qMat, rMat, err := qrRaw(aMat, m, n) + if err != nil { + return nil, nil, err + } + return floatsToArray(qMat, []int{m, m}), floatsToArray(rMat, []int{m, n}), nil +} + +// Cholesky returns the lower-triangular Cholesky factor L such that +// a = L * Lᵀ for a symmetric positive definite matrix a. +func Cholesky(a *core.Array) (*core.Array, error) { + if a.Dtype() == core.Complex { + return nil, base.Errf("Cholesky: complex matrices are not supported") + } + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, base.Errf("Cholesky: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + lMat, err := denseCholFactor(denseFloats(a, n, n), n) + if err != nil { + return nil, err + } + return floatsToArray(lMat, []int{n, n}), nil +} + +// denseCholBlock is the column-block width of the Cholesky sweep. A wider +// block shortens the serial diagonal factorisation and lengthens the +// parallel panel update's k run; 64 keeps the serial share near a +// twentieth of the arithmetic at the sizes this package serves while +// both panel passes stay inside a cache line's neighbourhood. +const denseCholBlock = 64 + +// denseCholMinWork is the element-touch floor one Cholesky panel worker must +// receive before the dispatch is worth its spawn cost, counted in +// update terms (rows times columns times the k run). Below it the panel +// runs on the calling goroutine. +const denseCholMinWork = 8192 + +// denseCholFactor factors a flat n×n row-major matrix into its lower +// triangular Cholesky factor, consuming the input. +// +// The sweep runs in column blocks of denseCholBlock: each block is first +// updated from the finished columns, then factored serially, then +// divided into the trapezoid below it. For an element (i, j) the three +// stages subtract the columns k < j in the one ascending order the +// unblocked dot always used, since the panel update covers exactly the +// k < j0 prefix and the factorisation or the divide covers the rest, +// so the running value the last stage divides is bit for bit the value +// the unblocked loop held. Only the order in which independent elements +// are visited changes, which is why the panel update and the trapezoid +// divide, both disjoint per row, run over a crew. +func denseCholFactor(mat []float64, n int) ([]float64, error) { + lMat := make([]float64, n*n) + for j0 := 0; j0 < n; j0 += denseCholBlock { + j1 := min(j0+denseCholBlock, n) + // Panel update: every element of rows j0..n−1 in columns + // j0..j1−1 loses the contribution of the finished columns k < j0. + // A row's elements depend on those columns alone, so rows split + // over the crew. + if j0 > 0 { + cols := j1 - j0 + panel := func(start, end int) { + for i := start; i < end; i++ { + for j := j0; j < j1 && j <= i; j++ { + s := mat[i*n+j] + for k := range j0 { + s -= lMat[i*n+k] * lMat[j*n+k] + } + mat[i*n+j] = s + } + } + } + if rows := n - j0; rows*cols*j0 >= denseCholMinWork { + engine.ParallelMin(rows, denseCholRows(denseCholMinWork, cols*j0), func(start, end int) { + panel(j0+start, j0+end) + }) + } else { + panel(j0, n) + } + } + // Diagonal block: columns in order, each finished before the + // next reads it. This is the only serial arithmetic left, and it + // shrinks with the block width: the block's own triangle is + // denseCholBlock³/6 of the matrix's n³/6. + for j := j0; j < j1; j++ { + for i := j; i < j1; i++ { + s := mat[i*n+j] + for k := j0; k < j; k++ { + s -= lMat[i*n+k] * lMat[j*n+k] + } + if i == j { + if s <= 0 { + return nil, base.Errf("Cholesky: matrix is not positive definite (pivot %d = %v)", i, s) + } + lMat[i*n+j] = math.Sqrt(s) + } else { + if lMat[j*n+j] == 0 { + return nil, base.Errf("Cholesky: zero diagonal at %d", j) + } + lMat[i*n+j] = s / lMat[j*n+j] + } + } + } + // Trapezoid below the block: row i of the panel needs row j of the + // diagonal block, which the sweep above has finished, and then its + // own earlier columns of the same block; both are in hand, so the + // rows divide independently. The blocks are disjoint per row, so + // one crew covers the divide. + if j1 < n { + divide := func(start, end int) { + for j := j0; j < j1; j++ { + d := lMat[j*n+j] + for i := start; i < end; i++ { + s := mat[i*n+j] + for k := j0; k < j; k++ { + s -= lMat[i*n+k] * lMat[j*n+k] + } + lMat[i*n+j] = s / d + } + } + } + if rows := n - j1; rows*(j1-j0)*(j1-j0)/2 >= denseCholMinWork { + engine.ParallelMin(rows, denseCholRows(denseCholMinWork, (j1-j0)*(j1-j0)/2), func(start, end int) { + divide(j1+start, j1+end) + }) + } else { + divide(j1, n) + } + } + } + return lMat, nil +} + +// denseCholRows returns the rows one worker needs for the given per-row +// cost to reach the dispatch floor, never below one. +func denseCholRows(floor, cost int) int { + if cost <= 0 { + return 1 + } + return max((floor+cost-1)/cost, 1) +} + +// LeastSquares solves Ax = b in the least-squares sense. a is m×n +// with m ≥ n. b is m or m×k. Implemented via QR. +func LeastSquares(a, b *core.Array) (*core.Array, error) { + if a.Dtype() == core.Complex { + return nil, base.Errf("LeastSquares: complex matrices are not supported") + } + // The QR route reads b through FloatAt as well: without the gate a + // complex right-hand side reaches the empty int payload and panics. + if b.Dtype() == core.Complex { + return nil, base.Errf("LeastSquares: complex right-hand sides are not supported") + } + if a.NDim() != 2 { + return nil, base.Errf("LeastSquares: 'a' must be 2-D, got shape %s", base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m < n { + return nil, base.Errf("LeastSquares: needs m ≥ n, got shape %s", base.ShapeText(a.Shape())) + } + if b.NDim() != 1 && b.NDim() != 2 { + return nil, base.Errf("LeastSquares: 'b' must be 1-D or 2-D, got shape %s", base.ShapeText(b.Shape())) + } + if b.Shape()[0] != m { + return nil, base.Errf("LeastSquares: b rows (%d) must match a rows (%d)", b.Shape()[0], m) + } + k := 1 + if b.NDim() == 2 { + k = b.Shape()[1] + } + aMat := denseFloats(a, m, n) + bMat := make([]float64, m*k) + if b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == m*k { + copy(bMat, b.RawFloats()) + } else if b.NDim() == 1 { + for i := range m { + bMat[i] = b.FloatAt(i) + } + } else { + for i := range m { + for j := range k { + bMat[i*k+j] = b.FloatAt(i*k + j) + } + } + } + // The reflectors go straight onto b; no Q is formed, because the + // only thing the solve wants from it is Qᵀ·b. + qtB, rMat := qrSolveRows(aMat, m, n, k, bMat) + x := make([]float64, n*k) + // Rank guard: with no column pivoting an exactly or nearly + // dependent column leaves a rounding-level R diagonal instead of an + // exact zero, which back-substitution would amplify into a huge + // silent answer. Pivots below the Householder-QR analogue of the + // SVD route's threshold, n·eps·max|R|, refuse the system. + rMax := 0.0 + for i := range n { + if v := math.Abs(rMat[i*n+i]); v > rMax { + rMax = v + } + } + rankTol := float64(n) * base.EpsF * rMax + for col := 0; col < k; col++ { + for i := n - 1; i >= 0; i-- { + s := qtB[i*k+col] + for j := i + 1; j < n; j++ { + s -= rMat[i*n+j] * x[j*k+col] + } + if math.Abs(rMat[i*n+i]) <= rankTol { + return nil, base.Errf("LeastSquares: rank-deficient system (R pivot %d = %g is at the rounding floor %g)", i, rMat[i*n+i], rankTol) + } + x[i*k+col] = s / rMat[i*n+i] + } + } + if b.NDim() == 1 { + return floatsToArray(x, []int{n}), nil + } + return floatsToArray(x, []int{n, k}), nil +} + +// householderReflector builds the reflector that zeroes the entries +// below the diagonal of column k of the row-major m×n matrix r and +// writes its scaled vector into v[:m−k]. It reports the vector length, +// beta = 2/(vᵀv) and whether a reflector was needed at all: a zero +// column, a zero norm or a zero vector leaves the column as the sweep +// wants it and the step is skipped. +// +// The scaling is the one the sweep has always used: entries near 1e154 +// would overflow the raw squares to +Inf, which zeroes beta and forms +// 0*Inf = NaN in the update, so the column is divided by its largest +// magnitude first and beta absorbs the square of that scale, leaving +// the update beta·(vᵀy)·v at the value the unscaled arithmetic would +// have had, with every intermediate O(1). +func householderReflector(r []float64, m, n, k int, v []float64) (ln int, beta float64, ok bool) { + scale := 0.0 + for i := k; i < m; i++ { + if a := math.Abs(r[i*n+k]); a > scale { + scale = a + } + } + if scale == 0 { + return 0, 0, false + } + sum := 0.0 + for i := k; i < m; i++ { + t := r[i*n+k] / scale + sum += t * t + } + norm := math.Sqrt(sum) + if norm == 0 { + return 0, 0, false + } + xk := r[k*n+k] / scale + s := -sign(xk) + if s == 0 { + s = -1 + } + u0 := xk - s*norm + vv := v[:m-k] + vv[0] = u0 + for i := 1; i < m-k; i++ { + vv[i] = r[(k+i)*n+k] / scale + } + vtv := u0 * u0 + for i := 1; i < m-k; i++ { + vtv += vv[i] * vv[i] + } + if vtv == 0 { + return 0, 0, false + } + return m - k, 2.0 / vtv, true +} + +// qrSolveRows runs the Householder sweep on a row-major m×n system and +// applies every reflector to the k right-hand sides as it goes, so the +// projection Qᵀ·b comes out without Q ever being formed: the sweep costs +// O(m·n²) for R and O(m·n·k) for the right-hand sides, where forming Q +// costs O(m²·n) and the dense product after it another O(m²·k). For the +// tall systems this serves, m ≫ n, that product is nearly the whole +// solve, and applying the reflectors is also the more accurate route: +// the projection stays orthogonal by construction instead of carrying +// the rounding of an explicitly accumulated Q. It returns the leading n +// rows of the transformed right-hand sides, concatenated by column, and +// R. +func qrSolveRows(aMat []float64, m, n, k int, rhs []float64) (qtB, rMat []float64) { + rMat = make([]float64, m*n) + copy(rMat, aMat) + y := make([]float64, m*k) + copy(y, rhs) + v := make([]float64, m) + for kk := range n { + ln, beta, ok := householderReflector(rMat, m, n, kk, v) + if !ok { + continue + } + vv := v[:ln] + // The right-hand sides first: each column's dot runs over i + // ascending and its update repeats the same operand order, and + // the columns are independent outputs, so a split over them + // cannot move a value. + if k*ln >= qrRHSMinWork { + engine.ParallelMin(k, 1, func(start, end int) { + for j := start; j < end; j++ { + applyReflectorToColumn(y, j, k, kk, ln, beta, vv) + } + }) + } else { + for j := range k { + applyReflectorToColumn(y, j, k, kk, ln, beta, vv) + } + } + // Then the trailing columns of R, blocked and dispatched exactly + // as the explicit-Q sweep does it, with the Q columns it also + // carries left out. + rBlocks := (n - kk + qrColumnBlock - 1) / qrColumnBlock + worker := func(gi, nw int) { + var d [qrColumnBlock]float64 + for b := gi * rBlocks / nw; b < (gi+1)*rBlocks/nw; b++ { + j0 := kk + b*qrColumnBlock + j1 := min(j0+qrColumnBlock, n) + wd := j1 - j0 + _ = d[wd-1] // bounds-check proof: the block never exceeds qrColumnBlock + clear(d[:wd]) + for jj := range wd { + jcol := j0 + jj + for i := range ln { + d[jj] += vv[i] * rMat[(kk+i)*n+jcol] + } + } + for j := range wd { + d[j] = beta * d[j] + } + for i := range ln { + vi := vv[i] + base := (kk+i)*n + j0 + for j := range wd { + rMat[base+j] -= vi * d[j] + } + } + } + } + weight := (n - kk) + k + w := min(weight*ln/qrWorkQuantum+1, engine.WorkersFor(rBlocks)) + if w < 2 { + worker(0, 1) + continue + } + var wg sync.WaitGroup + for gi := range w { + wg.Go(func() { + worker(gi, w) + }) + } + wg.Wait() + } + qtB = make([]float64, n*k) + for j := range k { + for i := range n { + qtB[i*k+j] = y[i*k+j] + } + } + return qtB, rMat +} + +// applyReflectorToColumn applies one reflector to column j of an m×k +// right-hand-side block: the dot runs over i ascending and the update +// repeats the same operand order. +func applyReflectorToColumn(y []float64, j, k, kk, ln int, beta float64, vv []float64) { + dot := 0.0 + base := kk * k + for i := range ln { + dot += vv[i] * y[base+i*k+j] + } + w := beta * dot + for i := range ln { + y[base+i*k+j] -= vv[i] * w + } +} + +// qrRHSMinWork is the element work one worker must carry before the +// right-hand-side columns of a reflector are split across goroutines. +const qrRHSMinWork = 1 << 15 + +// qrWorkQuantum is the element-touch budget one worker receives in a +// qrRaw reflector dispatch, counted as the reflector's reach (columns +// of R plus columns of Q) times the reflector length. A reflector whose +// application does not fill one quantum runs on the calling goroutine: +// the dispatch would cost more than the work. A fixed 32-way split +// measured 15 000 allocations and several milliseconds of scheduler +// time per 256² QR, because 32 goroutines spawn and synchronise per +// reflector however small the reflector's reach has grown; sizing the +// crew by the work keeps each dispatch's spawn cost proportional to the +// work it spreads. +const qrWorkQuantum = 8192 + +// qrColumnBlock is the number of R columns one blocked update pass +// covers. R is row-major, so a single column's update walks a stride-n +// pattern that touches one cache line per element; a block of columns +// walks the same rows linearly and reuses each line across the block. +// Eight columns span one 64-byte line of float64. +const qrColumnBlock = 8 + +// qrRaw performs Householder QR on a flat m×n matrix and returns the +// flat Q (m, m) and R (m, n). +// +// The reflectors apply one after another: their ORDER defines Q and R +// and is never changed. Within a single reflector every column of R +// and every column of Q is updated exactly once, reading only the +// frozen reflector v and beta, so the per-column work (dot along i +// ascending, then the same-operand rank-1 subtraction) is dispatched +// over the columns. Identical inputs, identical per-element sequence: +// the result is bit-identical to the serial sweep, whichever crew size +// or column blocking the dispatcher picks. +func qrRaw(aMat []float64, m, n int) (qMat, rMat []float64, err error) { + rMat = make([]float64, m*n) + copy(rMat, aMat) + qMat = eye(m) + // One reflector scratch reused across the sweep; each step uses its + // first m−k entries. + v := make([]float64, m) + // Per-reflector state the crew reads: the sweep writes it on the + // calling goroutine and the crew only reads it, so one body closure + // and its captures live for the whole sweep instead of a fresh pair + // of closures escaping to the heap per reflector. + var ( + k int // reflector's leading row and column + ln int // reflector length, m−k + beta float64 // 2/(vᵀv) of the current reflector + vv []float64 // reflector prefix view, re-sliced per step + ) + // worker runs the reflector's R columns and Q columns that fall to + // crew member gi of nw: a contiguous slab of the R blocks and a + // contiguous slab of the Q columns. Every member owns its own block + // scratch d, and slabs are disjoint, so no two members ever touch + // the same line; contiguous slabs keep each member's working set at + // its own rows of qMat and its own column stripe of rMat instead of + // walking the whole window. Each column's dot runs over i ascending + // and the subtraction repeats the same operand order, so the slab + // boundaries never move a bit. + worker := func(gi, nw int) { + rBlocks := (n - k + qrColumnBlock - 1) / qrColumnBlock + // R columns in blocks of qrColumnBlock. R is row-major, so a + // lone column's update walks a stride-n pattern that touches one + // cache line per element; a block keeps the hot window at one + // line per row. The per-column dot below keeps the exact shape + // the per-column sweep always had (same accumulator shape, same + // operand order), because reshaping the dot loop changes the + // compiler's fusing of the multiply into the add and that moves + // bits. Each column's w is beta·dot as before, and the blocked + // subtraction walks the same i order over the block linearly. + var d [qrColumnBlock]float64 + for b := gi * rBlocks / nw; b < (gi+1)*rBlocks/nw; b++ { + j0 := k + b*qrColumnBlock + j1 := min(j0+qrColumnBlock, n) + wd := j1 - j0 + _ = d[wd-1] // bounds-check proof: the block never exceeds qrColumnBlock + clear(d[:wd]) + for jj := range wd { + jcol := j0 + jj + for i := range ln { + d[jj] += vv[i] * rMat[(k+i)*n+jcol] + } + } + for j := range wd { + d[j] = beta * d[j] + } + for i := range ln { + vi := vv[i] + base := (k+i)*n + j0 + for j := range wd { + rMat[base+j] -= vi * d[j] + } + } + } + // Q columns, one at a time: per column the access is already + // linear in i (row-major, row col, columns k..m−1), and a slab + // of consecutive columns is a slab of consecutive rows. + for col := gi * m / nw; col < (gi+1)*m/nw; col++ { + dot := 0.0 + base := col*m + k + for i := range ln { + dot += vv[i] * qMat[base+i] + } + w := beta * dot + for i := range ln { + qMat[base+i] -= vv[i] * w + } + } + } + for k = range n { + // The reflector's construction and its scaling live in + // householderReflector, shared with the explicit-Q-free solve + // below. + var ok bool + ln, beta, ok = householderReflector(rMat, m, n, k, v) + if !ok { + continue + } + vv = v[:ln] + // Columns k..n-1 of R (in blocks, one item per block) and all m + // columns of Q: disjoint writes, read-only v and beta, so they + // can run as one parallel crew. The work weight counts a block + // as the qrColumnBlock columns it covers, which keeps the crew + // size honest about R's share. + items := (n-k+qrColumnBlock-1)/qrColumnBlock + m + weight := (n - k) + m + w := min(weight*ln/qrWorkQuantum+1, engine.WorkersFor(items)) + if w < 2 { + worker(0, 1) + continue + } + var wg sync.WaitGroup + for gi := range w { + wg.Go(func() { + worker(gi, w) + }) + } + wg.Wait() + } + return qMat, rMat, nil +} + +// denseFloats copies a 2-D array's payload into a flat m×n float64 +// matrix, converting from any numeric dtype via floatAt. Contiguous +// float64 payloads take a direct-copy fast path; floatAt gives the +// identical values for every dtype and layout either way. +func denseFloats(a *core.Array, m, n int) []float64 { + out := make([]float64, m*n) + if a.Dtype() == core.Float && !a.Strided() && len(a.RawFloats()) == m*n { + copy(out, a.RawFloats()) + return out + } + for i := range m { + for j := range n { + out[i*n+j] = a.FloatAt(i*n + j) + } + } + return out +} + +// floatsToArray wraps a flat m×n float64 matrix as a core.Array. +func floatsToArray(data []float64, shape []int) *core.Array { + out := core.New(core.Float, shape...) + copy(out.RawFloats(), data) + return out +} + +func eye(n int) []float64 { + out := make([]float64, n*n) + for i := range n { + out[i*n+i] = 1 + } + return out +} + +func sign(x float64) float64 { + switch { + case x > 0: + return 1 + case x < 0: + return -1 + default: + return 0 + } +} diff --git a/linalg/decomp2.go b/linalg/decomp2.go new file mode 100644 index 0000000..125d8a2 --- /dev/null +++ b/linalg/decomp2.go @@ -0,0 +1,1334 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +// SVD, Eigen, Pinverse, MatrixRank and Cond. Operates on +// float64 in and out; complex matrices are rejected. +// +// SVD: two-sided Jacobi rotation (Kogbetliantz / Press et al. NR3 §11.4). +// Both U and V are computed simultaneously by applying Jacobi +// rotations column-by-column on A and row-by-row on AᵀA; the +// algorithm is quadratic in the number of sweeps but converges +// robustly on ill-conditioned and rank-deficient matrices and is +// straightforward to keep correct. The matrix sizes tensor calls +// (n ≪ 256) keep this competitive with the QR approach. +// +// Eigen: implicit-shift QR on a symmetric tridiagonal matrix +// (Householder reduction first). Quadratic in sweeps but standard +// and stable. Both eigenvalues and eigenvectors are returned. + +import ( + "math" + "slices" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// svdMinWork is the measured parallel gate for one reflector +// application or matrix product slice, counted in element touches +// (items dispatched times the per-item inner length). Below it the +// work finishes before the worker dispatch pays off and runs on the +// calling goroutine. The bulge chase itself (symmetricQr, +// tridiagQrStep) stays serial: each rotation reads what the previous +// one wrote, and even the rounding-level entries far from the band +// feed a bulge a few rotations later, so no part of the tMat sweep may +// be reordered or narrowed. Only the eigenvector accumulation is +// order-safe: its mixes touch qMat alone, so they are recorded during +// the chase and replayed in batches over disjoint row ranges. +const svdMinWork = 16384 + +// svdMinItemsPerWorker is the floor on items per worker in the SVD +// dispatches. +const svdMinItemsPerWorker = 8 + +// qrFlushCap bounds the recorded rotation buffer between batched +// applications, keeping the replay scratch at a fixed size for any n. +const qrFlushCap = 65536 + +// qrRotation records one Givens rotation accumulated on the right of +// the eigenvector matrix: coordinates k and k+1 rotate by +// G = [[c, −s], [s, c]]. +type qrRotation struct { + k int + c, s float64 +} + +// qrSweep collects the rotations of one symmetricQr run and applies +// them to the eigenvector accumulator in batches. +// +// Deferring the mixes is bit-identical to applying them inside the +// chase. A qMat cell (i, k) is touched only by rotations whose band +// position is k or k-1, each computing c·old + s·old' and -s·old + +// c·old' from the cell's current value and the rotation's (c, s); the +// pair (c, s) is read off tMat alone. Recording the rotations and +// replaying them in the same chronological order therefore feeds every +// cell exactly the arithmetic sequence the inline loop produced, and a +// flush only cuts that sequence into consecutive chunks. The replay +// walks disjoint row ranges per worker, so the parallel batch and the +// serial fallback write the same bits. +type qrSweep struct { + qMat []float64 + n int + rots []qrRotation + // inline mirrors the pre-deferral behaviour: below qrDeferMin the + // chase applies each rotation to the accumulator as it runs, + // because the buffer, the flush and the worker dispatch cost more + // than the n-row mix they would save. + inline bool +} + +// qrDeferMin is the accumulator size from which the eigenvector mixes +// defer to batched replay. Measured on the square sweep: at 128 the +// deferral loses a third of the run to the buffer and the dispatch, +// at 192 it measures even, and at 256 the batched replay wins 1.6 +// times, with the gap widening as the mix grows cubically. +const qrDeferMin = 192 + +func newQrSweep(qMat []float64, n int) *qrSweep { + return &qrSweep{qMat: qMat, n: n, rots: make([]qrRotation, 0, 256), + inline: n < qrDeferMin} +} + +// record applies one rotation inline below qrDeferMin and otherwise +// appends it, flushing the buffer at the cap so the scratch stays +// bounded for large n. Both modes feed the cell the same arithmetic in +// the same chronological position, so the result does not depend on +// the mode. +func (qa *qrSweep) record(k int, c, s float64) { + if qa.inline { + qMat, n := qa.qMat, qa.n + for i := range n { + qi := qMat[i*n+k] + qi1 := qMat[i*n+k+1] + qMat[i*n+k] = c*qi + s*qi1 + qMat[i*n+k+1] = -s*qi + c*qi1 + } + return + } + qa.rots = append(qa.rots, qrRotation{k: k, c: c, s: s}) + if len(qa.rots) >= qrFlushCap { + qa.flush() + } +} + +// flush applies the recorded rotations to qMat, one worker per row +// range, and empties the buffer. Below the dispatch gate the batch +// runs on the calling goroutine. +func (qa *qrSweep) flush() { + rots := qa.rots + if len(rots) == 0 { + return + } + n := qa.n + qMat := qa.qMat + apply := func(start, end int) { + for i := start; i < end; i++ { + row := i * n + for _, r := range rots { + qi := qMat[row+r.k] + qi1 := qMat[row+r.k+1] + qMat[row+r.k] = r.c*qi + r.s*qi1 + qMat[row+r.k+1] = -r.s*qi + r.c*qi1 + } + } + } + if len(rots)*n >= svdMinWork { + engine.ParallelMin(n, 1, func(start, end int) { + apply(start, end) + }) + } else { + apply(0, n) + } + qa.rots = rots[:0] +} + +// scaleWinLo and scaleWinHi bound the magnitude window the squared +// arithmetic of the decompositions is exact and finite in. Every square, +// sum of squares, product and fourth power of an entry inside it stays a +// normal float64: 2^200 squared is 2^400 and its fourth power 2^800, far +// inside the 1.8e308 ceiling, and 2^-300 squared is 2^-600, far above +// the 2^-1022 normal floor. Outside the window the raw accumulations +// underflow to zero, which silently skips reflectors, or overflow to +// +Inf, which zeroes a reflector's beta and poisons the update with NaN. +const ( + scaleWinLo = -300 // 2^-300 ≈ 4.9e-91 + scaleWinHi = 200 // 2^200 ≈ 1.6e60 +) + +// windowScale returns the exact power of two that moves maxAbs into +// [2^scaleWinLo, 2^scaleWinHi], and 1 when it already lies inside. A +// decomposition multiplies its matrix by the factor and its spectrum by +// 1/f: that is the same problem in scaled units, because eigenvalues and +// singular values scale with the matrix while eigenvectors and the +// orthogonal factors do not. Powers of two carry no rounding, so an +// in-window input keeps every bit of the original arithmetic and a +// rescaled one keeps the accuracy its exponent leaves it. +func windowScale(maxAbs float64) float64 { + if maxAbs <= 0 || math.IsNaN(maxAbs) || math.IsInf(maxAbs, 0) { + return 1 + } + // Frexp splits maxAbs into m·2^e with m in [0.5, 1). + _, e := math.Frexp(maxAbs) + switch { + case e-1 > scaleWinHi: + return math.Ldexp(1, scaleWinHi-e) + case e-1 < scaleWinLo: + return math.Ldexp(1, scaleWinLo+1-e) + } + return 1 +} + +// maxMagF64 returns the largest absolute value in a real slice. +func maxMagF64(s []float64) float64 { + m := 0.0 + for _, v := range s { + if a := math.Abs(v); a > m { + m = a + } + } + return m +} + +// maxMagComplex returns the largest magnitude in a complex slice. The +// magnitudes come from math.Hypot, so an entry beyond ~1e154 keeps a +// finite magnitude where a sum of raw squares overflows. +func maxMagComplex(s []complex128) float64 { + m := 0.0 + for _, z := range s { + if a := cmplxAbs(z); a > m { + m = a + } + } + return m +} + +// scaleFloats multiplies every entry of s by f; f is an exact power of +// two, so the entries come back bit for bit from unscaleFloats. +func scaleFloats(s []float64, f float64) { + for i := range s { + s[i] *= f + } +} + +// unscaleFloats divides every entry of s by f, an exact power of two. +func unscaleFloats(s []float64, f float64) { + for i := range s { + s[i] /= f + } +} + +// scaleComplexes multiplies every entry of s by the real power of two f. +func scaleComplexes(s []complex128, f float64) { + for i := range s { + s[i] = complex(real(s[i])*f, imag(s[i])*f) + } +} + +// unscaleComplexes divides every entry of s by the real power of two f. +func unscaleComplexes(s []complex128, f float64) { + for i := range s { + s[i] = complex(real(s[i])/f, imag(s[i])/f) + } +} + +// SVD returns the thin singular value decomposition a = U · Σ · Vᵀ +// of an m×n matrix a with m ≥ n. For m < n the matrix is +// transposed first and the orthogonal factors are swapped on +// return. U is m×n with orthonormal columns (Uᵀ U = I_n); Σ is a +// 1-D length-n array of non-negative singular values sorted +// descending; Vᵀ is n×n orthogonal. +// +// Algorithm: +// 1. Bidiagonalise A, giving A = U₁ · B · V₁ᵀ via Householder. +// 2. Form C = Bᵀ B (n × n symmetric tridiagonal). +// 3. Symmetric implicit-shift QR with Wilkinson shift on C, +// accumulating rotations into V = V₁ · Q. +// 4. Σ = √(eigenvalues of C). Reconstruct U = A · V / Σ and +// orthogonalise via Gram-Schmidt to absorb numerical drift. +// +// This is the standard double-step SVD via Bᵀ B + symmetric QR: +// robust, well-understood, and reuses the symmetric tridiagonal +// eigensolver already shipped for `Eigen`. +// +// The double step prices the small end of the spectrum: the singular +// values come from the eigenvalues of BᵀB, so their accuracy on very +// small singular values is that of a squared condition number, fine +// for rank decisions, weaker for resolving near-null directions. +// `SVDComplex` states the same contract for its Aᴴ·A route. +func SVD(a *core.Array) (u, sigma, vt *core.Array, err error) { + if a.Dtype() == core.Complex { + return nil, nil, nil, base.Errf("SVD: complex matrices are not supported") + } + if a.NDim() != 2 { + return nil, nil, nil, base.Errf("SVD: needs a 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m == 0 || n == 0 { + return nil, nil, nil, base.Errf("SVD: zero-sized matrix, got shape %s", base.ShapeText(a.Shape())) + } + transposed := false + Amat := denseFloats(a, a.Shape()[0], a.Shape()[1]) + // Outside the safe window every squared intermediate of the + // bidiagonalisation, of C = BᵀB and of the σ recovery would leave + // the normal range. The matrix is moved into the window for the + // computation and the singular values, which carry the scale, are + // moved back once the factorisation is done; the orthogonal factors + // are scale-free. + ws := windowScale(maxMagF64(Amat)) + if ws != 1 { + scaleFloats(Amat, ws) + } + origM, origN := m, n + if m < n { + transposed = true + Amat = base.TransposeFlat(Amat, m, n) + m, n = n, m + } + bDiag, bSuper, v1 := bidiagonalise(Amat, m, n) + // Form C = Bᵀ B (n × n symmetric tridiagonal): AᵀA = V₁ C V₁ᵀ. + // The diagonal carries the superdiagonal's contribution: + // C[i,i] = B[i,i]² + B[i−1,i]²; the off-diagonal is + // C[i,i+1] = B[i,i]·B[i,i+1]. + cMat := make([]float64, n*n) + for i := 0; i < n; i++ { + cMat[i*n+i] = bDiag[i] * bDiag[i] + if i > 0 { + cMat[i*n+i] += bSuper[i-1] * bSuper[i-1] + } + if i+1 < n { + cMat[i*n+i+1] = bDiag[i] * bSuper[i] + cMat[(i+1)*n+i] = cMat[i*n+i+1] + } + } + // Eigendecompose C with the shipped symmetric solver: C = W Λ Wᵀ, + // so the right singular vectors are V = V₁·W with Σ = √Λ. C is + // (n × n) symmetric and tiny (n ≪ 1024 in the typical workspace), + // so the eigensolver is much cheaper than a QR sweep per call. + // C is exactly symmetric by construction, so the internals of + // `Eigen` run directly. + tMat, qMat := householderTridiag(cMat, n) + qT := base.TransposeFlat(qMat, n, n) + if err := symmetricQr(tMat, qT, n); err != nil { + return nil, nil, nil, base.Errf("SVD: %w", err) + } + eVals := make([]float64, n) + for i := range n { + eVals[i] = tMat[i*n+i] + } + eIdx := sortAscIndices(eVals) + cVals := make([]float64, n) // ascending eigenvalues, as Eigen returns them + eVecs := make([]float64, n*n) + for j := range n { + cVals[j] = eVals[eIdx[j]] + for i := range n { + eVecs[i*n+j] = qT[i*n+eIdx[j]] + } + } + vMat := make([]float64, n*n) // V = V₁·W + // Rows are independent; each element keeps its serial dot over k. + if n*n*n >= svdMinWork { + engine.ParallelMin(n, 1, func(start, end int) { + for i := start; i < end; i++ { + for j := range n { + s := 0.0 + for k := range n { + s += v1[i*n+k] * eVecs[k*n+j] + } + vMat[i*n+j] = s + } + } + }) + } else { + for i := range n { + for j := range n { + s := 0.0 + for k := range n { + s += v1[i*n+k] * eVecs[k*n+j] + } + vMat[i*n+j] = s + } + } + } + // Eigen returns ascending order; the SVD contract wants Σ + // descending, so permute the columns of V to match. + idx := sortDescIndices(cVals) + permuteCols(vMat, n, n, idx) + // Recover Σ directly from A as σₖ = ‖A·vₖ‖: the BᵀB round trip + // squares the condition number, so reading √λ back would lose half + // the digits on ill-conditioned inputs. Columns are independent and + // each norm keeps its serial accumulation order over i, then j. + sorted := make([]float64, n) + if m*n*n >= svdMinWork { + engine.ParallelMin(n, 1, func(start, end int) { + for k := start; k < end; k++ { + norm := 0.0 + for i := range m { + acc := 0.0 + for j := range n { + acc += Amat[i*n+j] * vMat[j*n+k] + } + norm += acc * acc + } + sorted[k] = math.Sqrt(norm) + } + }) + } else { + for k := range n { + norm := 0.0 + for i := range m { + acc := 0.0 + for j := range n { + acc += Amat[i*n+j] * vMat[j*n+k] + } + norm += acc * acc + } + sorted[k] = math.Sqrt(norm) + } + } + // The recovered Σ can disagree with the eigenvalue order by + // rounding on near-degenerate values; settle on one descending + // order before building U, permuting V along with it. + if ord := sortDescIndices(sorted); !isIdentityPerm(ord) { + permuteCols(vMat, n, n, ord) + perm := make([]float64, n) + copy(perm, sorted) + for i := range n { + sorted[i] = perm[ord[i]] + } + } + // Thin U = A·V·Σ⁻¹ from the working (m ≥ n) matrix, with A and Σ in + // the same scaled units. + uThin := reconstructUThin(Amat, vMat, sorted, m, n) + // The singular values carry the matrix's scale and the orthogonal + // factors do not, so only Σ is moved back out of the scaled units + // the factorisation ran in. + if ws != 1 { + unscaleFloats(sorted, ws) + } + U := floatsToArray(uThin, []int{m, n}) + S := floatsToArray(sorted, []int{n}) + Vt := floatsToArray(base.TransposeFlat(vMat, n, n), []int{n, n}) + if transposed { + // The input had m < n; the factorisation ran on the transpose + // Aᵀ = U' Σ V'ᵀ, so the original reads A = V' Σ U'ᵀ: + // U_orig = V' itself (origM × origM) and Vᵀ_orig = (thin U')ᵀ. + return floatsToArray(vMat, []int{origM, origM}), + S, + floatsToArray(base.TransposeFlat(uThin, m, n), []int{origM, origN}), + nil + } + return U, S, Vt, nil +} + +// bidiagonalise reduces A (m×n with m ≥ n, row-major) to upper +// bidiagonal form via Golub-Kahan Householder reflectors: B = +// H_p…H₁·A·K₁…K_q, so A = U₁·B·V₁ᵀ with V₁ = K₁K₂…K_q accumulated by +// right-multiplying the reflectors into its columns. The left factor U₁ +// is not accumulated: the only caller recovers U from A·V after the +// sweep, and carrying the m×m accumulator cost an eye(m) matrix and an +// m-row update per left reflector without feeding any returned value. +// The bidiagonal and the right accumulator are exactly the values the +// accumulation carried: neither reads uMat. +// +// Reflectors apply strictly in sweep order; within one reflector every +// touched column or row is updated exactly once from the frozen +// reflector vector and beta, so the per-line work (dot along the line, +// then the same-operand rank-1 subtraction) dispatches over the lines +// of bFull and the accumulator in one parallel range. The per-element +// sequence is the serial one, so the factorisation is bit-identical +// across crew sizes. +func bidiagonalise(aMat []float64, m, n int) (bDiag, bSuper []float64, vMat []float64) { + bDiag = make([]float64, n) + if n > 1 { + bSuper = make([]float64, n-1) + } + bFull := make([]float64, m*n) + copy(bFull, aMat) + vMat = eye(n) + // Reflector scratch reused across the sweep; each step uses the + // prefix it needs. + vec := make([]float64, m) + vecR := make([]float64, n) + for k := range n { + // Left reflector: rows k..m-1 of column k, zeroing below (k, k). + lv := vec[:m-k] + for i := range lv { + lv[i] = bFull[(k+i)*n+k] + } + hh := householderVectorInto(lv, lv) + if hh.beta != 0 { + ln := m - k + bColumn := func(j int) { + dot := 0.0 + for i := range ln { + dot += hh.v[i] * bFull[(k+i)*n+j] + } + w := hh.beta * dot + for i := range ln { + bFull[(k+i)*n+j] -= hh.v[i] * w + } + } + if (n-k)*ln >= svdMinWork { + engine.ParallelMin(n-k, svdMinItemsPerWorker, func(start, end int) { + for j := start; j < end; j++ { + bColumn(k + j) + } + }) + } else { + for j := k; j < n; j++ { + bColumn(j) + } + } + } + if k+1 >= n { + break + } + // Right reflector: columns k+1..n-1 of row k, zeroing right of + // (k, k+1). + rv := vecR[:n-k-1] + for j := range rv { + rv[j] = bFull[k*n+(k+1+j)] + } + hhR := householderVectorInto(rv, rv) + if hhR.beta != 0 { + ln := len(rv) + bRow := func(i int) { + dot := 0.0 + for j := range ln { + dot += hhR.v[j] * bFull[i*n+(k+1+j)] + } + w := hhR.beta * dot + for j := range ln { + bFull[i*n+(k+1+j)] -= hhR.v[j] * w + } + } + vRow := func(j int) { + dot := 0.0 + for i := range ln { + dot += hhR.v[i] * vMat[j*n+(k+1+i)] + } + w := hhR.beta * dot + for i := range ln { + vMat[j*n+(k+1+i)] -= hhR.v[i] * w + } + } + // V₁ = K₁K₂…: likewise from the right, columns k+1..n-1. + // Rows k..m-1 of bFull and rows of vMat: disjoint lines. + items := (m - k) + n + if items*ln >= svdMinWork { + engine.ParallelMin(items, svdMinItemsPerWorker, func(start, end int) { + for it := start; it < end; it++ { + if it < m-k { + bRow(k + it) + } else { + vRow(it - (m - k)) + } + } + }) + } else { + for i := k; i < m; i++ { + bRow(i) + } + for j := range n { + vRow(j) + } + } + } + } + for i := range n { + bDiag[i] = bFull[i*n+i] + if i+1 < n { + bSuper[i] = bFull[i*n+i+1] + } + } + return bDiag, bSuper, vMat +} + +// smallOff is the deflation threshold for the tridiagonal QR sweep. +func smallOff(d1, d2 float64) float64 { + return base.EpsF * (math.Abs(d1) + math.Abs(d2)) +} + +// reconstructUThin builds thin U = A_orig · V · diag(1/Σ) when Σ is +// non-zero, and falls back to the kernel projection for small +// σᵢ. aOrig is (post-transpose) input shape (m, n). vMat is (n, n). +// sVals are the singular values in matched column order. Result +// has shape (m, n). +// +// The candidate projection A_orig·V is computed once into a scratch, +// dispatched over the rows of aOrig: every element is the same serial +// dot over j (ascending) the per-column loop used to run, so the +// values are bit-identical and only their computing order moved. The +// modified Gram-Schmidt pass stays serial: column k reads every +// earlier column, and each projection's dot is an order-bound +// reduction over the column being updated. +func reconstructUThin(aOrig, vMat, sVals []float64, m, n int) []float64 { + biggest := 0.0 + for _, v := range sVals { + if v > biggest { + biggest = v + } + } + thresh := base.EpsF * float64(m) * biggest + // The Gram-Schmidt pass tests the norm of a column that the division + // by σ makes unit-length, so its floor is an absolute one relative to + // that unit and not to the largest singular value: with the singular + // values themselves as the reference, a matrix whose spectrum is + // above ~1e16 would route every healthy column to the kernel rebuild + // below and lose the factorisation. + unitTol := base.EpsF * float64(m) + // The un-orthonormalised candidate is A_orig·V/Σ on the columns + // where σ is non-tiny (and A_orig·V otherwise), re-orthonormalised + // via modified Gram-Schmidt below so that UᵀU = I_n holds even when + // the reduction leaves small numerical drift. A column whose + // residual is essentially zero (σ at rounding level) is rebuilt from + // the best orthogonalised coordinate candidate e_j: the largest + // residual is at least 1/√m, because a unit vector of the + // orthogonal complement has a coordinate at least that large. + uMat := make([]float64, m*n) + // Scratch reused across columns: the working column and, for kernel + // columns, the candidate buffers (fully overwritten before use). + best := make([]float64, m) + cand := make([]float64, m) + // proj holds the un-orthonormalised projection A_orig·V, one value + // per (row, singular vector) pair, each from the same serial dot. + proj := engine.GetFloat64Buf(m * n) + defer engine.PutFloat64Buf(proj) + if m*n*n >= svdMinWork { + engine.ParallelMin(m, 1, func(start, end int) { + for i := start; i < end; i++ { + row := aOrig[i*n : i*n+n] + out := i * n + for k := range n { + s := 0.0 + for j := range n { + s += row[j] * vMat[j*n+k] + } + proj[out+k] = s + } + } + }) + } else { + for i := range m { + for k := range n { + s := 0.0 + for j := range n { + s += aOrig[i*n+j] * vMat[j*n+k] + } + proj[i*n+k] = s + } + } + } + for k := range n { + if sVals[k] > thresh { + inv := 1 / sVals[k] + for i := range m { + uMat[i*n+k] = proj[i*n+k] * inv + } + } else { + for i := range m { + uMat[i*n+k] = proj[i*n+k] + } + } + // Subtract projections on the earlier columns. + for j := range k { + dot := 0.0 + for i := range m { + dot += uMat[i*n+j] * uMat[i*n+k] + } + for i := range m { + uMat[i*n+k] -= dot * uMat[i*n+j] + } + } + nrm := 0.0 + for i := range m { + nrm += uMat[i*n+k] * uMat[i*n+k] + } + nrm = math.Sqrt(nrm) + if nrm > unitTol { + inv := 1 / nrm + for i := range m { + uMat[i*n+k] *= inv + } + continue + } + // Kernel column: orthogonalise every coordinate candidate and + // keep the healthiest residual. + bestRes := -1.0 + for j := range m { + for i := range m { + cand[i] = 0 + } + cand[j] = 1 + for range 2 { + for j2 := range k { + d := 0.0 + for i := range m { + d += uMat[i*n+j2] * cand[i] + } + for i := range m { + cand[i] -= d * uMat[i*n+j2] + } + } + } + rn := 0.0 + for i := range m { + rn += cand[i] * cand[i] + } + if rn > bestRes { + bestRes = rn + copy(best, cand) + } + } + if bestRes > 0 { + inv := 1 / math.Sqrt(bestRes) + for i := range m { + uMat[i*n+k] = best[i] * inv + } + } + } + return uMat +} + +// Eigen returns the eigenvalues and orthonormal eigenvectors of a +// real symmetric n×n matrix a. Eigenvalues are returned in a 1-D +// float array, sorted in ascending order; the eigenvectors are +// the columns of an n×n orthogonal array. Asymmetric matrices are +// not supported. +func Eigen(a *core.Array) (values, vectors *core.Array, err error) { + if a.Dtype() == core.Complex { + return nil, nil, base.Errf("Eigen: complex matrices are not supported") + } + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, nil, base.Errf("Eigen: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + mat := denseFloats(a, n, n) + // A matrix outside the safe window would drive the reflector norms + // and the sweep's squared accumulation past the representable range; + // the tridiagonalisation and the QR sweep run on it scaled into the + // window and the eigenvalues are scaled back. The eigenvectors are + // the same for every positive multiple of the matrix. + ws := windowScale(maxMagF64(mat)) + if ws != 1 { + scaleFloats(mat, ws) + } + if !isSymmetric(mat, n) { + return nil, nil, base.Errf("Eigen: matrix is not symmetric within 1e-12 tolerance") + } + tMat, qMat := householderTridiag(mat, n) + // The reduction maintains T = qMat·A·qMatᵀ (reflectors applied on + // the left of qMat), so the eigenvector matrix of A after the QR + // sweep is qMatᵀ·G_total. The sweep accumulates rotations on the + // right of its accumulator, hence the transpose before the call. + qT := base.TransposeFlat(qMat, n, n) + if err := symmetricQr(tMat, qT, n); err != nil { + return nil, nil, base.Errf("Eigen: %w", err) + } + vals := make([]float64, n) + for i := range n { + vals[i] = tMat[i*n+i] + } + idx := sortAscIndices(vals) + out := make([]float64, n) + for i := range n { + out[i] = vals[idx[i]] + } + sortedCols := make([]float64, n*n) + for j := range n { + for i := range n { + sortedCols[i*n+j] = qT[i*n+idx[j]] + } + } + if ws != 1 { + unscaleFloats(out, ws) + } + return floatsToArray(out, []int{n}), floatsToArray(sortedCols, []int{n, n}), nil +} + +// Pinverse returns the Moore-Penrose pseudoinverse of a 2-D matrix, +// built from the SVD by inverting singular values strictly greater +// than ε. When eps ≤ 0 the default is max(m, n) · max(Σ) · base.EpsF. +func Pinverse(a *core.Array, eps float64) (*core.Array, error) { + if a.Dtype() == core.Complex { + return nil, base.Errf("Pinverse: complex matrices are not supported") + } + if a.NDim() != 2 { + return nil, base.Errf("Pinverse: needs a 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m == 0 || n == 0 { + return nil, base.Errf("Pinverse: zero-sized matrix, got shape %s", base.ShapeText(a.Shape())) + } + u, sigma, vt, err := SVD(a) + if err != nil { + return nil, err + } + sVals := sigma.RawFloats() + // The thin factorisation carries min(m, n) singular values: U is + // (m, r) and Vᵀ is (r, n), whatever the input aspect ratio. + r := min(m, n) + maxS := 0.0 + for _, s := range sVals { + if s > maxS { + maxS = s + } + } + if eps <= 0 { + dim := max(n, m) + eps = float64(dim) * maxS * base.EpsF + } + dInv := make([]float64, r) + for i, s := range sVals { + if s > eps { + dInv[i] = 1 / s + } + } + uMat := denseFloats(u, m, r) + vtMat := denseFloats(vt, r, n) + out := make([]float64, n*m) + // A⁺ = V · Σ⁻¹ · Uᵀ = (V · Σ⁻¹) · Uᵀ, shape (n, m): + // out[i,j] = Σ_k V[i,k] · Σ⁻¹[k] · U[j,k], with V[i,k] = Vᵀ[k,i]. + for i := range n { + for j := range m { + s := 0.0 + for k := range r { + s += vtMat[k*n+i] * dInv[k] * uMat[j*r+k] + } + out[i*m+j] = s + } + } + return floatsToArray(out, []int{n, m}), nil +} + +// MatrixRank returns the number of singular values of a 2-D matrix +// strictly greater than eps. With eps ≤ 0 the default is +// max(m, n) · max(Σ) · base.EpsF. +func MatrixRank(a *core.Array, eps float64) (int, error) { + if a.Dtype() == core.Complex { + return 0, base.Errf("MatrixRank: complex matrices are not supported") + } + if a.NDim() != 2 { + return 0, base.Errf("MatrixRank: needs a 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + _, sigma, _, err := SVD(a) + if err != nil { + return 0, err + } + return countAboveThreshold(sigma.RawFloats(), eps, a.Shape()[0], a.Shape()[1]), nil +} + +// Cond returns the 2-norm condition number σ_max / σ_min of a 2-D +// matrix. If any singular value is zero (or at-or-below a positive +// eps) the result is +Inf, mirroring the standard convention. +func Cond(a *core.Array, eps float64) (float64, error) { + if a.Dtype() == core.Complex { + return 0, base.Errf("Cond: complex matrices are not supported") + } + if a.NDim() != 2 { + return 0, base.Errf("Cond: needs a 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + _, sigma, _, err := SVD(a) + if err != nil { + return 0, err + } + return condValue(sigma.RawFloats(), eps, a.Shape()[0], a.Shape()[1]), nil +} + +// householderTridiag reduces a symmetric matrix to tridiagonal form via +// Householder reflectors, accumulating the orthogonal factor into qMat. +// +// The sweep order is untouched. Within one reflector the LEFT +// application to T (columns) and the accumulation into qMat (columns) +// run over disjoint buffers, so they share one dispatch; the RIGHT +// application to T (rows) must follow, because it reads the entries the +// left application has just written; the drift symmetrisation follows +// it serially. Each line is updated exactly once from the frozen +// reflector with the serial dot order, so the result is bit-identical. +func householderTridiag(aMat []float64, n int) (tMat, qMat []float64) { + tMat = make([]float64, n*n) + copy(tMat, aMat) + qMat = eye(n) + // Reflector scratch reused across the sweep. + v := make([]float64, n) + for k := 0; k < n-2; k++ { + vv := v[:n-k-1] + for i := 0; i < n-k-1; i++ { + vv[i] = tMat[(k+1+i)*n+k] + } + u := householderVectorInto(vv, vv) + if u.beta == 0 { + continue + } + ln := n - k - 1 + tColumn := func(j int) { + dot := 0.0 + for i := range ln { + dot += u.v[i] * tMat[(k+1+i)*n+j] + } + w := u.beta * dot + for i := range ln { + tMat[(k+1+i)*n+j] -= u.v[i] * w + } + } + qColumn := func(j int) { + dot := 0.0 + for i := range ln { + dot += u.v[i] * qMat[(k+1+i)*n+j] + } + w := u.beta * dot + for i := range ln { + qMat[(k+1+i)*n+j] -= u.v[i] * w + } + } + // Apply H from both sides to T: left pass over the columns of T + // together with the qMat accumulation (disjoint buffers). + items := (n - k) + n + if items*ln >= svdMinWork { + engine.ParallelMin(items, svdMinItemsPerWorker, func(start, end int) { + for it := start; it < end; it++ { + if it < n-k { + tColumn(k + it) + } else { + qColumn(it - (n - k)) + } + } + }) + } else { + for j := k; j < n; j++ { + tColumn(j) + } + for j := range n { + qColumn(j) + } + } + // Right pass over the rows of T: reads the left pass's writes, + // so it stays strictly after it. + tRow := func(i int) { + dot := 0.0 + for j := range ln { + dot += u.v[j] * tMat[i*n+(k+1)+j] + } + w := u.beta * dot + for j := range ln { + tMat[i*n+(k+1)+j] -= u.v[j] * w + } + } + if (n-k)*ln >= svdMinWork { + engine.ParallelMin(n-k, svdMinItemsPerWorker, func(start, end int) { + for i := start; i < end; i++ { + tRow(k + i) + } + }) + } else { + for i := k; i < n; i++ { + tRow(i) + } + } + // Symmetrise against numerical drift. + for i := k + 1; i < n; i++ { + for j := i + 1; j < n; j++ { + avg := (tMat[i*n+j] + tMat[j*n+i]) / 2 + tMat[i*n+j] = avg + tMat[j*n+i] = avg + } + } + } + return tMat, qMat +} + +// symmetricQr runs the implicit-shift symmetric QR algorithm with +// the Wilkinson shift on a symmetric tridiagonal tMat. Returns the +// updated tridiagonal (whose diagonal is the eigenvalues) and the +// eigenvector matrix Q, or an error when the sweep exhausts its +// passes without deflating every subdiagonal entry: an unconverged +// spectrum is never returned as an answer. +func symmetricQr(tMat, qMat []float64, n int) error { + maxIters := max(30*n, 30) + // Deflation floor relative to the matrix norm: a purely + // neighbour-relative threshold never triggers when both diagonal + // entries are small, even though the off-diagonal is already at + // rounding level of the whole transform (backward stability + // guarantees nothing better than eps·‖T‖ anyway). + scale := 0.0 + for i := range n { + scale += tMat[i*n+i] * tMat[i*n+i] + if i+1 < n { + e := tMat[i*n+i+1] + scale += 2 * e * e + } + } + tolAbs := base.EpsF * math.Sqrt(scale) + negligible := func(i int) bool { // tests the subdiagonal entry (i, i-1) + e := math.Abs(tMat[i*n+(i-1)]) + // An exactly zero entry is deflated by definition: a zero + // matrix has tolAbs 0 and relative floors of 0, and a strict + // inequality would never deflate it. + return e == 0 || e < smallOff(tMat[i*n+i], tMat[(i-1)*n+(i-1)]) || e < tolAbs + } + acc := newQrSweep(qMat, n) + for range maxIters { + h := n - 1 + for h > 0 && negligible(h) { + tMat[h*n+(h-1)] = 0 + tMat[(h-1)*n+h] = 0 + h-- + } + if h <= 0 { + break + } + // The active block containing h reaches up to the first + // negligible subdiagonal entry (its decoupling boundary): + // expanding l past that boundary would leave a zero entry + // inside the window, where the bulge chase would abort and + // strand the block below it. + l := h - 1 + for l > 0 && !negligible(l) { + l-- + } + if l > 0 { + tMat[l*n+(l-1)] = 0 + tMat[(l-1)*n+l] = 0 + } + if h <= l { + continue + } + if h-l == 1 { + // A 2×2 block is diagonalised exactly by one Jacobi + // rotation; the shift chase leaves a rounding-level + // off-diagonal that can sit forever above the deflation + // threshold, so finish it off directly. + jacobi2x2(tMat, acc, l, n) + tMat[l*n+h] = 0 + tMat[h*n+l] = 0 + continue + } + d := tMat[(h-1)*n+(h-1)] + e := tMat[(h-1)*n+h] + f := tMat[h*n+h] + shift := wilkinsonShift2x2(d, e, f) + tridiagQrStep(tMat, acc, l, h, n, shift) + } + // An exhausted sweep must not pass its tridiagonal off as a + // spectrum: every sibling iterator here (hessenbergQr, the + // Golub-Reinsch SVD, schurQR) errors on exhaustion, and so does + // this one. + for i := 1; i < n; i++ { + if !negligible(i) { + return base.Errf("the QR sweep failed to deflate a %d×%d matrix in %d passes", n, n, maxIters) + } + } + acc.flush() + return nil +} + +// applyGivens applies the rotation G = [[c, −s], [s, c]] on coordinates +// (k, k+1) as a similarity transform of tMat and records it for the +// batched accumulation on the right of the eigenvector matrix. +func applyGivens(tMat []float64, acc *qrSweep, k, n int, c, s float64) { + for j := range n { + tk := tMat[k*n+j] + tk1 := tMat[(k+1)*n+j] + tMat[k*n+j] = c*tk + s*tk1 + tMat[(k+1)*n+j] = -s*tk + c*tk1 + } + for i := range n { + ti := tMat[i*n+k] + ti1 := tMat[i*n+k+1] + tMat[i*n+k] = c*ti + s*ti1 + tMat[i*n+k+1] = -s*ti + c*ti1 + } + acc.record(k, c, s) +} + +// jacobi2x2 diagonalises the symmetric 2×2 block at (k, k+1) exactly +// with one Jacobi rotation, recording the rotation into acc. +func jacobi2x2(tMat []float64, acc *qrSweep, k, n int) { + a := tMat[k*n+k] + b := tMat[k*n+k+1] + d := tMat[(k+1)*n+(k+1)] + if b == 0 { + return + } + tau := (d - a) / (2 * b) + // Smaller root of t² + ((a−d)/b)·t − 1 = 0: t = −sign(τ)/(|τ|+√(1+τ²)). + t := -1.0 / (math.Abs(tau) + math.Sqrt(1+tau*tau)) + if tau < 0 { + t = -t + } + c := 1 / math.Sqrt(1+t*t) + s := t * c + applyGivens(tMat, acc, k, n, c, s) +} +func wilkinsonShift2x2(d, e, f float64) float64 { + delta := (d - f) / 2 + if delta == 0 { + return f - math.Abs(e) + } + sq := math.Sqrt(delta*delta + e*e) + if delta > 0 { + return f - e*e/(delta+sq) + } + return f + e*e/(sq-delta) +} + +// tridiagQrStep performs one implicit-shift QR step on the symmetric +// tridiagonal block [l, h] of tMat, recording the rotations into acc +// for the batched eigenvector accumulation. Every rotation is a +// similarity transform applied to full rows and columns, so tMat stays +// exactly similar to its input; the pair (c, s) comes from the +// subdiagonal and the bulge the previous rotation left behind, the +// bulge chase of Golub & Van Loan §8.5. +// +// The tMat passes are strictly sequential and must stay so: rotation +// k+1 annihilates the bulge rotation k wrote, and the rounding-level +// entries far from the band drift inward one diagonal per pass, so +// narrowing or reordering any pass would move bits. Only the qMat mix +// is free of that chain; it is replayed from the recording, in order, +// by acc.flush. +func tridiagQrStep(tMat []float64, acc *qrSweep, l, h, n int, shift float64) { + // The first rotation is read from the leading entry of (T − μI); + // later ones annihilate the chased bulge. + x := tMat[l*n+l] - shift + y := tMat[l*n+(l+1)] + for k := l; k < h; k++ { + if k > l { + x = tMat[(k-1)*n+k] // subdiagonal entry + y = tMat[(k-1)*n+(k+1)] // bulge to annihilate + } + r := math.Hypot(x, y) + if r == 0 { + return // nothing to rotate; the band is already clean + } + c, s := x/r, y/r + // Rows k, k+1: left multiplication by Gᵀ, G = [[c, −s], [s, c]]. + for j := range n { + tk := tMat[k*n+j] + tk1 := tMat[(k+1)*n+j] + tMat[k*n+j] = c*tk + s*tk1 + tMat[(k+1)*n+j] = -s*tk + c*tk1 + } + // Columns k, k+1: right multiplication by G. + for i := range n { + ti := tMat[i*n+k] + ti1 := tMat[i*n+k+1] + tMat[i*n+k] = c*ti + s*ti1 + tMat[i*n+k+1] = -s*ti + c*ti1 + } + // Eigenvector accumulation on the right: Q = Q·G, replayed in + // order by acc.flush. + acc.record(k, c, s) + } +} + +// householder represents a Householder reflection v and its beta. +type householder struct { + v []float64 + beta float64 +} + +// householderVectorInto builds the standard reflect-to-first-coordinate +// Householder vector for x, where H = I - beta v vᵀ with H x = sign(x₀)‖x‖ e₁, +// writing the reflector into dst (which must have room for len(x) +// entries and may alias x) so sweep callers can reuse one buffer +// instead of allocating per reflection. Squared magnitudes are summed +// relative to the largest entry so a vector with entries near 1e154 +// does not overflow on the way to its norm. +// +// beta = 2/(vᵀv) needs the square of v's largest entry to stay in the +// normal range. Past about 1.3e154 that square is +Inf and beta comes +// out as +0, which callers read as "no reflection at all", and below +// about 1e-154 it is subnormal or 0 and beta comes out as +Inf, whose +// product with a zero dot is NaN. Either way the reflector is silently +// lost. When the largest entry falls outside that range the reflector is +// re-expressed in units of the exact power of two that brings it into +// the safe window: v is scaled and beta carries the compensatory square, +// so H is the same reflection with every intermediate O(1). +func householderVectorInto(dst, x []float64) householder { + maxAbs := 0.0 + for _, xi := range x { + if v := math.Abs(xi); v > maxAbs { + maxAbs = v + } + } + if maxAbs == 0 { + return householder{v: x, beta: 0} + } + scaled2 := 0.0 + for _, xi := range x { + xi /= maxAbs + scaled2 += xi * xi + } + norm := maxAbs * math.Sqrt(scaled2) + dst[0] = x[0] + signOrNonZero(x[0])*norm + copy(dst[1:], x[1:]) + vMax := 0.0 + for _, vi := range dst { + if v := math.Abs(vi); v > vMax { + vMax = v + } + } + if vMax == 0 { + return householder{v: dst, beta: 0} + } + v2 := 0.0 + for _, vi := range dst { + vi /= vMax + v2 += vi * vi + } + if beta := 2 / (vMax * vMax * v2); beta > 0 && beta < math.MaxFloat64 { + return householder{v: dst, beta: beta} + } + ws := windowScale(vMax) + scaleFloats(dst, ws) + // w keeps max|v| inside the safe window, so w*w is a normal finite + // number and 2/(w*w*v2) is finite and non-zero: the same reflection, + // still built so that H x = sign(x₀)‖x‖ e₁. + w := vMax * ws + return householder{v: dst, beta: 2 / (w * w * v2)} +} + +// signOrNonZero returns the sign of x, with 0 mapping to +1 (Householder +// convention). +func signOrNonZero(x float64) float64 { + if x >= 0 { + return 1 + } + return -1 +} + +// isSymmetric reports whether the matrix is symmetric within a purely +// relative tolerance: a floor would admit an asymmetry that is a large +// fraction of a small-scale matrix, so the guard would call a matrix +// symmetric that is not, whatever the absolute size of its entries. +func isSymmetric(mat []float64, n int) bool { + scale := 0.0 + for i := range n { + for j := range n { + a := math.Abs(mat[i*n+j]) + if a > scale { + scale = a + } + } + } + tol := 1e-12 * scale + for i := range n { + for j := i + 1; j < n; j++ { + if math.Abs(mat[i*n+j]-mat[j*n+i]) > tol { + return false + } + } + } + return true +} + +func countAboveThreshold(s []float64, eps float64, m, n int) int { + if eps <= 0 { + dim := max(n, m) + biggest := 0.0 + for _, v := range s { + if v > biggest { + biggest = v + } + } + eps = float64(dim) * biggest * base.EpsF + } + count := 0 + for _, v := range s { + if v > eps { + count++ + } + } + return count +} + +func condValue(s []float64, eps float64, m, n int) float64 { + if len(s) == 0 { + return 0 + } + biggest := 0.0 + smallest := math.Inf(1) + for _, v := range s { + if v > biggest { + biggest = v + } + if v < smallest { + smallest = v + } + } + if biggest == 0 { + // The zero matrix has no usable inverse direction; the condition + // number is infinite, matching the doc contract. + return math.Inf(1) + } + if smallest <= 0 { + return math.Inf(1) + } + if eps > 0 && smallest <= eps { + return math.Inf(1) + } + return biggest / smallest +} + +func permuteCols(a []float64, rows, cols int, idx []int) { + out := make([]float64, rows*cols) + for j := range cols { + for i := range rows { + out[i*cols+j] = a[i*cols+idx[j]] + } + } + copy(a, out) +} + +// sortDescIndices orders the indices so the values descend. The sort is +// stable, matching the insertion sort it replaced bit for bit: ties keep +// their original order, so the returned permutation is identical. +func sortDescIndices(s []float64) []int { + idx := make([]int, len(s)) + for i := range idx { + idx[i] = i + } + slices.SortStableFunc(idx, func(a, b int) int { + switch { + case s[a] > s[b]: + return -1 + case s[a] < s[b]: + return 1 + default: + return 0 + } + }) + return idx +} + +// sortAscIndices orders the indices so the values ascend, stable like +// sortDescIndices. +func sortAscIndices(s []float64) []int { + idx := make([]int, len(s)) + for i := range idx { + idx[i] = i + } + slices.SortStableFunc(idx, func(a, b int) int { + switch { + case s[a] < s[b]: + return -1 + case s[a] > s[b]: + return 1 + default: + return 0 + } + }) + return idx +} + +// isIdentityPerm reports whether the permutation is the identity. +func isIdentityPerm(p []int) bool { + for i, v := range p { + if v != i { + return false + } + } + return true +} diff --git a/linalg/decomp3.go b/linalg/decomp3.go new file mode 100644 index 0000000..1c8197e --- /dev/null +++ b/linalg/decomp3.go @@ -0,0 +1,343 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// core.Complex decompositions. The real `Eigen` tridiagonalises a symmetric +// matrix with Householder reflections, a construction that has no +// direct complex analogue; `EigenComplex` instead runs the Jacobi +// iteration, which generalises cleanly to Hermitian matrices through a +// two-step rotation: a complex phase turn that makes the offending +// off-diagonal entry real, then a real plane rotation that zeroes it. +// Every step preserves the spectrum exactly and the iteration inherits +// the real Jacobi convergence argument, at the usual price of +// O(n³) per sweep. +// +// `SVDComplex` builds on `EigenComplex`: the singular vectors of A are +// the eigenvectors of the Hermitian A^H·A, the singular values are +// their square roots, and the left vectors follow as A·v/σ, completed +// through the null space by Gram-Schmidt. Squaring the spectrum costs +// accuracy on very small singular values, which the doc contract +// states plainly. + +// EigenComplex returns the eigenvalues and eigenvectors of a complex +// Hermitian matrix. Values are real, sorted ascending; vectors has +// shape (n, n) with column j the unit eigenvector for values[j]. +// The matrix must be square, complex and Hermitian within a scale- +// relative 1e-12 tolerance. Real symmetric matrices should use the +// faster `Eigen`. +func EigenComplex(a *core.Array) (values, vectors *core.Array, err error) { + if a.Dtype() != core.Complex { + return nil, nil, base.Errf("EigenComplex: needs a complex matrix, got dtype %s", a.Dtype()) + } + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, nil, base.Errf("EigenComplex: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + if n == 0 { + return nil, nil, base.Errf("EigenComplex: zero-sized matrix, got shape %s", base.ShapeText(a.Shape())) + } + h := make([]complex128, n*n) + if a.Dtype() == core.Complex && !a.Strided() && len(a.RawComplexes()) == n*n { + copy(h, a.RawComplexes()) + } else { + for i := range n * n { + h[i] = a.ComplexAt(i) + } + } + if !hermitianOK(h, n) { + return nil, nil, base.Errf("EigenComplex: matrix is not Hermitian within 1e-12 tolerance") + } + // The Jacobi sweep's convergence test is a sum of squared + // magnitudes: outside the safe window it overflows to +Inf, which + // makes every sweep look already converged, or underflows to 0, + // which makes the threshold 0 and stops on the first check, and + // either way the raw diagonal is returned as the spectrum. The + // matrix is moved into the window and the values, which carry the + // scale, are moved back before they are returned; the vectors are + // scale-free. + ws := windowScale(maxMagComplex(h)) + if ws != 1 { + scaleComplexes(h, ws) + } + // The accumulated similarity: columns of v are eigenvectors. + v := eyeComplex(n) + if err := hermitianJacobi(h, v, n); err != nil { + return nil, nil, base.Errf("EigenComplex: %w", err) + } + + vals := make([]float64, n) + for i := range n { + vals[i] = real(h[i*n+i]) + } + idx := sortAscIndices(vals) + outVals := make([]float64, n) + outVecs := make([]complex128, n*n) + for j := range n { + outVals[j] = vals[idx[j]] + for i := range n { + outVecs[i*n+j] = v[i*n+idx[j]] + } + } + if ws != 1 { + unscaleFloats(outVals, ws) + } + valuesArr := floatsToArray(outVals, []int{n}) + vecArr := core.New(core.Complex, []int{n, n}...) + copy(vecArr.RawComplexes(), outVecs) + return valuesArr, vecArr, nil +} + +// hermitianJacobi diagonalises a Hermitian matrix in place, sweeping +// all pairs until the off-diagonal mass is at rounding level. The +// eigenvectors accumulate into v. An exhausted sweep without reaching +// the threshold is an error, never a silently unconverged spectrum. +func hermitianJacobi(h []complex128, v []complex128, n int) error { + norm := 0.0 + for i := range n * n { + norm += real(h[i])*real(h[i]) + imag(h[i])*imag(h[i]) + } + norm = math.Sqrt(norm) + // A purely relative convergence threshold: an absolute floor such + // as max(1, norm) would declare a Hermitian matrix of norm below + // the floor already diagonal and return its raw diagonal, identity + // eigenvectors included. + threshold := 1e-13 * norm + offMass := func() float64 { + off := 0.0 + for p := range n { + for q := p + 1; q < n; q++ { + z := h[p*n+q] + off += real(z)*real(z) + imag(z)*imag(z) + } + } + return math.Sqrt(off) + } + for range 60 { + if offMass() <= threshold { + return nil + } + for p := range n { + for q := p + 1; q < n; q++ { + z := h[p*n+q] + if base.AbsComplex(z) <= threshold { + continue + } + // Step 1: a diagonal phase turn makes the pair entry + // real and positive. + if imag(z) != 0 || real(z) < 0 { + alpha := cmplx.Phase(z) + rot := cmplx.Exp(complex(0, -alpha)) + // The similarity h = Dᴴ·h·D turns the pair entry + // real; V = V·D scales the p-th component of every + // vector. + for j := range n { + h[p*n+j] *= rot + } + for i := range n { + h[i*n+p] *= cmplx.Conj(rot) + } + for i := range n { + v[i*n+p] *= cmplx.Conj(rot) + } + z = h[p*n+q] + } + // Step 2: the real Jacobi rotation that zeroes the now + // real entry, applied to rows and columns p, q. + tau := (real(h[q*n+q]) - real(h[p*n+p])) / (2 * real(z)) + var t float64 + if tau >= 0 { + t = 1 / (tau + math.Sqrt(1+tau*tau)) + } else { + t = -1 / (-tau + math.Sqrt(1+tau*tau)) + } + c := 1 / math.Sqrt(1+t*t) + s := t * c + cr, sr := complex(c, 0), complex(s, 0) + // h = Gᴴ·h·G mixes rows and columns p, q; V = V·G + // mixes the p, q components of every eigenvector. + for j := range n { + hp, hq := h[p*n+j], h[q*n+j] + h[p*n+j] = cr*hp - sr*hq + h[q*n+j] = sr*hp + cr*hq + } + for i := range n { + hp, hq := h[i*n+p], h[i*n+q] + h[i*n+p] = cr*hp - sr*hq + h[i*n+q] = sr*hp + cr*hq + } + for i := range n { + vp, vq := v[i*n+p], v[i*n+q] + v[i*n+p] = cr*vp - sr*vq + v[i*n+q] = sr*vp + cr*vq + } + h[p*n+q] = 0 + h[q*n+p] = 0 + } + } + } + if offMass() <= threshold { + return nil + } + return base.Errf("the Jacobi sweep failed to converge on a %d×%d Hermitian matrix in 60 passes", n, n) +} + +// hermitianOK reports whether every entry satisfies a[i][j] equals +// conj(a[j][i]) within a purely relative tolerance: an absolute floor +// would admit an asymmetry that is a large fraction of a small-scale +// matrix, so the guard would approve a matrix that is not Hermitian +// whatever the size of its entries. +func hermitianOK(h []complex128, n int) bool { + scale := 0.0 + for _, z := range h { + if a := base.AbsComplex(z); a > scale { + scale = a + } + } + tol := 1e-12 * scale + for i := range n { + for j := i; j < n; j++ { + z, m := h[i*n+j], h[j*n+i] + if d := base.AbsComplex(z - cmplx.Conj(m)); d > tol { + return false + } + } + } + return true +} + +func eyeComplex(n int) []complex128 { + out := make([]complex128, n*n) + for i := range n { + out[i*n+i] = 1 + } + return out +} + +// SVDComplex returns the thin singular value decomposition of a +// complex m×n matrix: A = U·diag(Σ)·Vᴴ with Σ sorted descending. U +// has shape (m, n), sigma (n,) and vᴴ (n, n), matching the real `SVD`'s +// shapes. Wide inputs (m < n) decompose the conjugate transpose and +// swap the factors. Because the singular values come from the +// eigenvalues of Aᴴ·A, their accuracy on very small singular values is +// that of a squared condition, fine for rank decisions, weaker for +// resolving near-null directions. +func SVDComplex(a *core.Array) (u, sigma, vH *core.Array, err error) { + if a.Dtype() != core.Complex { + return nil, nil, nil, base.Errf("SVDComplex: needs a complex matrix, got dtype %s", a.Dtype()) + } + if a.NDim() != 2 { + return nil, nil, nil, base.Errf("SVDComplex: needs a 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m == 0 || n == 0 { + return nil, nil, nil, base.Errf("SVDComplex: zero-sized matrix, got shape %s", base.ShapeText(a.Shape())) + } + flat := make([]complex128, m*n) + if a.Dtype() == core.Complex && !a.Strided() && len(a.RawComplexes()) == m*n { + copy(flat, a.RawComplexes()) + } else { + for i := range m * n { + flat[i] = a.ComplexAt(i) + } + } + // The bidiagonalisation's reflector norms are sums of raw squares: + // outside the safe window a tiny matrix has norm 0, which skips the + // reflector and silently zeroes the below-diagonal mass, and a huge + // one overflows to +Inf, which puts NaN in beta and the update. The + // matrix is moved into the window and the singular values, which + // carry the scale, are moved back on return; the unitary factors are + // scale-free. + ws := windowScale(maxMagComplex(flat)) + if ws != 1 { + scaleComplexes(flat, ws) + } + if m < n { + // A = V₂·Σ·U₂ᴴ obtained from the tall decomposition of Aᴴ. + // Shapes mirror the real `SVD`: U (m, m), Σ (m,), Vᴴ (m, n). + ah := transposeConj(flat, m, n) + u2, sigma2, v2, err := svdComplexTall(ah, n, m) + if err != nil { + return nil, nil, nil, err + } + unscaleFloats(sigma2, ws) + uOut := core.New(core.Complex, []int{m, m}...) + copy(uOut.RawComplexes(), v2) + sigOut := floatsToArray(sigma2, []int{m}) + vhOut := core.New(core.Complex, []int{m, n}...) + copy(vhOut.RawComplexes(), transposeConj(u2, n, m)) + return uOut, sigOut, vhOut, nil + } + uMat, sigVals, vMat, err := svdComplexTall(flat, m, n) + if err != nil { + return nil, nil, nil, err + } + unscaleFloats(sigVals, ws) + uOut := core.New(core.Complex, []int{m, n}...) + copy(uOut.RawComplexes(), uMat) + sigOut := floatsToArray(sigVals, []int{n}) + vhOut := core.New(core.Complex, []int{n, n}...) + copy(vhOut.RawComplexes(), transposeConj(vMat, n, n)) + return uOut, sigOut, vhOut, nil +} + +// svdComplexTall decomposes the complex m×n matrix with m ≥ n by +// Golub-Kahan bidiagonalisation and the Golub-Reinsch shifted QR +// iteration on the real bidiagonal: A = U·diag(σ)·Vᴴ with σ sorted +// descending, U thin (m×n) and V square (n×n). +func svdComplexTall(flat []complex128, m, n int) ([]complex128, []float64, []complex128, error) { + work := append([]complex128(nil), flat...) + d, e, u, v := svdBidiagonalise(work, m, n) + if err := svdGolubReinsch(d, e, u, v, m, n); err != nil { + return nil, nil, nil, err + } + // Non-negative singular values, flipping the matching U column. + for i := range n { + if d[i] < 0 { + d[i] = -d[i] + for r := range m { + u[r*m+i] = -u[r*m+i] + } + } + } + // Descending order, permuting the factor columns along. + for i := range n { + big := i + for j := i + 1; j < n; j++ { + if d[j] > d[big] { + big = j + } + } + if big != i { + d[i], d[big] = d[big], d[i] + for r := range m { + u[r*m+i], u[r*m+big] = u[r*m+big], u[r*m+i] + } + for r := range n { + v[r*n+i], v[r*n+big] = v[r*n+big], v[r*n+i] + } + } + } + thin := make([]complex128, m*n) + for i := range m { + copy(thin[i*n:(i+1)*n], u[i*m:i*m+n]) + } + return thin, d, v, nil +} + +func transposeConj(a []complex128, m, n int) []complex128 { + out := make([]complex128, n*m) + for i := range m { + for j := range n { + out[j*m+i] = cmplx.Conj(a[i*n+j]) + } + } + return out +} diff --git a/linalg/decomp3_test.go b/linalg/decomp3_test.go new file mode 100644 index 0000000..76b3c08 --- /dev/null +++ b/linalg/decomp3_test.go @@ -0,0 +1,293 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +func mustComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array { + t.Helper() + a, err := core.FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + return a +} + +// matmulComplex multiplies flat complex matrices a (m×k) by b (k×n). +func matmulComplex(a, b []complex128, m, k, n int) []complex128 { + out := make([]complex128, m*n) + for i := range m { + for p := range k { + aip := a[i*k+p] + for j := range n { + out[i*n+j] += aip * b[p*n+j] + } + } + } + return out +} + +// flatNorm returns the Frobenius norm of a flat complex matrix. +func flatNorm(a []complex128) float64 { + s := 0.0 + for _, z := range a { + s += real(z)*real(z) + imag(z)*imag(z) + } + return math.Sqrt(s) +} + +// TestEigenComplexPauli checks the Hermitian solver on matrices whose +// spectrum is known exactly: a·I + b·σx + c·σy + d·σz has eigenvalues +// a ± √(b²+c²+d²). +func TestEigenComplexPauli(t *testing.T) { + // [[2, 1−i],[1+i, 3]] = 2.5·I + 1·σx + 1·σy + 0.5·σz, so the + // eigenvalues are 2.5 ± 1.5. + h := mustComplexes(t, []complex128{ + 2, complex(1, -1), + complex(1, 1), 3, + }, 2, 2) + values, vectors, err := EigenComplex(h) + if err != nil { + t.Fatalf("EigenComplex: %v", err) + } + want := []float64{1, 4} + for i := range 2 { + if math.Abs(values.FloatAt(i)-want[i]) > 1e-12 { + t.Fatalf("eigenvalue[%d] = %v, want %v", i, values.FloatAt(i), want[i]) + } + } + // Eigenvector residuals ‖H·v − λ·v‖, column by column. + hFlat := []complex128{2, complex(1, -1), complex(1, 1), 3} + for j := range 2 { + col := []complex128{vectors.ComplexAt(0*2 + j), vectors.ComplexAt(1*2 + j)} + hv := matmulComplex(hFlat, col, 2, 2, 1) + for i := range 2 { + res := hv[i] - complex(values.FloatAt(j), 0)*col[i] + if cmplx.Abs(res) > 1e-12 { + t.Fatalf("residual ‖Hv−λv‖[%d,%d] = %v", i, j, cmplx.Abs(res)) + } + } + } + // Unitarity: Vᴴ·V = I. + for i := range 2 { + for j := range 2 { + s := complex(0, 0) + for k := range 2 { + s += cmplx.Conj(vectors.ComplexAt(k*2+i)) * vectors.ComplexAt(k*2+j) + } + want := 0.0 + if i == j { + want = 1 + } + if cmplx.Abs(s-complex(want, 0)) > 1e-12 { + t.Fatalf("VᴴV[%d,%d] = %v, want %v", i, j, s, want) + } + } + } +} + +// TestEigenComplexLarger runs the solver on a 5×5 Hermitian matrix and +// verifies every Ritz pair by residual and every column by +// orthogonality, the properties users actually consume. +func TestEigenComplexLarger(t *testing.T) { + const n = 5 + // H = B + Bᴴ for a pseudorandom complex B, Hermitian by + // construction. + flat := make([]complex128, n*n) + seed := uint64(88172645463325252) + next := func() complex128 { + seed ^= seed << 13 + seed ^= seed >> 7 + seed ^= seed << 17 + return complex(float64(int64(seed%2000)-1000)/1000, float64(int64(seed%2000)-1000)/1000) + } + for i := range n * n { + flat[i] = next() + } + h := make([]complex128, n*n) + for i := range n { + for j := range n { + h[i*n+j] = flat[i*n+j] + cmplx.Conj(flat[j*n+i]) + } + } + hArr := mustComplexes(t, h, n, n) + values, vectors, err := EigenComplex(hArr) + if err != nil { + t.Fatalf("EigenComplex: %v", err) + } + // Ascending order. + for i := 1; i < n; i++ { + if values.FloatAt(i) < values.FloatAt(i-1) { + t.Fatalf("eigenvalues not ascending: %v then %v", values.FloatAt(i-1), values.FloatAt(i)) + } + } + for j := range n { + // Residual column: H·v_j − λ_j·v_j. + col := make([]complex128, n) + for i := range n { + col[i] = vectors.ComplexAt(i*n + j) + } + hv := matmulComplex(h, col, n, n, 1) + for i := range n { + res := hv[i] - complex(values.FloatAt(j), 0)*col[i] + if cmplx.Abs(res) > 1e-10*(1+math.Abs(values.FloatAt(j))) { + t.Fatalf("residual [%d,%d] = %v", i, j, cmplx.Abs(res)) + } + } + } +} + +// TestEigenComplexRejectsInvalid pins the input contract. +func TestEigenComplexRejectsInvalid(t *testing.T) { + real := mustFloats(t, []float64{1, 0, 0, 1}, 2, 2) + if _, _, err := EigenComplex(real); err == nil { + t.Fatal("expected an error for a real input") + } + nonsq := mustComplexes(t, []complex128{1, 0, 0, 1, 0, 0}, 2, 3) + if _, _, err := EigenComplex(nonsq); err == nil { + t.Fatal("expected an error for a non-square matrix") + } + asym := mustComplexes(t, []complex128{1, 2, 0, 1}, 2, 2) + if _, _, err := EigenComplex(asym); err == nil { + t.Fatal("expected an error for a non-Hermitian matrix") + } +} + +// TestSVDComplexKnown checks the decomposition on a rank-1 outer +// product with exactly known singular values, plus the reconstruction, +// orthogonality and ordering contracts on tall and wide inputs. +func TestSVDComplexKnown(t *testing.T) { + // A = u·vᵀ with ‖u‖=1, ‖v‖=√2, so σ = {√2, 0}. + rt2 := 1 / math.Sqrt2 + u := []complex128{complex(rt2, 0), complex(0, rt2)} + v := []float64{1, 1} + a := make([]complex128, 4) + for i := range 2 { + for j := range 2 { + a[i*2+j] = u[i] * complex(v[j], 0) + } + } + uOut, sigma, vH, err := SVDComplex(mustComplexes(t, a, 2, 2)) + if err != nil { + t.Fatalf("SVDComplex: %v", err) + } + if math.Abs(sigma.FloatAt(0)-math.Sqrt2) > 1e-12 { + t.Fatalf("σ₁ = %v, want √2", sigma.FloatAt(0)) + } + if sigma.FloatAt(1) > 1e-12 { + t.Fatalf("σ₂ = %v, want 0", sigma.FloatAt(1)) + } + // Reconstruction A ≈ U·Σ·Vᴴ. + recon := make([]complex128, 4) + for i := range 2 { + for j := range 2 { + s := complex(0, 0) + for k := range 2 { + s += uOut.ComplexAt(i*2+k) * complex(sigma.FloatAt(k), 0) * + vH.ComplexAt(k*2+j) + } + recon[i*2+j] = s + } + } + if d := flatNorm(subComplex(a, recon)); d > 1e-12 { + t.Fatalf("reconstruction error %v", d) + } +} + +// TestSVDComplexTallAndWide checks the tall and the wide path on the +// same content: reconstruction, orthogonality of both factors and +// descending singular values. +func TestSVDComplexTallAndWide(t *testing.T) { + build := func(m, n int) *core.Array { + flat := make([]complex128, m*n) + seed := uint64(11400714819323198485) + for i := range m * n { + seed ^= seed << 13 + seed ^= seed >> 7 + seed ^= seed << 17 + flat[i] = complex(float64(int64(seed%400)-200)/100, float64(int64(seed%400)-200)/100) + } + return mustComplexes(t, flat, m, n) + } + check := func(t *testing.T, a *core.Array) { + m, n := a.Shape()[0], a.Shape()[1] + u, sigma, vH, err := SVDComplex(a) + if err != nil { + t.Fatalf("SVDComplex(%dx%d): %v", m, n, err) + } + // Shapes mirror the real SVD: U (m, min), Σ (min,), Vᴴ (m, n). + rank := min(m, n) + if u.Shape()[0] != m || u.Shape()[1] != rank { + t.Fatalf("U shape %s, want [%d %d]", base.ShapeText(u.Shape()), m, rank) + } + if sigma.Len() != rank || vH.Shape()[0] != rank || vH.Shape()[1] != n { + t.Fatalf("sigma len %d, Vᴴ shape %s", sigma.Len(), base.ShapeText(vH.Shape())) + } + for i := 1; i < rank; i++ { + if sigma.FloatAt(i) > sigma.FloatAt(i-1)+1e-12 { + t.Fatalf("singular values not descending: %v then %v", sigma.FloatAt(i-1), sigma.FloatAt(i)) + } + } + // U·Σ·Vᴴ. + recon := make([]complex128, m*n) + for i := range m { + for j := range n { + s := complex(0, 0) + for k := range rank { + s += u.ComplexAt(i*rank+k) * complex(sigma.FloatAt(k), 0) * + vH.ComplexAt(k*n+j) + } + recon[i*n+j] = s + } + } + aFlat := make([]complex128, m*n) + for i := range m * n { + aFlat[i] = a.ComplexAt(i) + } + if d := flatNorm(subComplex(aFlat, recon)); d > 1e-9*float64(m) { + t.Fatalf("%dx%d reconstruction error %v", m, n, d) + } + // Orthogonality of U's columns and of Vᴴᴴ (i.e. VᴴV). + for i := range rank { + for j := range rank { + su := complex(0, 0) + for k := range m { + su += cmplx.Conj(u.ComplexAt(k*rank+i)) * u.ComplexAt(k*rank+j) + } + // Vᴴ has orthonormal ROWS in every convention. + sv2 := complex(0, 0) + for k := range n { + sv2 += vH.ComplexAt(i*n+k) * cmplx.Conj(vH.ComplexAt(j*n+k)) + } + want := 0.0 + if i == j { + want = 1 + } + if cmplx.Abs(su-complex(want, 0)) > 1e-9 { + t.Fatalf("UᴴU[%d,%d] = %v", i, j, su) + } + if cmplx.Abs(sv2-complex(want, 0)) > 1e-9 { + t.Fatalf("(VᴴVᴴ*)[%d,%d] = %v", i, j, sv2) + } + } + } + } + t.Run("tall", func(t *testing.T) { check(t, build(6, 4)) }) + t.Run("wide", func(t *testing.T) { check(t, build(4, 6)) }) +} + +// subComplex subtracts two flat complex matrices of equal length. +func subComplex(a, b []complex128) []complex128 { + out := make([]complex128, len(a)) + for i := range a { + out[i] = a[i] - b[i] + } + return out +} diff --git a/linalg/decomp_test.go b/linalg/decomp_test.go new file mode 100644 index 0000000..0effd7b --- /dev/null +++ b/linalg/decomp_test.go @@ -0,0 +1,479 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestSVDReconstruction checks that A = U · Σ · Vᵀ reconstructs A +// for several rank profiles (full-rank square, tall, wide, and a +// rank-1 matrix where the one-sided Jacobi prototype failed). +func TestSVDReconstruction(t *testing.T) { + cases := []struct { + name string + vals []float64 + m, n int + hasSwap bool + }{ + { + name: "square_full_rank", + vals: []float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, + m: 3, n: 3, + }, + { + name: "tall_rank2", + vals: []float64{1, 1, 1, 2, 2, 2, 3, 3, 3, 1, 2, 3}, + m: 4, n: 3, + }, + { + name: "wide_rank2", + vals: []float64{1, 2, 3, 4, 5, 1, 2, 3, 4, 5, 1, 2}, + m: 3, n: 4, + hasSwap: true, + }, + { + name: "rank1", + vals: []float64{2, 4, 6, 8, 10, 12}, + m: 3, n: 2, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + a := mustFromFloats(t, tc.vals, tc.m, tc.n) + uOut, sigma, vt, err := SVD(a) + if err != nil { + t.Fatalf("SVD: %v", err) + } + // Thin SVD: U is (m, k), Vᵀ is (k, n) where k = min(m, n). + k := min(tc.m, tc.n) + wantU := [2]int{tc.m, k} + wantVt := [2]int{k, tc.n} + if got := uOut.Shape(); got[0] != wantU[0] || got[1] != wantU[1] { + t.Fatalf("U shape: got %v want %v", got, wantU) + } + if got := vt.Shape(); got[0] != wantVt[0] || got[1] != wantVt[1] { + t.Fatalf("Vᵀ shape: got %v want %v", got, wantVt) + } + wantSigma := k + if got := sigma.Shape()[0]; got != wantSigma { + t.Fatalf("sigma shape: got %d want %d", got, wantSigma) + } + aRecon := matmulSVDReconstruct(uOut, sigma, vt, tc.m, tc.n) + if !matricesClose(aRecon, denseFloats(a, tc.m, tc.n), tc.m, tc.n, 1e-9) { + t.Fatalf("SVD reconstruction [%s]: max diff > 1e-9\norig=%v\nrecon=%v", + tc.name, a, aRecon) + } + if !matricesThinOrthogonal(uOut.RawFloats(), uOut.Shape()[0], U_SVD_COLS(uOut), 1e-9) { + t.Errorf("U not column-orthonormal in %s (shape=%v)", tc.name, uOut.Shape()) + } + for i := 1; i < wantSigma; i++ { + if sigma.RawFloats()[i-1] < sigma.RawFloats()[i] { + t.Errorf("Σ not sorted descending: %v", sigma.RawFloats()) + break + } + } + }) + } +} + +// TestSVDRankDeficient checks the SVD on a deliberately rank-deficient +// matrix, the case the one-sided Jacobi prototype failed on. +func TestSVDRankDeficient(t *testing.T) { + u := []float64{1, 2, 3, 4} + v := []float64{1, -1, 1} + a := outerProduct(u, v, 4, 3) + uOut, sigma, vt, err := SVD(a) + if err != nil { + t.Fatalf("SVD: %v", err) + } + if len(sigma.RawFloats()) != 3 { + t.Fatalf("σ shape: %v", sigma.Shape()) + } + if !(sigma.RawFloats()[0] > 1e-9) { + t.Errorf("σ₀: got %g, want > 1e-9", sigma.RawFloats()[0]) + } + if math.Abs(sigma.RawFloats()[1]) > 1e-9 || math.Abs(sigma.RawFloats()[2]) > 1e-9 { + t.Errorf("σ₁ and σ₂ should be ~0, got %g %g", + sigma.RawFloats()[1], sigma.RawFloats()[2]) + } + aRecon := matmulSVDReconstruct(uOut, sigma, vt, 4, 3) + if !matricesClose(aRecon, denseFloats(a, 4, 3), 4, 3, 1e-9) { + t.Fatalf("rank-deficient SVD reconstruction: max diff > 1e-9\norig=%v\nrecon=%v", + a, aRecon) + } +} + +// TestEigenDiagonal checks that a diagonal matrix returns the +// diagonal entries as eigenvalues. +func TestEigenDiagonal(t *testing.T) { + a := mustFromFloats(t, []float64{ + 3, 0, 0, + 0, 1, 0, + 0, 0, 5, + }, 3, 3) + vals, vecs, err := Eigen(a) + if err != nil { + t.Fatalf("Eigen: %v", err) + } + if got := vals.Shape()[0]; got != 3 { + t.Fatalf("vals shape: %v", vals.Shape()) + } + want := []float64{1, 3, 5} + for i, w := range want { + if math.Abs(vals.RawFloats()[i]-w) > 1e-9 { + t.Errorf("eigenvalue[%d]: got %g want %g", i, vals.RawFloats()[i], w) + } + } + recon := eigenReconstruct(vecs, vals, 3) + wantMat := denseFloats(a, 3, 3) + if !matricesClose(recon, wantMat, 3, 3, 1e-9) { + t.Fatalf("Eigen reconstruction: max diff > 1e-9\nwant=%v\ngot=%v", + wantMat, recon) + } +} + +// TestEigenAsymmetric verifies the asymmetric input path errors. +func TestEigenAsymmetric(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 2, + 3, 4, + }, 2, 2) + if _, _, err := Eigen(a); err == nil { + t.Fatal("Eigen: expected error on asymmetric matrix") + } +} + +// TestPinverse checks Moore-Penrose pseudoinverse on the canonical +// example and verifies A · A⁺ · A = A on a rank-deficient matrix. +func TestPinverse(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + inv, err := Pinverse(a, 0) + if err != nil { + t.Fatalf("Pinverse: %v", err) + } + if !matricesClose(denseFloats(inv, 2, 2), []float64{-2, 1, 1.5, -0.5}, 2, 2, 1e-9) { + t.Fatalf("Pinverse of invertible: got %v", inv) + } + u := []float64{1, 2, 3} + v := []float64{2, -1} + aRank1 := outerProduct(u, v, 3, 2) + pinv, err := Pinverse(aRank1, 0) + if err != nil { + t.Fatalf("Pinverse(rank1): %v", err) + } + _ = pinv + recon := matmulMul(matmulMul(denseFloats(aRank1, 3, 2), denseFloats(pinv, 2, 3), 3, 2, 3), denseFloats(aRank1, 3, 2), 3, 3, 2) + want := denseFloats(aRank1, 3, 2) + if !matricesClose(recon, want, 3, 2, 1e-9) { + t.Fatalf("A · A⁺ · A != A for rank-1 input") + } +} + +// TestPinverseWideThinShapes pins the aspect-ratio plumbing of the +// pseudoinverse: a wide (m < n) input used to read the thin Vᵀ as an +// (n, n) matrix and panic past its payload. Values ride on the SVD +// track; the shapes must hold on their own. +func TestPinverseWideThinShapes(t *testing.T) { + diag := mustFromFloats(t, []float64{ + 3, 0, 0, + 0, 2, 0, + }, 2, 3) + inv, err := Pinverse(diag, 0) + if err != nil { + t.Fatalf("Pinverse wide diagonal: %v", err) + } + if inv.Shape()[0] != 3 || inv.Shape()[1] != 2 { + t.Fatalf("Pinverse wide diagonal shape: %v, want (3, 2)", inv.Shape()) + } + + general := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + inv2, err := Pinverse(general, 0) + if err != nil { + t.Fatalf("Pinverse wide general: %v", err) + } + if inv2.Shape()[0] != 3 || inv2.Shape()[1] != 2 { + t.Fatalf("Pinverse wide general shape: %v, want (3, 2)", inv2.Shape()) + } +} + +// TestPinverseWideDiagonalValues checks the wide-matrix values against +// the exact pseudoinverse once the SVD track lands. +func TestPinverseWideDiagonalValues(t *testing.T) { + a := mustFromFloats(t, []float64{ + 3, 0, 0, + 0, 2, 0, + }, 2, 3) + inv, err := Pinverse(a, 0) + if err != nil { + t.Fatalf("Pinverse wide: %v", err) + } + want := []float64{1.0 / 3, 0, 0, 0.5, 0, 0} + got := denseFloats(inv, 3, 2) + for i, w := range want { + if diff := got[i] - w; diff > 1e-12 || diff < -1e-12 { + t.Fatalf("Pinverse wide [%d]: got %v, want %v", i, got[i], w) + } + } +} + +// TestMatrixRankProperties checks rank on a couple of cases. +func TestMatrixRankProperties(t *testing.T) { + // Identity padded with a zero row: rows [1,0],[0,1],[0,0]. + a := mustFromFloats(t, []float64{1, 0, 0, 1, 0, 0}, 3, 2) + if got, err := MatrixRank(a, 0); err != nil || got != 2 { + t.Errorf("MatrixRank(full-rank): got %d err %v, want 2", got, err) + } + // Every row a multiple of [1,2]: rank 1. + rank1 := mustFromFloats(t, []float64{1, 2, 2, 4, 3, 6}, 3, 2) + if got, err := MatrixRank(rank1, 0); err != nil || got != 1 { + t.Errorf("MatrixRank(rank-1): got %d err %v, want 1", got, err) + } +} + +// TestCondProperties checks condition number on a diagonal matrix +// (singular values are the diagonal entries). +func TestCondProperties(t *testing.T) { + a := mustFromFloats(t, []float64{ + 3, 0, + 0, 0.5, + }, 2, 2) + c, err := Cond(a, 0) + if err != nil { + t.Fatalf("Cond: %v", err) + } + if math.Abs(c-6) > 1e-9 { + t.Errorf("Cond([[3, 0], [0, 0.5]]): got %g want 6", c) + } + sing := mustFromFloats(t, []float64{ + 2, 0, + 0, 0, + }, 2, 2) + c, err = Cond(sing, 0) + if err != nil { + t.Fatalf("Cond(singular): %v", err) + } + if !math.IsInf(c, 1) { + t.Errorf("Cond(singular): got %g want +Inf", c) + } +} + +// TestSVDWideTransposed checks that a wide matrix (m < n) decomposes +// correctly with the transpose swap. +func TestSVDWideTransposed(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + }, 3, 4) + uOut, sigma, vt, err := SVD(a) + if err != nil { + t.Fatalf("SVD: %v", err) + } + if uOut.Shape()[0] != 3 || uOut.Shape()[1] != 3 { + t.Fatalf("U shape: got %v want (3, 3)", uOut.Shape()) + } + if vt.Shape()[0] != 3 || vt.Shape()[1] != 4 { + t.Fatalf("Vᵀ shape: got %v want (3, 4)", vt.Shape()) + } + if sigma.Shape()[0] != 3 { + t.Fatalf("Σ shape: got %v want (3,)", sigma.Shape()) + } + recon := matmulSVDReconstruct(uOut, sigma, vt, 3, 4) + if !matricesClose(recon, denseFloats(a, 3, 4), 3, 4, 1e-9) { + t.Fatalf("wide SVD reconstruction: max diff > 1e-9") + } +} + +// --- helpers (test-only) --- + +// U_SVD_COLS returns the number of columns of a thin SVD U +// (min(m, n)), the rank-n bound the orthogonal columns reach. +func U_SVD_COLS(u *core.Array) int { + return u.Shape()[1] +} + +func matmulSVDReconstruct(u, sigma, vt *core.Array, m, n int) []float64 { + // Thin SVD: U is (m, k); Σ is (k,); Vᵀ is (k, n) where k = min(m, n). + k := min(m, n) + uMat := denseFloats(u, m, k) + vtMat := denseFloats(vt, k, n) + sVals := sigma.RawFloats() + out := make([]float64, m*n) + for i := range m { + for j := range n { + s := 0.0 + for kk := range k { + s += uMat[i*k+kk] * sVals[kk] * vtMat[kk*n+j] + } + out[i*n+j] = s + } + } + return out +} + +func eigenReconstruct(q, vals *core.Array, n int) []float64 { + qMat := denseFloats(q, n, n) + out := make([]float64, n*n) + for i := range n { + for j := range n { + s := 0.0 + for k := range n { + s += qMat[i*n+k] * vals.RawFloats()[k] * qMat[j*n+k] + } + out[i*n+j] = s + } + } + return out +} + +func outerProduct(u, v []float64, m, n int) *core.Array { + out := make([]float64, m*n) + for i := range m { + for j := range n { + out[i*n+j] = u[i] * v[j] + } + } + return floatsToArray(out, []int{m, n}) +} + +func matmulMul(a, b []float64, rows, mid, cols int) []float64 { + out := make([]float64, rows*cols) + for i := range rows { + for j := range cols { + s := 0.0 + for k := range mid { + s += a[i*mid+k] * b[k*cols+j] + } + out[i*cols+j] = s + } + } + return out +} + +func matricesClose(a, b []float64, m, n int, tol float64) bool { + if len(a) != m*n || len(b) != m*n { + return false + } + for i := 0; i < m*n; i++ { + if math.Abs(a[i]-b[i]) > tol { + return false + } + } + return true +} + +// matricesThinOrthogonal checks that a (m, n) matrix a satisfies +// aᵀ a = I_n. Used for the thin SVD factor U. +func matricesThinOrthogonal(a []float64, m, n int, tol float64) bool { + for i := range n { + for j := range n { + s := 0.0 + want := 0.0 + if i == j { + want = 1 + } + for k := range m { + s += a[k*n+i] * a[k*n+j] + } + if math.Abs(s-want) > tol { + return false + } + } + } + return true +} + +// TestLeastSquaresCollinear pins the rank guard: exactly collinear +// columns are refused, and near-collinear ones whose R pivot sits at +// the rounding floor of max|R| are refused too, where the old +// exact-zero check returned a silently huge x. +func TestLeastSquaresCollinear(t *testing.T) { + // Exactly collinear, integer dtype: the historical case. + ints, err := core.FromInts([]int64{1, 2, 2, 4}, 2, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + if _, err := LeastSquares(ints, mustFloats(t, []float64{3, 7})); err == nil { + t.Fatal("exactly collinear integer columns: want an error") + } + // Near-collinear float columns: the second column differs from the + // first by 2^-45, so its R pivot (2^-45) sits below the relative + // rank floor n·eps·max|R| = 2·eps·2^10, where back-substitution + // would amplify it into a ~1e13 answer. Powers of two keep every + // elimination step exact, so the trip is deterministic. + a := mustFloats(t, []float64{ + 1024, 1024, + 0, math.Ldexp(1, -45), + 0, 0, + }, 3, 2) + if _, err := LeastSquares(a, mustFloats(t, []float64{1, 1, 1})); err == nil { + t.Fatal("near-collinear columns: want an error") + } + // A well-conditioned system of the same shape still solves. + good := mustFloats(t, []float64{ + 1, 1, + 1, 2, + 1, 3, + }, 3, 2) + x, err := LeastSquares(good, mustFloats(t, []float64{2, 3, 4})) + if err != nil { + t.Fatalf("LeastSquares full-rank: %v", err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-9 || math.Abs(x.FloatAt(1)-1) > 1e-9 { + t.Fatalf("x = (%.12g, %.12g), want (1, 1)", x.FloatAt(0), x.FloatAt(1)) + } +} + +// TestCondZeroMatrix pins the zero-matrix contract: no usable inverse +// direction, so the condition number is +Inf, not 0. +func TestCondZeroMatrix(t *testing.T) { + c, err := Cond(mustFloats(t, []float64{0, 0, 0, 0}, 2, 2), 0) + if err != nil { + t.Fatalf("Cond: %v", err) + } + if !math.IsInf(c, 1) { + t.Fatalf("Cond(zero matrix) = %g, want +Inf", c) + } +} + +// TestCholeskyUpdateRejectsUpperTriangle pins the factor validation: +// a nonzero strict upper triangle is refused, matching what the error +// message has always claimed. +func TestCholeskyUpdateRejectsUpperTriangle(t *testing.T) { + l := mustFloats(t, []float64{ + 2, 1, + 0, 1, + }, 2, 2) + if _, err := CholeskyUpdate(l, mustFloats(t, []float64{1, 1})); err == nil { + t.Fatal("CholeskyUpdate with a nonzero upper triangle: want an error") + } + if _, err := CholeskyDowndate(l, mustFloats(t, []float64{1, 1})); err == nil { + t.Fatal("CholeskyDowndate with a nonzero upper triangle: want an error") + } +} + +// TestHouseholderVectorIntoLargeScale pins the overflow-safe reflector: +// entries near 1e154 keep a finite norm and beta instead of squaring +// their way to +Inf. +func TestHouseholderVectorIntoLargeScale(t *testing.T) { + x := []float64{1e154, 1e154, 1e154} + dst := make([]float64, 3) + hh := householderVectorInto(dst, x) + if math.IsInf(hh.beta, 0) || math.IsNaN(hh.beta) { + t.Fatalf("beta = %v, want finite", hh.beta) + } + for i, v := range hh.v { + if math.IsInf(v, 0) || math.IsNaN(v) { + t.Fatalf("v[%d] = %v, want finite", i, v) + } + } + // The zero vector still answers a zero-beta identity reflector. + empty := householderVectorInto(make([]float64, 2), []float64{0, 0}) + if empty.beta != 0 { + t.Fatalf("beta = %v for a zero vector, want 0", empty.beta) + } +} diff --git a/linalg/decompositions_test.go b/linalg/decompositions_test.go new file mode 100644 index 0000000..ff35cab --- /dev/null +++ b/linalg/decompositions_test.go @@ -0,0 +1,693 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/rand/v2" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +func TestQR(t *testing.T) { + // A 3×2 matrix with a known QR decomposition. + a := mustFromFloats(t, []float64{ + 12, -51, + 6, 167, + -4, 24, + }, 3, 2) + q, r, err := QR(a) + if err != nil { + t.Fatal(err) + } + if q.Shape()[0] != 3 || q.Shape()[1] != 3 { + t.Errorf("QR Q shape: %v", q.Shape()) + } + if r.Shape()[0] != 3 || r.Shape()[1] != 2 { + t.Errorf("QR R shape: %v", r.Shape()) + } + // QᵀQ should be I. + for i := range 3 { + for j := range 3 { + var want float64 + if i == j { + want = 1 + } + // QᵀQ: (Qᵀ Q)[i,j] = sum_k Q[k,i] * Q[k,j]. + dot := 0.0 + for k := range 3 { + qi, _ := core.FloatAt(q, k, i) + qj, _ := core.FloatAt(q, k, j) + dot += qi * qj + } + if math.Abs(dot-want) > 1e-9 { + t.Errorf("QᵀQ[%d,%d]: got %v, want %v", i, j, dot, want) + } + } + } + // R should be upper triangular (entries below the diagonal ~0). + for i := 1; i < 3; i++ { + for j := 0; j < minInt(i, 2); j++ { + v, _ := core.FloatAt(r, i, j) + if math.Abs(v) > 1e-9 { + t.Errorf("R[%d,%d] should be ~0, got %v", i, j, v) + } + } + } + // Q * R should reconstruct A. + recon := mustFromFloats(t, []float64{0, 0, 0, 0, 0, 0}, 3, 2) + for i := range 3 { + for j := range 2 { + sum := 0.0 + for k := range 3 { + qik, _ := core.FloatAt(q, i, k) + rkj, _ := core.FloatAt(r, k, j) + sum += qik * rkj + } + recon.RawFloats()[i*2+j] = sum + } + } + for i := range 3 { + for j := range 2 { + got, _ := core.FloatAt(recon, i, j) + orig, _ := core.FloatAt(a, i, j) + if math.Abs(got-orig) > 1e-9 { + t.Errorf("QR reconstruction [%d,%d]: got %v, want %v", i, j, got, orig) + } + } + } + // Square matrix case. + a2 := mustFromFloats(t, []float64{ + 2, 1, + 1, 3, + }, 2, 2) + q2, r2, err := QR(a2) + if err != nil { + t.Fatal(err) + } + if q2.Shape()[0] != 2 || q2.Shape()[1] != 2 { + t.Errorf("QR square Q shape: %v", q2.Shape()) + } + if r2.Shape()[0] != 2 || r2.Shape()[1] != 2 { + t.Errorf("QR square R shape: %v", r2.Shape()) + } + // Shape errors. + bad, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if _, _, err := QR(bad); err == nil { + t.Error("QR: expected error for m < n") + } + _, _, err = QR(mustFromFloats(t, []float64{1, 2, 3, 4}, 4)) // 1-D + if err == nil { + t.Error("QR: expected error for 1-D input") + } +} + +func TestCholesky(t *testing.T) { + // A symmetric positive definite matrix. + a := mustFromFloats(t, []float64{ + 4, 12, -16, + 12, 37, -43, + -16, -43, 98, + }, 3, 3) + l, err := Cholesky(a) + if err != nil { + t.Fatal(err) + } + if l.Shape()[0] != 3 || l.Shape()[1] != 3 { + t.Errorf("Cholesky L shape: %v", l.Shape()) + } + // L * Lᵀ should reconstruct A. + for i := range 3 { + for j := range 3 { + sum := 0.0 + for k := 0; k <= minInt(i, j); k++ { + lik, _ := core.FloatAt(l, i, k) + ljk, _ := core.FloatAt(l, j, k) + sum += lik * ljk + } + orig, _ := core.FloatAt(a, i, j) + if math.Abs(sum-orig) > 1e-9 { + t.Errorf("Cholesky reconstruction [%d,%d]: got %v, want %v", i, j, sum, orig) + } + } + } + // Non-PD matrix: singular means an error. + notPD := mustFromFloats(t, []float64{1, 2, 2, 1}, 2, 2) + if _, err := Cholesky(notPD); err == nil { + t.Error("Cholesky: expected error for non-PD input") + } +} + +func TestSVD(t *testing.T) { + // 3×2 rank-2 matrix with known singular values. + a := mustFromFloats(t, []float64{ + 1, 0, + 0, 2, + 0, 0, + }, 3, 2) + u, sigma, vt, err := SVD(a) + if err != nil { + t.Fatalf("SVD: %v", err) + } + if u.Shape()[0] != 3 || u.Shape()[1] != 2 { + t.Errorf("U shape: %v", u.Shape()) + } + if sigma.Shape()[0] != 2 { + t.Errorf("Σ shape: %v", sigma.Shape()) + } + if vt.Shape()[0] != 2 || vt.Shape()[1] != 2 { + t.Errorf("Vᵀ shape: %v", vt.Shape()) + } + // Singular values should be {2, 1} in descending order. + if math.Abs(sigma.RawFloats()[0]-2) > 1e-9 { + t.Errorf("σ₀: got %g want 2", sigma.RawFloats()[0]) + } + if math.Abs(sigma.RawFloats()[1]-1) > 1e-9 { + t.Errorf("σ₁: got %g want 1", sigma.RawFloats()[1]) + } + if _, _, _, err := SVD(mustFromFloats(t, []float64{1, 2}, 2)); err == nil { + t.Error("SVD: expected error for 1-D input") + } +} + +func TestEigen(t *testing.T) { + // Simple diagonal symmetric matrix. + a := mustFromFloats(t, []float64{ + 4, 0, + 0, 9, + }, 2, 2) + vals, vecs, err := Eigen(a) + if err != nil { + t.Fatalf("Eigen: %v", err) + } + if len(vals.RawFloats()) != 2 { + t.Fatalf("vals shape: %v", vals.Shape()) + } + if math.Abs(vals.RawFloats()[0]-4) > 1e-9 || math.Abs(vals.RawFloats()[1]-9) > 1e-9 { + t.Errorf("eigenvalues: got %v want [4, 9]", vals.RawFloats()) + } + if vecs.Shape()[0] != 2 || vecs.Shape()[1] != 2 { + t.Errorf("vecs shape: %v", vecs.Shape()) + } + // 1-D input errors. + if _, _, err := Eigen(mustFromFloats(t, []float64{1, 2}, 2)); err == nil { + t.Error("Eigen: expected error for 1-D input") + } +} + +func TestLeastSquares(t *testing.T) { + // Solve [[3, 1], [1, 2]] x = [9, 8]: solution is x = [2, 3]. + a := mustFromFloats(t, []float64{ + 3, 1, + 1, 2, + }, 2, 2) + b := mustFromFloats(t, []float64{9, 8}, 2) + x, err := LeastSquares(a, b) + if err != nil { + t.Fatal(err) + } + if x.NDim() != 1 || x.Len() != 2 { + t.Errorf("LeastSquares x shape: %v", x.Shape()) + } + v0, _ := core.FloatAt(x, 0) + v1, _ := core.FloatAt(x, 1) + if math.Abs(v0-2) > 1e-9 { + t.Errorf("LeastSquares [0]: got %v, want 2", v0) + } + if math.Abs(v1-3) > 1e-9 { + t.Errorf("LeastSquares [1]: got %v, want 3", v1) + } + // Over-determined system: 3×2 with rank 2. + overA := mustFromFloats(t, []float64{ + 1, 1, + 1, 2, + 1, 3, + }, 3, 2) + overB := mustFromFloats(t, []float64{1, 2, 2}, 3) + xOver, err := LeastSquares(overA, overB) + if err != nil { + t.Fatal(err) + } + // Solution minimises ||Ax - b||₂. + if xOver.Len() != 2 { + t.Errorf("LeastSquares over: shape = %v", xOver.Shape()) + } + // Shape errors. + if _, err := LeastSquares(mustFromFloats(t, []float64{1, 2, 3}, 3), overB); err == nil { + t.Error("LeastSquares: expected error for 1-D 'a'") + } + underA := mustFromFloats(t, []float64{ + 1, 2, 3, + 4, 5, 6, + }, 2, 3) + if _, err := LeastSquares(underA, overB); err == nil { + t.Error("LeastSquares: expected error for m < n") + } +} + +func TestMatrixRank(t *testing.T) { + // Identity has full rank. + id, _ := core.Identity(core.Float, 4) + got, err := MatrixRank(id, 0) + if err != nil { + t.Fatalf("MatrixRank: %v", err) + } + if got != 4 { + t.Errorf("MatrixRank(I_4): got %d want 4", got) + } + if _, err := MatrixRank(mustFromFloats(t, []float64{1, 2}, 2), 0); err == nil { + t.Error("MatrixRank: expected error for 1-D input") + } +} + +func TestCond(t *testing.T) { + id, _ := core.Identity(core.Float, 3) + c, err := Cond(id, 0) + if err != nil { + t.Fatalf("Cond: %v", err) + } + if math.Abs(c-1) > 1e-9 { + t.Errorf("Cond(I_3): got %g want 1", c) + } + if _, err := Cond(mustFromFloats(t, []float64{1, 2}, 2), 0); err == nil { + t.Error("Cond: expected error for 1-D input") + } +} + +func TestEinsumMatmul(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + b := mustFromFloats(t, []float64{5, 6, 7, 8}, 2, 2) + got, err := core.Einsum("ij,jk->ik", a, b) + if err != nil { + t.Fatal(err) + } + // Result should be MatMul(a, b) = [[19, 22], [43, 50]]. + expect := mustFromFloats(t, []float64{19, 22, 43, 50}, 2, 2) + for i := range 4 { + v, _ := core.FloatAt(got, i/2, i%2) + w, _ := core.FloatAt(expect, i/2, i%2) + if math.Abs(v-w) > 1e-9 { + t.Errorf("einsum matmul [%d]: got %v, want %v", i, v, w) + } + } +} + +func TestEinsumDot(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3}, 3) + b := mustFromFloats(t, []float64{4, 5, 6}, 3) + got, err := core.Einsum("i,i->", a, b) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0) + if math.Abs(v-32) > 1e-9 { + t.Errorf("einsum dot: got %v, want 32", v) + } +} + +func TestEinsumTranspose(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + got, err := core.Einsum("ij->ji", a) + if err != nil { + t.Fatal(err) + } + // Transpose: [[1,2],[3,4]] -> [[1,3],[2,4]] + v00, _ := core.FloatAt(got, 0, 0) + v01, _ := core.FloatAt(got, 0, 1) + v10, _ := core.FloatAt(got, 1, 0) + v11, _ := core.FloatAt(got, 1, 1) + if v00 != 1 || v01 != 3 || v10 != 2 || v11 != 4 { + t.Errorf("einsum transpose: got [[%v,%v],[%v,%v]], want [[1,3],[2,4]]", v00, v01, v10, v11) + } +} + +func TestEinsumDiagonal(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 2, 3, + 4, 5, 6, + 7, 8, 9, + }, 3, 3) + got, err := core.Einsum("ii->i", a) + if err != nil { + t.Fatal(err) + } + v0, _ := core.FloatAt(got, 0) + v1, _ := core.FloatAt(got, 1) + v2, _ := core.FloatAt(got, 2) + if v0 != 1 || v1 != 5 || v2 != 9 { + t.Errorf("einsum diagonal: got %v %v %v, want 1 5 9", v0, v1, v2) + } +} + +func TestEinsumTrace(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 2, 3, + 4, 5, 6, + 7, 8, 9, + }, 3, 3) + got, err := core.Einsum("ii->", a) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0) + if v != 15 { // 1 + 5 + 9 + t.Errorf("einsum trace: got %v, want 15", v) + } +} + +func TestEinsumSum(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + got, err := core.Einsum("ij->", a) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0) + if v != 10 { + t.Errorf("einsum sum: got %v, want 10", v) + } +} + +func TestEinsumElementWiseMul(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + b := mustFromFloats(t, []float64{5, 6, 7, 8}, 2, 2) + got, err := core.Einsum("ij,ij->ij", a, b) + if err != nil { + t.Fatal(err) + } + for i, w := range []float64{5, 12, 21, 32} { + v, _ := core.FloatAt(got, i/2, i%2) + if v != w { + t.Errorf("einsum elwise [%d]: got %v, want %v", i, v, w) + } + } +} + +func TestEinsumOuter(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2}, 2) + b := mustFromFloats(t, []float64{3, 4, 5}, 3) + got, err := core.Einsum("i,j->ij", a, b) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 2 || got.Shape()[1] != 3 { + t.Errorf("einsum outer shape: %v", got.Shape()) + } + // [1*3, 1*4, 1*5; 2*3, 2*4, 2*5] = [[3,4,5],[6,8,10]] + expect := []float64{3, 4, 5, 6, 8, 10} + for i, w := range expect { + v, _ := core.FloatAt(got, i/3, i%3) + if v != w { + t.Errorf("einsum outer [%d]: got %v, want %v", i, v, w) + } + } +} + +// TestEinsumTransposedInner pins the label alignment of full +// contractions: "ij,ji->" must pair a's columns with b's rows. It used +// to multiply positionally, computing "ij,ij->" and failing outright +// on non-square shapes. +func TestEinsumTransposedInner(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + b := mustFromFloats(t, []float64{5, 6, 7, 8}, 2, 2) + got, err := core.Einsum("ij,ji->", a, b) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0) + if v != 69 { // 1*5 + 2*7 + 3*6 + 4*8 + t.Errorf("einsum ij,ji->: got %v, want 69", v) + } + + // Non-square operands must work: sum of a * b^T. + w := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + x := mustFromFloats(t, []float64{7, 8, 9, 10, 11, 12}, 3, 2) + got2, err := core.Einsum("ij,ji->", w, x) + if err != nil { + t.Fatal(err) + } + v2, _ := core.FloatAt(got2, 0) + if v2 != 212 { // 58 + 154, the sum of w * x^T + t.Errorf("einsum ij,ji-> non-square: got %v, want 212", v2) + } +} + +// TestEinsumOuterDtypes pins the promotion ladder of the outer product: +// int vectors stay int (exact products), complex vectors stay complex. +func TestEinsumOuterDtypes(t *testing.T) { + ia := mustFromInts(t, []int64{1, 2}, 2) + ib := mustFromInts(t, []int64{3, 4}, 2) + got, err := core.Einsum("i,j->ij", ia, ib) + if err != nil { + t.Fatal(err) + } + if got.Dtype() != core.Int { + t.Fatalf("einsum outer int dtype: %s", got.Dtype()) + } + want := mustFromInts(t, []int64{3, 4, 6, 8}, 2, 2) + if !core.Equal(want, got) { + t.Errorf("einsum outer int: %s", got) + } + + ca := mustFromComplexes(t, []complex128{1 + 2i}, 1) + cb := mustFromComplexes(t, []complex128{3 + 4i}, 1) + gotC, err := core.Einsum("i,j->ij", ca, cb) + if err != nil { + t.Fatal(err) + } + if gotC.Dtype() != core.Complex { + t.Fatalf("einsum outer complex dtype: %s", gotC.Dtype()) + } + vc, err := core.ComplexAt(gotC, 0, 0) + if err != nil { + t.Fatal(err) + } + if vc != -5+10i { // (1+2i)(3+4i) + t.Errorf("einsum outer complex: got %v, want (-5+10i)", vc) + } +} + +func TestEinsumError(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2}, 2) + // Wrong operand count. + if _, err := core.Einsum("i,j->ij", a); err == nil { + t.Error("Einsum: expected error for wrong operand count") + } + // Unsupported pattern. + b := mustFromFloats(t, []float64{3, 4}, 2) + if _, err := core.Einsum("ii,jj->ij", a, b); err == nil { + t.Error("Einsum: expected error for unsupported pattern") + } + // Bad spec (no ->). + if _, err := core.Einsum("ii", a); err == nil { + t.Error("Einsum: expected error for missing ->") + } +} + +func minInt(a, b int) int { + if a < b { + return a + } + return b +} + +// TestEigenNonDiagonalReconstruction pins the corrected symmetric QR +// sweep: eigenpairs of non-diagonal matrices used to come back as +// garbage (near-zero eigenvalues for [[2,1],[1,2]]) because the QR step +// updated the tridiagonal with incoherent formulas. Eigenpairs are now +// checked by the residual ‖A·v − λ·v‖ and eigenvector orthogonality +// across a range of ranks. +func TestEigenNonDiagonalReconstruction(t *testing.T) { + // Golden: eigenvalues of [[2,1],[1,2]] are 1 and 3. + em := mustFromFloats(t, []float64{2, 1, 1, 2}, 2, 2) + vals, _, err := Eigen(em) + if err != nil { + t.Fatalf("Eigen golden: %v", err) + } + if v := vals.FloatAt(0); math.Abs(v-1) > 1e-12 { + t.Errorf("Eigen golden [0]: %v, want 1", v) + } + if v := vals.FloatAt(1); math.Abs(v-3) > 1e-12 { + t.Errorf("Eigen golden [1]: %v, want 3", v) + } + + for _, n := range []int{3, 5, 8, 20, 50} { + mat := make([]float64, n*n) + for i := range n { + for j := i; j < n; j++ { + v := math.Sin(float64(n)+float64(i*7+j*13)) * 2.0 + mat[i*n+j] = v + mat[j*n+i] = v + } + } + a := floatsToArray(mat, []int{n, n}) + vs, vc, err := Eigen(a) + if err != nil { + t.Fatalf("Eigen n=%d: %v", n, err) + } + maxRes, maxOrtho := 0.0, 0.0 + for k := range n { + for i := range n { + av := 0.0 + for j := range n { + av += mat[i*n+j] * vc.FloatAt(j*n+k) + } + if d := math.Abs(av - vs.FloatAt(k)*vc.FloatAt(i*n+k)); d > maxRes { + maxRes = d + } + } + for k2 := k + 1; k2 < n; k2++ { + dot := 0.0 + for i := range n { + dot += vc.FloatAt(i*n+k) * vc.FloatAt(i*n+k2) + } + if math.Abs(dot) > maxOrtho { + maxOrtho = math.Abs(dot) + } + } + } + if maxRes > 1e-11 || maxOrtho > 1e-12 { + t.Errorf("Eigen n=%d: residual %v, orthogonality %v", n, maxRes, maxOrtho) + } + } +} + +// TestSVDGeneralReconstruction pins the corrected Golub-Kahan +// pipeline: the old bidiagonalisation zeroed the wrong direction and +// dropped the upper triangle, the right singular vectors ignored V₁, +// and the wide return path transposed U. General matrices of every +// aspect ratio must now reconstruct with orthonormal factors. +func TestSVDGeneralReconstruction(t *testing.T) { + for _, tc := range []struct{ m, n int }{{4, 4}, {6, 3}, {3, 6}, {8, 5}, {5, 8}, {2, 3}} { + m, n := tc.m, tc.n + mat := make([]float64, m*n) + for i := range mat { + mat[i] = math.Sin(float64(i)*0.7)*1.5 + 0.3 + } + a := floatsToArray(mat, []int{m, n}) + u, s, vt, err := SVD(a) + if err != nil { + t.Fatalf("SVD %dx%d: %v", m, n, err) + } + if got := s.Shape()[0]; got != min(m, n) { + t.Fatalf("SVD %dx%d: %d singular values, want %d", m, n, got, min(m, n)) + } + ucols := u.Shape()[1] + if ucols != min(m, n) || vt.Shape()[0] != min(m, n) { + t.Fatalf("SVD %dx%d thin shapes: u=%v vt=%v", m, n, u.Shape(), vt.Shape()) + } + maxRec := 0.0 + for i := range m { + for j := range n { + sum := 0.0 + for k := range s.Len() { + sum += u.FloatAt(i*ucols+k) * s.FloatAt(k) * vt.FloatAt(k*n+j) + } + if d := math.Abs(sum - mat[i*n+j]); d > maxRec { + maxRec = d + } + } + } + orthU := 0.0 + for k1 := range ucols { + for k2 := range ucols { + dot := 0.0 + for i := range m { + dot += u.FloatAt(i*ucols+k1) * u.FloatAt(i*ucols+k2) + } + want := 0.0 + if k1 == k2 { + want = 1 + } + if d := math.Abs(dot - want); d > orthU { + orthU = d + } + } + } + if maxRec > 1e-9 || orthU > 1e-12 { + t.Errorf("SVD %dx%d: reconstruction %v, U orthogonality %v", m, n, maxRec, orthU) + } + // Σ must come back descending. + for k := 1; k < s.Len(); k++ { + if s.FloatAt(k) > s.FloatAt(k-1) { + t.Errorf("SVD %dx%d: σ not descending at %d", m, n, k) + } + } + } +} + +// TestFitPolynomialRefusesExtremeDegree pins the degree guard: degree +// MaxInt would wrap degree+1 negative and slip past the sample count +// check. +func TestFitPolynomialRefusesExtremeDegree(t *testing.T) { + x := mustFloats(t, []float64{0, 1, 2}, 3) + y := mustFloats(t, []float64{0, 1, 4}, 3) + if _, err := FitPolynomial(x, y, math.MaxInt); err == nil || !strings.Contains(err.Error(), "too large") { + t.Fatalf("FitPolynomial with degree MaxInt: %v", err) + } +} + +// TestCholeskyBlockedReconstruction pins the factorisation past the +// width of its column block, the sizes only the benchmarks otherwise +// reach: 320 is five whole blocks and 200 three blocks and a partial +// one, so the panel update, the diagonal block and the trapezoid divide +// all run, the last of them over a block that does not end on a +// boundary. A divide that reaches back over the finished columns, or a +// panel update that misses them, leaves L·Lᵀ away from A by far more +// than rounding, so the reconstruction is checked against the original +// matrix rather than against another run of the same sweep. +func TestCholeskyBlockedReconstruction(t *testing.T) { + for _, n := range []int{320, 200, 96} { + rng := rand.New(rand.NewPCG(3, 5)) + b := make([]float64, n*n) + for i := range b { + b[i] = rng.NormFloat64() + } + // A = B·Bᵀ + n·I is symmetric positive definite, with the + // diagonal held well above the rounding floor. + a := make([]float64, n*n) + for i := range n { + for j := range i + 1 { + s := 0.0 + for k := range n { + s += b[i*n+k] * b[j*n+k] + } + if i == j { + s += float64(n) + } + a[i*n+j], a[j*n+i] = s, s + } + } + l, err := Cholesky(mustFromFloats(t, a, n, n)) + if err != nil { + t.Fatalf("Cholesky(%d): %v", n, err) + } + lf := l.RawFloats() + worst, scale := 0.0, 0.0 + for i := range n { + if lf[i*n+i] <= 0 { + t.Fatalf("n=%d: the factor's diagonal at %d is %g, want a positive real root", n, i, lf[i*n+i]) + } + for j := range n { + if j > i && lf[i*n+j] != 0 { + t.Fatalf("n=%d: the factor holds %g above the diagonal at [%d,%d]", n, lf[i*n+j], i, j) + } + s := 0.0 + for k := range min(i, j) + 1 { + s += lf[i*n+k] * lf[j*n+k] + } + if d := math.Abs(s - a[i*n+j]); d > worst { + worst = d + } + if v := math.Abs(a[i*n+j]); v > scale { + scale = v + } + } + } + if worst/scale > 1e-12 { + t.Fatalf("n=%d: L·Lᵀ misses A by %.6g (scale %.6g), relative %.3g", n, worst, scale, worst/scale) + } + } +} diff --git a/linalg/det_factor_dispatch_test.go b/linalg/det_factor_dispatch_test.go new file mode 100644 index 0000000..722ff5f --- /dev/null +++ b/linalg/det_factor_dispatch_test.go @@ -0,0 +1,34 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" +) + +// TestDetFactorsThroughParallelDispatch drives base.Factor's rank-1 +// update dispatch through the public Det entry: a 200 by 200 upper +// triangular matrix keeps every pivot on the diagonal, and the first +// pivot's update spans 199 rows times 200 columns, past the dispatch +// quantum, so the crewed path runs rather than the inline one. No row +// changes under the elimination, since the multiplier of every row is +// zero, so the determinant is exactly 2 to the 200th, which float64 +// holds as an exact power of two. +func TestDetFactorsThroughParallelDispatch(t *testing.T) { + const n = 200 + vals := make([]float64, n*n) + for i := range n { + for j := i; j < n; j++ { + vals[i*n+j] = 2 + } + } + d, err := Det(mustFloats(t, vals, n, n)) + if err != nil { + t.Fatalf("Det: %v", err) + } + if want := math.Ldexp(1, 200); d != want { + t.Fatalf("Det = %v, want exactly %v", d, want) + } +} diff --git a/linalg/detfloat32_test.go b/linalg/detfloat32_test.go new file mode 100644 index 0000000..14c126a --- /dev/null +++ b/linalg/detfloat32_test.go @@ -0,0 +1,17 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "testing" +) + +// TestDetFloat32Input moved with Det from the root package: the linear +// algebra kernels compute in float64 even from float32 inputs. +func TestDetFloat32Input(t *testing.T) { + det, err := Det(mustFromFloat32s(t, []float32{1, 2, 3, 4}, 2, 2)) + if err != nil || det != -2 { + t.Fatalf("Det float32: %v %v", det, err) + } +} diff --git a/linalg/doc.go b/linalg/doc.go new file mode 100644 index 0000000..d6adb11 --- /dev/null +++ b/linalg/doc.go @@ -0,0 +1,84 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package linalg is the dense and sparse linear algebra of the library: +// the factorisations, the solvers, the eigensolvers, the matrix +// functions and the Krylov methods, over the array types the core +// provides. Every symbol here is re-exported by the root package, so +// user code may import that facade alone and write tensor.Solve. +// +// # Dense +// +// The dense surface takes a 2-D array and answers arrays. Solve, Inv +// and Det run one LU decomposition with partial pivoting; Cholesky, +// QR, LeastSquares, RRQR and the tridiagonal pair SolveTridiagonal and +// SolveCyclicTridiagonal are the other factorisations and their solves. +// Eigen, EigenComplex, EigenGeneral, EigenGeneralised, SVD, SVDComplex +// and SchurComplex are the spectral decompositions, and Pinverse, +// MatrixRank and Cond read the singular spectrum that SVD produces. +// SolveTikhonov, SolveTruncated and SolveRRQR answer a system whose +// data does not determine the solution. MatrixExp, MatrixSqrt and +// MatrixLog are the matrix functions; FitPolynomial, PolynomialRoots +// and CubicSpline carry the polynomial work; GMRES solves a system +// presented as an operator rather than as a matrix. +// +// # Sparse +// +// The sparse surface starts from a COO matrix (the root package's +// tensor.SparseCOO) and converts it once into the compressed view the +// algorithm wants: SparseCSR for the iterative solvers, the Lanczos +// eigensolver and the row-wise products, SparseCSC for the direct +// factorisations and the column-wise scatter. The direct route is +// NewSparseCholesky for a symmetric positive definite matrix, under a +// fill-reducing SparseOrdering, and NewSparseLU for a general square +// matrix; NewSparseILU builds the incomplete factorisation the Krylov +// solvers take as a preconditioner. SpSolve and SpSolveBiCGSTAB, with +// their two complex counterparts, are the iterative solves; SpLSQR and +// SpLSMR are the least-squares iterations; SpEigen, SpEigenGeneral and +// the two complex forms are the Krylov eigensolvers; SpExpApply +// applies a matrix exponential to a vector without forming the matrix. +// +// # The pipeline +// +// Pipe and Pipeline read a sequence of array transformations as one +// expression. The evaluation is eager and each step allocates a new +// array. There is no lazy graph, no autograd and no backpropagation +// here: a failed step is recorded, the steps after it become no-ops, +// and Result returns the first error alongside the array. +// +// # Conventions +// +// The dense routines compute in float64 and complex128: int and float32 +// inputs promote on the way in, and one complex operand promotes the +// whole call. A sparse routine with no complex form refuses a complex +// input instead of promoting it, and the four complex entry points name +// their element type in the name: SpSolveComplexCG, +// SpSolveComplexBiCGSTAB, SpEigenComplex and SpEigenGeneralComplex. +// +// Symmetry is checked where the algorithm depends on it, within a +// scale-relative 1e-12 tolerance by Eigen, EigenComplex, SpEigen, +// SpEigenComplex, SpSolve, SpSolveComplexCG and SpExpApply. It is not +// checked where the algorithm does not need it: EigenGeneralised +// requires a symmetric a and does not verify it, and the two general +// sparse eigensolvers do not screen for non-finite entries, which come +// back as NaN Ritz pairs. +// +// A shape mismatch, a singular matrix, a rank-deficient system, a +// matrix that leaves the positive definite cone during a rank-one +// update, and a Krylov iteration that exhausts its budget with the +// tolerance unmet are all errors naming themselves. An unconverged +// iterative solve returns no estimate, never a silent approximation. +// +// Results are reproducible: the kernels that dispatch over workers +// split the output space rather than the reduction, so an element is +// computed by the same arithmetic sequence whatever the worker count. +// +// # Constructors +// +// The array type belongs to the core, and its general constructors are +// re-exported by the root package as tensor.FromFloats, tensor.Zeros, +// tensor.Identity and their neighbours. A caller that imports this +// package on its own, without the root facade, has ArrayFromFloatsSafe +// for building the input to a call: it copies the slice it is given, +// so the caller keeps ownership of the values it passed in. +package linalg diff --git a/linalg/dtype_shape_refusal_pins_test.go b/linalg/dtype_shape_refusal_pins_test.go new file mode 100644 index 0000000..f298750 --- /dev/null +++ b/linalg/dtype_shape_refusal_pins_test.go @@ -0,0 +1,435 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression tests: public entries that +// accepted a dtype or a shape the implementation then tripped over. +// Every case is reachable through the facade with well-typed input. +// The dtype and shape gaps panicked, and for the COO and CSR entries +// the read happened in an engine worker goroutine, which aborts the +// process instead of returning to the caller; the CSRFromCOO ordering +// and the matrix-function routing produced a wrong answer silently. + +// cOO builds a real SparseCOO from a flat (index tuple, value) +// list. +func cOO(t *testing.T, shape []int, idx []int64, vals []float64) *core.SparseCOO { + t.Helper() + coo, err := core.NewSparseCOO(mustFromInts(t, idx, len(vals), len(shape)), floatsToArray(vals, []int{len(vals)}), shape) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// pinComplexCOO is cOO with complex128 values. +func pinComplexCOO(t *testing.T, shape []int, idx []int64, vals []complex128) *core.SparseCOO { + t.Helper() + coo, err := core.NewSparseCOO(mustFromInts(t, idx, len(vals), len(shape)), mustFromComplexes(t, vals, len(vals)), shape) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// overdetermined is the real 4x2 matrix the least-squares entry +// points below solve, and the exact x it yields for b = [1 2 3 4]: +// AᵀA = [[3,2],[2,3]], Aᵀb = [8,9], so x = [6/5, 11/5]. +func overdetermined(t *testing.T) *core.Array { + t.Helper() + return mustFloats(t, []float64{1, 0, 0, 1, 1, 1, 1, 1}, 4, 2) +} + +// leastSquaresSol is the exact solution above. +var leastSquaresSol = []float64{6.0 / 5, 11.0 / 5} + +// TestLeastSquaresComplexRightHandSideRefused pins the b-side dtype +// gate of LeastSquares: a real a with a complex b used to fall into +// b.FloatAt, whose complex payload is empty, and panicked with an +// index-out-of-range instead of returning the dtype refusal. +func TestLeastSquaresComplexRightHandSideRefused(t *testing.T) { + a := overdetermined(t) + for _, tc := range []struct { + name string + b *core.Array + }{ + {"vector", mustFromComplexes(t, []complex128{1 + 1i, 2, 3, 4}, 4)}, + {"matrix", mustFromComplexes(t, []complex128{1 + 1i, 2, 3, 4, 5, 6, 7, 8}, 4, 2)}, + } { + t.Run(tc.name, func(t *testing.T) { + if _, err := LeastSquares(a, tc.b); err == nil { + t.Fatal("want an error for a complex right-hand side") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-dtype refusal", err) + } + }) + } + // A real right-hand side still takes the QR route: the gate refuses + // the dtype, not the call. + b := mustFloats(t, []float64{1, 2, 3, 4}, 4) + x, err := LeastSquares(a, b) + if err != nil { + t.Fatalf("LeastSquares: %v", err) + } + for i, want := range leastSquaresSol { + if math.Abs(x.FloatAt(i)-want) > 1e-12 { + t.Fatalf("x[%d] = %v, want %v", i, x.FloatAt(i), want) + } + } +} + +// TestSvdSolveComplexRightHandSideRefused pins the same gap at the two +// SVD-based solvers, which both fill their b copy through FloatAt. +func TestSvdSolveComplexRightHandSideRefused(t *testing.T) { + a := overdetermined(t) + solvers := []struct { + name string + solve func(*core.Array, *core.Array) (*core.Array, error) + }{ + {"SolveTruncated", func(a, b *core.Array) (*core.Array, error) { return SolveTruncated(a, b, 2) }}, + {"SolveTikhonov", func(a, b *core.Array) (*core.Array, error) { return SolveTikhonov(a, b, 1e-9) }}, + } + for _, s := range solvers { + t.Run(s.name+"/vector", func(t *testing.T) { + b := mustFromComplexes(t, []complex128{1 + 1i, 2, 3, 4}, 4) + if _, err := s.solve(a, b); err == nil { + t.Fatal("want an error for a complex right-hand side") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-dtype refusal", err) + } + }) + t.Run(s.name+"/matrix", func(t *testing.T) { + b := mustFromComplexes(t, []complex128{1 + 1i, 2, 3, 4, 5, 6, 7, 8}, 4, 2) + if _, err := s.solve(a, b); err == nil { + t.Fatal("want an error for a complex right-hand side") + } + }) + t.Run(s.name+"/real", func(t *testing.T) { + b := mustFloats(t, []float64{1, 2, 3, 4}, 4) + x, err := s.solve(a, b) + if err != nil { + t.Fatalf("%s: %v", s.name, err) + } + // Full rank and light damping: the answer matches the exact + // least-squares solution of the system above. + for i, want := range leastSquaresSol { + if math.Abs(x.FloatAt(i)-want) > 1e-6 { + t.Fatalf("x[%d] = %v, want %v", i, x.FloatAt(i), want) + } + } + }) + } +} + +// TestCubicSplineEvaluateComplexQueryRefused pins the Evaluate side of +// the dtype contract NewCubicSpline already enforced: a complex query +// array used to reach x.FloatAt and panic. +func TestCubicSplineEvaluateComplexQueryRefused(t *testing.T) { + s, err := NewCubicSpline(mustFloats(t, []float64{0, 1, 2, 3}), mustFloats(t, []float64{0, 1, 4, 9})) + if err != nil { + t.Fatalf("NewCubicSpline: %v", err) + } + for _, tc := range []struct { + name string + x *core.Array + }{ + {"vector", mustFromComplexes(t, []complex128{1 + 1i, 2}, 2)}, + {"matrix", mustFromComplexes(t, []complex128{1 + 1i, 2, 3, 4}, 2, 2)}, + } { + t.Run(tc.name, func(t *testing.T) { + if _, err := s.Evaluate(tc.x); err == nil { + t.Fatal("want an error for a complex query array") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-dtype refusal", err) + } + }) + } + // A real query still evaluates, and at a knot it returns the knot + // value. + out, err := s.Evaluate(mustFloats(t, []float64{1, 2, 3}, 3)) + if err != nil { + t.Fatalf("Evaluate: %v", err) + } + for i, want := range []float64{1, 4, 9} { + if math.Abs(out.FloatAt(i)-want) > 1e-12 { + t.Fatalf("Evaluate[%d] = %v, want %v", i, out.FloatAt(i), want) + } + } +} + +// TestSpEigenGeneralRejectsNonSquareShapes pins the square 2-D +// contract at the general sparse eigensolver. Without the entry guard a +// 3x5 COO built its CSR fine (each index is valid for its own +// dimension) and the Krylov product then read x[4] with len(x) == 3 +// inside an engine worker: a process abort no caller can recover. A +// rank above two was accepted too, with every dimension past the first +// two silently ignored. +func TestSpEigenGeneralRejectsNonSquareShapes(t *testing.T) { + cases := []struct { + name string + mk func(t *testing.T) *core.SparseCOO + }{ + {"real/3x5", func(t *testing.T) *core.SparseCOO { + return cOO(t, []int{3, 5}, []int64{0, 4, 1, 3}, []float64{1, 2}) + }}, + {"real/rank3", func(t *testing.T) *core.SparseCOO { + return cOO(t, []int{2, 2, 2}, []int64{0, 0, 0, 1, 1, 1}, []float64{1, 2}) + }}, + {"complex/3x5", func(t *testing.T) *core.SparseCOO { + return pinComplexCOO(t, []int{3, 5}, []int64{0, 4, 1, 3}, []complex128{1, 2}) + }}, + {"complex/rank3", func(t *testing.T) *core.SparseCOO { + return pinComplexCOO(t, []int{2, 2, 2}, []int64{0, 0, 0, 1, 1, 1}, []complex128{1, 2}) + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + s := tc.mk(t) + general := strings.HasPrefix(tc.name, "real") + call := func() error { + if general { + _, _, err := SpEigenGeneral(s, 1, nil) + return err + } + _, _, err := SpEigenGeneralComplex(s, 1, nil) + return err + } + err := call() + if err == nil { + t.Fatal("want an error for a non-square or n-D matrix") + } + // Every shape is refused before any engine dispatch, so the + // call returns instead of aborting inside a worker. + if !strings.Contains(err.Error(), "shape") && !strings.Contains(err.Error(), "square") && + !strings.Contains(err.Error(), "2-D") { + t.Fatalf("error = %v, want a shape refusal", err) + } + }) + } + // A square 2-D real COO still decomposes: [[2,1],[1,2]] has + // eigenvalues 3 and 1, so k = 1 returns 3 with a unit eigenvector. + sq := cOO(t, []int{2, 2}, []int64{0, 0, 0, 1, 1, 0, 1, 1}, []float64{2, 1, 1, 2}) + vals, vecs, err := SpEigenGeneral(sq, 1, nil) + if err != nil { + t.Fatalf("SpEigenGeneral: %v", err) + } + if math.Abs(real(vals.ComplexAt(0))-3) > 1e-10 { + t.Fatalf("values[0] = %v, want 3", vals.ComplexAt(0)) + } + norm := math.Hypot(real(vecs.ComplexAt(0)), imag(vecs.ComplexAt(0))) + norm = math.Hypot(norm, math.Hypot(real(vecs.ComplexAt(1)), imag(vecs.ComplexAt(1)))) + if math.Abs(norm-1) > 1e-12 { + t.Fatalf("eigenvector norm = %v, want 1", norm) + } +} + +// TestCSRFromCOOMergesDuplicatesBeforeDroppingZeros pins the ordering +// inside the conversion: the duplicate sum must run on the merged +// coordinate list, never against the last slot of the output. The old +// order summed into Values[len(Values)-1] after a zero had been +// dropped, which either indexed -1 or added the duplicate onto a +// different coordinate's entry. +func TestCSRFromCOOMergesDuplicatesBeforeDroppingZeros(t *testing.T) { + t.Run("duplicate of a dropped zero", func(t *testing.T) { + // [(0,0)=0, (0,0)=3]: the zero is dropped but its duplicate + // still sums to 3, so the coordinate must carry 3. + csr, err := CSRFromCOO(cOO(t, []int{2, 2}, []int64{0, 0, 0, 0}, []float64{0, 3})) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + if csr.NNZ() != 1 || csr.Values[0] != 3 || csr.ColIdx[0] != 0 { + t.Fatalf("CSR = values %v colIdx %v, want one entry 3 at column 0", csr.Values, csr.ColIdx) + } + wantRowStart := []int{0, 1, 1} + for i, want := range wantRowStart { + if csr.RowStart[i] != want { + t.Fatalf("RowStart = %v, want %v", csr.RowStart, wantRowStart) + } + } + }) + t.Run("duplicate in a later row", func(t *testing.T) { + // [(0,1)=5, (1,0)=0, (1,0)=3] used to add the 3 onto the (0,1) + // slot: [[0,5],[3,0]] read as [[0,8],[0,0]]. + csr, err := CSRFromCOO(cOO(t, []int{2, 2}, []int64{0, 1, 1, 0, 1, 0}, []float64{5, 0, 3})) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + wantVals := []float64{5, 3} + wantCols := []int{1, 0} + wantRowStart := []int{0, 1, 2} + for i, want := range wantVals { + if len(csr.Values) != 2 || csr.Values[i] != want { + t.Fatalf("Values = %v, want %v", csr.Values, wantVals) + } + if csr.ColIdx[i] != wantCols[i] { + t.Fatalf("ColIdx = %v, want %v", csr.ColIdx, wantCols) + } + } + for i, want := range wantRowStart { + if csr.RowStart[i] != want { + t.Fatalf("RowStart = %v, want %v", csr.RowStart, wantRowStart) + } + } + }) + t.Run("duplicates merging to zero", func(t *testing.T) { + // [(0,0)=1, (0,0)=-1, (1,1)=2]: the merged coordinate cancels + // and only the canonical survivor is stored. + csr, err := CSRFromCOO(cOO(t, []int{2, 2}, []int64{0, 0, 0, 0, 1, 1}, []float64{1, -1, 2})) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + if csr.NNZ() != 1 || csr.Values[0] != 2 || csr.ColIdx[0] != 1 { + t.Fatalf("CSR = values %v colIdx %v, want one entry 2 at column 1", csr.Values, csr.ColIdx) + } + }) + t.Run("duplicates summing to a value", func(t *testing.T) { + // [(0,0)=1, (0,0)=2, (1,1)=4] leaves a sorted, unique pattern. + csr, err := CSRFromCOO(cOO(t, []int{2, 2}, []int64{0, 0, 0, 0, 1, 1}, []float64{1, 2, 4})) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + if csr.NNZ() != 2 || csr.Values[0] != 3 || csr.Values[1] != 4 { + t.Fatalf("CSR = values %v colIdx %v, want [3 4] at columns [0 1]", csr.Values, csr.ColIdx) + } + }) +} + +// TestSparseCSRComplexOperandRefused pins the dtype gate of the CSR +// kernels: MatVec panicked inside an engine worker (a process abort), +// MatMulDense on the calling goroutine, both because a complex operand +// has no float payload. +func TestSparseCSRComplexOperandRefused(t *testing.T) { + csr := sparseCSRFromEntries(t, 2, 2, [][3]float64{{0, 0, 1}, {1, 1, 1}}) + t.Run("MatVec", func(t *testing.T) { + if _, err := csr.MatVec(mustFromComplexes(t, []complex128{1, 2}, 2)); err == nil { + t.Fatal("want an error for a complex vector") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-dtype refusal", err) + } + }) + t.Run("MatMulDense", func(t *testing.T) { + if _, err := csr.MatMulDense(mustFromComplexes(t, []complex128{1, 2, 3, 4}, 2, 2)); err == nil { + t.Fatal("want an error for a complex operand") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-dtype refusal", err) + } + }) + t.Run("real operands", func(t *testing.T) { + y, err := csr.MatVec(mustFloats(t, []float64{3, 4}, 2)) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + for i, want := range []float64{3, 4} { + if y.FloatAt(i) != want { + t.Fatalf("y[%d] = %v, want %v", i, y.FloatAt(i), want) + } + } + }) +} + +// TestSparseILURefusesMissingOrZeroDiagonal pins the documented +// contract for the whole matrix: zero pivots were only detected for +// the diagonals a later row pivots on, so a row nothing below it +// references (the last row in particular) was accepted and Apply then +// divided by the defensive 1. +func TestSparseILURefusesMissingOrZeroDiagonal(t *testing.T) { + cases := []struct { + name string + coo func(t *testing.T) *core.SparseCOO + }{ + {"last row without a diagonal", func(t *testing.T) *core.SparseCOO { + // [[1,1],[0,0]]: no later row pivots on row 1. + return cOO(t, []int{2, 2}, []int64{0, 0, 0, 1}, []float64{1, 1}) + }}, + {"diagonal matrix missing its second diagonal", func(t *testing.T) *core.SparseCOO { + return cOO(t, []int{2, 2}, []int64{0, 0}, []float64{2}) + }}, + {"explicit zero diagonal", func(t *testing.T) *core.SparseCOO { + return cOO(t, []int{2, 2}, []int64{0, 0, 1, 1}, []float64{0, 0}) + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if _, err := NewSparseILU(tc.coo(t)); err == nil { + t.Fatal("want an error for a missing or zero diagonal") + } else if !strings.Contains(err.Error(), "diagonal") { + t.Fatalf("error = %v, want a diagonal refusal", err) + } + }) + } + // A complete non-zero diagonal still factors, and on a tridiagonal + // matrix ILU(0) is the exact LU: Apply reproduces the dense solve. + a := mustFloats(t, []float64{3, -1, 0, -1, 3, -1, 0, -1, 3}, 3, 3) + coo := cOO(t, []int{3, 3}, + []int64{0, 0, 0, 1, 1, 0, 1, 1, 1, 2, 2, 1, 2, 2}, + []float64{3, -1, -1, 3, -1, -1, 3}) + ilu, err := NewSparseILU(coo) + if err != nil { + t.Fatalf("NewSparseILU: %v", err) + } + r := mustFloats(t, []float64{1, 0, 0}, 3) + x, err := Solve(a, r) + if err != nil { + t.Fatalf("Solve: %v", err) + } + got := ilu.Apply([]float64{1, 0, 0}) + for i := range got { + if math.Abs(got[i]-x.FloatAt(i)) > 1e-12 { + t.Fatalf("Apply[%d] = %v, want the dense solve %v", i, got[i], x.FloatAt(i)) + } + } +} + +// TestSmallScaleRoutingAgreesWithEigenGuard pins the routing decision +// at tiny scale. The matrix below carries an absolute asymmetry of +// 5e-13, which is 5e-3 relative to its 1e-10 scale: the old +// `1e-12 * max(1, scale)` floor called it symmetric, the Eigen route +// applied its own relative check and refused, and the Schur route +// never ran. The routing tolerance is now purely relative, so any +// matrix the route approves stays within the fraction of the scale the +// route's own guard tolerates. +func TestSmallScaleRoutingAgreesWithEigenGuard(t *testing.T) { + const s = 1e-10 + asymmetric := mustFloats(t, []float64{4 * s, s, s + 5e-13, 3 * s}, 2, 2) + if isSymmetricMatrix(asymmetric) { + t.Fatal("a matrix with 5e-3 relative asymmetry is reported symmetric") + } + r, err := MatrixSqrt(asymmetric) + if err != nil { + t.Fatalf("MatrixSqrt refused a nonsymmetric small-scale matrix instead of routing it to the Schur path: %v", err) + } + // The answer must be a square root: r·r = A, relatively to the + // matrix scale. + prod, err := core.MatMul2D(r, r) + if err != nil { + t.Fatalf("MatMul2D: %v", err) + } + for i := range 4 { + if math.Abs(prod.FloatAt(i)-asymmetric.FloatAt(i)) > 1e-6*4*s { + t.Fatalf("(√A)²[%d] = %v, want %v", i, prod.FloatAt(i), asymmetric.FloatAt(i)) + } + } + // A genuinely symmetric matrix of the same scale keeps the cheaper + // symmetric route and still answers. + symmetric := mustFloats(t, []float64{4 * s, s, s, 3 * s}, 2, 2) + if !isSymmetricMatrix(symmetric) { + t.Fatal("an exactly symmetric matrix is reported nonsymmetric") + } + if _, err := MatrixSqrt(symmetric); err != nil { + t.Fatalf("MatrixSqrt on the symmetric matrix: %v", err) + } + // Rounding-level asymmetry, by contrast, still counts as symmetric: + // the relative tolerance is 1e-12 of the scale. + near := mustFloats(t, []float64{4 * s, s, s + 1e-22, 3 * s}, 2, 2) + if !isSymmetricMatrix(near) { + t.Fatal("asymmetry at 1e-12 of the scale is reported nonsymmetric") + } +} diff --git a/linalg/dtypes_census_test.go b/linalg/dtypes_census_test.go new file mode 100644 index 0000000..02e10db --- /dev/null +++ b/linalg/dtypes_census_test.go @@ -0,0 +1,1000 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The dtype census for linalg: every public entry that takes an array +// is probed with Bool, the narrow integers and the Int anchor, against +// a float64 baseline carrying exactly the widened probe values. An +// entry either refuses a dtype by name (pinned wording), computes +// bit-identically through the accessor walks (the Int treatment), or +// rides a core narrow-dtype refusal through. Nothing may panic or silently +// read a nil payload. + +// cenDtypes is the probe set: the narrow dtypes plus the Int anchor. +var cenDtypes = []core.Dtype{core.Bool, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32, core.Int} + +// cenMaker builds arrays for one probe dtype; the baseline maker +// carries the same cast values in float64. +type cenMaker func(vals []float64, shape ...int) *core.Array + +// cenCast narrows v into dt's value space the way the payload store +// does, so the baseline sees exactly what the probe dtype carries. +func cenCast(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 + } +} + +// cenMakers returns the probe maker for dt and the float64 baseline +// maker over the same cast values. +func cenMakers(t *testing.T, dt core.Dtype) (probe, base cenMaker) { + t.Helper() + build := func(vals []float64, shape ...int) *core.Array { + cast := make([]float64, len(vals)) + for i, v := range vals { + cast[i] = cenCast(dt, v) + } + 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...) + if err != nil { + t.Fatalf("FromBools: %v", err) + } + return a + case core.Int8: + vs := make([]int8, len(cast)) + for i, v := range cast { + vs[i] = int8(int64(v)) + } + a, err := core.FromInt8s(vs, shape...) + if err != nil { + t.Fatalf("FromInt8s: %v", err) + } + return a + case core.Uint8: + vs := make([]uint8, len(cast)) + for i, v := range cast { + vs[i] = uint8(int64(v)) + } + a, err := core.FromUint8s(vs, shape...) + if err != nil { + t.Fatalf("FromUint8s: %v", err) + } + return a + case core.Int16: + vs := make([]int16, len(cast)) + for i, v := range cast { + vs[i] = int16(int64(v)) + } + a, err := core.FromInt16s(vs, shape...) + if err != nil { + t.Fatalf("FromInt16s: %v", err) + } + return a + case core.Uint16: + vs := make([]uint16, len(cast)) + for i, v := range cast { + vs[i] = uint16(int64(v)) + } + a, err := core.FromUint16s(vs, shape...) + if err != nil { + t.Fatalf("FromUint16s: %v", err) + } + return a + case core.Int32: + vs := make([]int32, len(cast)) + for i, v := range cast { + vs[i] = int32(int64(v)) + } + a, err := core.FromInt32s(vs, shape...) + if err != nil { + t.Fatalf("FromInt32s: %v", err) + } + return a + case core.Uint32: + vs := make([]uint32, len(cast)) + for i, v := range cast { + vs[i] = uint32(int64(v)) + } + a, err := core.FromUint32s(vs, shape...) + if err != nil { + t.Fatalf("FromUint32s: %v", err) + } + return a + case core.Int: + vs := make([]int64, len(cast)) + for i, v := range cast { + vs[i] = int64(v) + } + a, err := core.FromInts(vs, shape...) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + return a + default: + fs := make([]float64, len(cast)) + copy(fs, cast) + a, err := core.FromFloats(fs, shape...) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a + } + } + baseBuild := func(vals []float64, shape ...int) *core.Array { + cast := make([]float64, len(vals)) + for i, v := range vals { + cast[i] = cenCast(dt, v) + } + a, err := core.FromFloats(cast, shape...) + if err != nil { + t.Fatalf("FromFloats baseline: %v", err) + } + return a + } + return build, baseBuild +} + +// cenWiden reads an array through the widening accessor: complex +// outputs (the general eigensolver spectra, polynomial roots) read +// through ComplexAt, everything else through FloatAt. +func cenWiden(t *testing.T, a *core.Array) []complex128 { + t.Helper() + out := make([]complex128, a.Len()) + for i := range out { + if a.Dtype() == core.Complex { + out[i] = a.ComplexAt(i) + continue + } + out[i] = complex(a.FloatAt(i), 0) + } + return out +} + +// cenArrays pins probe outputs against baseline outputs: either both +// runs take the same value-driven error, or both succeed with equal +// dtype and bit-identical widened values. +func cenArrays(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.Shape()[0] != b.Shape()[0] || len(p.Shape()) != len(b.Shape()) || p.Len() != b.Len() { + t.Fatalf("%s(%s): output %d shape %v, want %v", label, dt, k, p.Shape(), b.Shape()) + } + pv, bv := cenWiden(t, p), cenWiden(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]) + } + } + } +} + +// cenScalar pins a float64 or int result against the baseline's. +func cenFloat(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 pv != bv { + t.Fatalf("%s(%s) = %v, want the baseline %v", label, dt, pv, bv) + } +} + +// cenInt pins an int result against the baseline's. +func cenInt(t *testing.T, label string, dt core.Dtype, pv int, perr error, bv int, 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 pv != bv { + t.Fatalf("%s(%s) = %d, want the baseline %d", label, dt, pv, bv) + } +} + +// cenWantErr pins a refusal whose text must carry the given fragments. +func cenWantErr(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) + } + } +} + +// cenSparse builds a diagonal sparse matrix: values take the probe +// dtype where the core constructors allow it (the census pins their +// refusal where they do not) and the float baseline otherwise. +func cenSparse(t *testing.T, maker cenMaker, vals []float64, n int) *core.SparseCOO { + t.Helper() + indVals := make([]float64, 0, 2*n) + for i := range n { + indVals = append(indVals, float64(i), float64(i)) + } + ind, err := core.FromInts(int64s(indVals), n, 2) + if err != nil { + t.Fatalf("FromInts indices: %v", err) + } + s, err := core.NewSparseCOO(ind, maker(vals, n), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return s +} + +func int64s(vals []float64) []int64 { + out := make([]int64, len(vals)) + for i, v := range vals { + out[i] = int64(v) + } + return out +} + +// cenSparseErr is cenSparse without the fatal: a core constructor +// refusal on narrow values comes back as the error the row pins. +func cenSparseErr(maker cenMaker, vals []float64, n int) (*core.SparseCOO, error) { + indVals := make([]float64, 0, 2*n) + for i := range n { + indVals = append(indVals, float64(i), float64(i)) + } + ind, err := core.FromInts(int64s(indVals), n, 2) + if err != nil { + return nil, err + } + return core.NewSparseCOO(ind, maker(vals, n), []int{n, n}) +} + +// cenSPDFull builds the 2x2 sparse SPD matrix [[4,1],[1,2]] in float64: +// its Cholesky factor holds the [1,0] entry the rank-1 modification +// surface needs, whatever dtype the update vector carries. +func cenSPDFull() *core.SparseCOO { + ind, err := core.FromInts([]int64{0, 0, 0, 1, 1, 0, 1, 1}, 4, 2) + if err != nil { + panic(err) + } + vals, err := core.FromFloats([]float64{4, 1, 1, 2}, 4) + if err != nil { + panic(err) + } + s, err := core.NewSparseCOO(ind, vals, []int{2, 2}) + if err != nil { + panic(err) + } + return s +} + +// TestDtypesCensusLinalg probes every array-taking public entry. +func TestDtypesCensusLinalg(t *testing.T) { + rows := []struct { + name string + run func(t *testing.T, probe, base cenMaker, dt core.Dtype) + }{ + {"Det", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := Det(probe([]float64{4, 1, 1, 3}, 2, 2)) + b, berr := Det(base([]float64{4, 1, 1, 3}, 2, 2)) + cenFloat(t, "Det", dt, p, perr, b, berr) + }}, + {"DetComplex refuses every real dtype", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + _, perr := DetComplex(probe([]float64{1, 0, 0, 1}, 2, 2)) + cenWantErr(t, "DetComplex/"+dt.String(), perr, "DetComplex", "complex") + _, berr := DetComplex(base([]float64{1, 0, 0, 1}, 2, 2)) + cenWantErr(t, "DetComplex float baseline", berr, "DetComplex", "complex") + }}, + {"Solve", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := Solve(probe([]float64{4, 1, 1, 3}, 2, 2), probe([]float64{1, 2}, 2)) + b, berr := Solve(base([]float64{4, 1, 1, 3}, 2, 2), base([]float64{1, 2}, 2)) + cenArrays(t, "Solve", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Inv", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := Inv(probe([]float64{4, 1, 1, 3}, 2, 2)) + b, berr := Inv(base([]float64{4, 1, 1, 3}, 2, 2)) + cenArrays(t, "Inv", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"QR", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + q, r, perr := QR(probe([]float64{2, 1, 1, 3, 0, 1}, 3, 2)) + bq, br, berr := QR(base([]float64{2, 1, 1, 3, 0, 1}, 3, 2)) + cenArrays(t, "QR", dt, []*core.Array{q, r}, perr, []*core.Array{bq, br}, berr) + }}, + {"Cholesky", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := Cholesky(probe([]float64{4, 1, 1, 3}, 2, 2)) + b, berr := Cholesky(base([]float64{4, 1, 1, 3}, 2, 2)) + cenArrays(t, "Cholesky", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SVD", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + u, s, vt, perr := SVD(probe([]float64{2, 0, 1, 2, 0, 1}, 3, 2)) + bu, bs, bvt, berr := SVD(base([]float64{2, 0, 1, 2, 0, 1}, 3, 2)) + cenArrays(t, "SVD", dt, []*core.Array{u, s, vt}, perr, []*core.Array{bu, bs, bvt}, berr) + }}, + {"Eigen", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + v, e, perr := Eigen(probe([]float64{4, 1, 1, 3}, 2, 2)) + bv, be, berr := Eigen(base([]float64{4, 1, 1, 3}, 2, 2)) + cenArrays(t, "Eigen", dt, []*core.Array{v, e}, perr, []*core.Array{bv, be}, berr) + }}, + {"EigenGeneral", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + v, e, perr := EigenGeneral(probe([]float64{3, 1, 0, 2}, 2, 2)) + bv, be, berr := EigenGeneral(base([]float64{3, 1, 0, 2}, 2, 2)) + cenArrays(t, "EigenGeneral", dt, []*core.Array{v, e}, perr, []*core.Array{bv, be}, berr) + }}, + {"EigenComplex refuses real dtypes", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + _, _, perr := EigenComplex(probe([]float64{1, 0, 0, 1}, 2, 2)) + cenWantErr(t, "EigenComplex/"+dt.String(), perr, "EigenComplex", "complex") + }}, + {"SVDComplex refuses real dtypes", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + _, _, _, perr := SVDComplex(probe([]float64{1, 0, 0, 1}, 2, 2)) + cenWantErr(t, "SVDComplex/"+dt.String(), perr, "SVDComplex", "complex") + }}, + {"SchurComplex widens real dtypes", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + // SchurComplex accepts real input by design: MatrixSqrt + // and MatrixLog route real spectra through it, reading the + // elements via ComplexAt. Narrow dtypes follow Int. + pt, pq, perr := SchurComplex(probe([]float64{3, 1, 0, 2}, 2, 2)) + bt, bq, berr := SchurComplex(base([]float64{3, 1, 0, 2}, 2, 2)) + cenArrays(t, "SchurComplex", dt, []*core.Array{pt, pq}, perr, []*core.Array{bt, bq}, berr) + }}, + {"MatrixExp", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := MatrixExp(probe([]float64{1, 0, 0, 2}, 2, 2)) + b, berr := MatrixExp(base([]float64{1, 0, 0, 2}, 2, 2)) + cenArrays(t, "MatrixExp", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MatrixSqrt", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := MatrixSqrt(probe([]float64{4, 0, 0, 9}, 2, 2)) + b, berr := MatrixSqrt(base([]float64{4, 0, 0, 9}, 2, 2)) + cenArrays(t, "MatrixSqrt", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MatrixLog", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := MatrixLog(probe([]float64{2, 0, 0, 3}, 2, 2)) + b, berr := MatrixLog(base([]float64{2, 0, 0, 3}, 2, 2)) + cenArrays(t, "MatrixLog", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MatrixRank", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := MatrixRank(probe([]float64{2, 0, 0, 3, 0, 0}, 3, 2), 0) + b, berr := MatrixRank(base([]float64{2, 0, 0, 3, 0, 0}, 3, 2), 0) + cenInt(t, "MatrixRank", dt, p, perr, b, berr) + }}, + {"Cond", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := Cond(probe([]float64{4, 1, 1, 3}, 2, 2), 0) + b, berr := Cond(base([]float64{4, 1, 1, 3}, 2, 2), 0) + cenFloat(t, "Cond", dt, p, perr, b, berr) + }}, + {"LeastSquares", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := LeastSquares(probe([]float64{2, 0, 1, 2, 0, 1}, 3, 2), probe([]float64{1, 2, 2}, 3)) + b, berr := LeastSquares(base([]float64{2, 0, 1, 2, 0, 1}, 3, 2), base([]float64{1, 2, 2}, 3)) + cenArrays(t, "LeastSquares", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"FitPolynomial", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := FitPolynomial(probe([]float64{0, 1, 2, 3}, 4), probe([]float64{0, 1, 4, 9}, 4), 2) + b, berr := FitPolynomial(base([]float64{0, 1, 2, 3}, 4), base([]float64{0, 1, 4, 9}, 4), 2) + cenArrays(t, "FitPolynomial", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Pinverse", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := Pinverse(probe([]float64{4, 1, 1, 3}, 2, 2), 1e-10) + b, berr := Pinverse(base([]float64{4, 1, 1, 3}, 2, 2), 1e-10) + cenArrays(t, "Pinverse", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"PolynomialRoots", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := PolynomialRoots(probe([]float64{2, -3, 1}, 3)) + b, berr := PolynomialRoots(base([]float64{2, -3, 1}, 3)) + cenArrays(t, "PolynomialRoots", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"RRQR", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + q, r, _, _, perr := RRQR(probe([]float64{2, 1, 1, 3, 0, 1}, 3, 2)) + bq, br, _, _, berr := RRQR(base([]float64{2, 1, 1, 3, 0, 1}, 3, 2)) + cenArrays(t, "RRQR", dt, []*core.Array{q, r}, perr, []*core.Array{bq, br}, berr) + }}, + {"RRQRRank", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := RRQRRank(probe([]float64{2, 0, 0, 3, 0, 0}, 3, 2)) + b, berr := RRQRRank(base([]float64{2, 0, 0, 3, 0, 0}, 3, 2)) + cenInt(t, "RRQRRank", dt, p, perr, b, berr) + }}, + {"SolveRRQR", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := SolveRRQR(probe([]float64{2, 0, 0, 3, 0, 0}, 3, 2), probe([]float64{2, 6, 0}, 3)) + b, berr := SolveRRQR(base([]float64{2, 0, 0, 3, 0, 0}, 3, 2), base([]float64{2, 6, 0}, 3)) + cenArrays(t, "SolveRRQR", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SolveTridiagonal", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := SolveTridiagonal(probe([]float64{1, 1}, 2), probe([]float64{4, 4}, 2), + probe([]float64{1, 1}, 2), probe([]float64{1, 2, 3}, 3)) + b, berr := SolveTridiagonal(base([]float64{1, 1}, 2), base([]float64{4, 4}, 2), + base([]float64{1, 1}, 2), base([]float64{1, 2, 3}, 3)) + cenArrays(t, "SolveTridiagonal", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SolveCyclicTridiagonal", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := SolveCyclicTridiagonal(probe([]float64{1, 1, 1}, 3), probe([]float64{4, 4, 4}, 3), + probe([]float64{1, 1, 1}, 3), probe([]float64{1, 2, 3}, 3)) + b, berr := SolveCyclicTridiagonal(base([]float64{1, 1, 1}, 3), base([]float64{4, 4, 4}, 3), + base([]float64{1, 1, 1}, 3), base([]float64{1, 2, 3}, 3)) + cenArrays(t, "SolveCyclicTridiagonal", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SolveTruncated", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := SolveTruncated(probe([]float64{4, 1, 1, 3}, 2, 2), probe([]float64{1, 2}, 2), 2) + b, berr := SolveTruncated(base([]float64{4, 1, 1, 3}, 2, 2), base([]float64{1, 2}, 2), 2) + cenArrays(t, "SolveTruncated", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SolveTikhonov", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := SolveTikhonov(probe([]float64{4, 1, 1, 3}, 2, 2), probe([]float64{1, 2}, 2), 0.1) + b, berr := SolveTikhonov(base([]float64{4, 1, 1, 3}, 2, 2), base([]float64{1, 2}, 2), 0.1) + cenArrays(t, "SolveTikhonov", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"CholeskyUpdate", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := CholeskyUpdate(probe([]float64{2, 0, 0.5, 1.6583123951777}, 2, 2), probe([]float64{1, 1}, 2)) + b, berr := CholeskyUpdate(base([]float64{2, 0, 0.5, 1.6583123951777}, 2, 2), base([]float64{1, 1}, 2)) + cenArrays(t, "CholeskyUpdate", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"CholeskyDowndate", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + p, perr := CholeskyDowndate(probe([]float64{2, 0, 0.5, 1.6583123951777}, 2, 2), probe([]float64{0.5, 0.5}, 2)) + b, berr := CholeskyDowndate(base([]float64{2, 0, 0.5, 1.6583123951777}, 2, 2), base([]float64{0.5, 0.5}, 2)) + cenArrays(t, "CholeskyDowndate", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"NewCubicSpline", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + ps, perr := NewCubicSpline(probe([]float64{0, 1, 2, 3}, 4), probe([]float64{0, 1, 4, 9}, 4)) + bs, berr := NewCubicSpline(base([]float64{0, 1, 2, 3}, 4), base([]float64{0, 1, 4, 9}, 4)) + if berr != nil || perr != nil { + cenArrays(t, "NewCubicSpline", dt, nil, perr, nil, berr) + return + } + for _, x := range []float64{0.5, 1.25, 2.5} { + if ps.At(x) != bs.At(x) { + t.Fatalf("NewCubicSpline(%s): At(%g) = %v, want the baseline %v", dt, x, ps.At(x), bs.At(x)) + } + } + pe, perr := ps.Evaluate(probe([]float64{0.5, 1.25, 2.5}, 3)) + be, berr := bs.Evaluate(base([]float64{0.5, 1.25, 2.5}, 3)) + cenArrays(t, "CubicSpline.Evaluate", dt, []*core.Array{pe}, perr, []*core.Array{be}, berr) + }}, + // GMRES: the standing real-dtype gate refuses the narrow + // widths and bool by name (Int is accepted), and a narrow + // operator output is refused by the same family of wording. + {"GMRES b dtype gate", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + op := func(v *core.Array) (*core.Array, error) { return core.Add(v, v) } + p, perr := GMRES(op, probe([]float64{2, 4}, 2), 2, 5, 1e-10) + b, berr := GMRES(op, base([]float64{2, 4}, 2), 2, 5, 1e-10) + if dt == core.Int { + cenArrays(t, "GMRES int b", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + return + } + cenWantErr(t, "GMRES/"+dt.String(), perr, "GMRES", "b must be a real dtype", dt.String()) + if berr != nil { + t.Fatalf("GMRES float baseline failed: %v", berr) + } + }}, + {"GMRES op dtype gate", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt == core.Float { + return + } + narrowOp := func(v *core.Array) (*core.Array, error) { + return probe([]float64{v.FloatAt(0) * 2, v.FloatAt(1) * 2}, 2), nil + } + _, perr := GMRES(narrowOp, base([]float64{2, 4}, 2), 2, 5, 1e-10) + if dt == core.Int { + // An Int operator output passes the real-dtype gate + // and reads through the accessors: the Int treatment. + if perr != nil { + t.Fatalf("GMRES with an Int op output: %v", perr) + } + return + } + cenWantErr(t, "GMRES op/"+dt.String(), perr, "GMRES", "op returned", "want a real dtype", dt.String()) + }}, + // The pipeline surface: core elementwise dispatch answers the + // narrow widths, and the deferred core entries refuse loudly + // through Result. + {"Pipe AddF", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + pp, perr := Pipe(probe([]float64{1, 2}, 2)).AddF(1).Result() + bp, berr := Pipe(base([]float64{1, 2}, 2)).AddF(1).Result() + cenArrays(t, "Pipe.AddF", dt, []*core.Array{pp}, perr, []*core.Array{bp}, berr) + }}, + {"Pipe Add arrays", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt == core.Bool { + // Bool arithmetic is refused by the core throat; the + // pipeline records the refusal for Result. + _, err := Pipe(probe([]float64{1, 0}, 2)).Add(probe([]float64{1, 1}, 2)).Result() + cenWantErr(t, "Pipe.Add bool", err, "bool arrays have no arithmetic") + return + } + // Same-dtype arithmetic keeps the operand dtype on the + // narrow widths (promote(dt, dt) = dt); only the values + // are pinned against the float baseline. + pp, perr := Pipe(probe([]float64{1, 2}, 2)).Add(probe([]float64{3, 4}, 2)).Result() + bp, berr := Pipe(base([]float64{1, 2}, 2)).Add(base([]float64{3, 4}, 2)).Result() + if perr != nil || berr != nil { + t.Fatalf("Pipe.Add: probe %v, base %v", perr, berr) + } + if pp.Dtype() != dt { + t.Fatalf("Pipe.Add(%s) dtype = %s, want %s", dt, pp.Dtype(), dt) + } + pv, bv := cenWiden(t, pp), cenWiden(t, bp) + for i := range pv { + if pv[i] != bv[i] { + t.Fatalf("Pipe.Add(%s) element %d = %v, want %v", dt, i, pv[i], bv[i]) + } + } + }}, + {"Pipe Sqrt refusal ride-through", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + _, perr := Pipe(probe([]float64{4, 9}, 2)).Sqrt().Result() + bp, berr := Pipe(base([]float64{4, 9}, 2)).Sqrt().Result() + if berr != nil { + t.Fatalf("Pipe.Sqrt float baseline: %v", berr) + } + if dt == core.Int { + // Int reaches the float kernel: the Int treatment. + if perr != nil { + t.Fatalf("Pipe.Sqrt on Int: %v", perr) + } + _ = bp + return + } + cenWantErr(t, "Pipe.Sqrt/"+dt.String(), perr, "Sqrt", dt.String(), "convert with Astype") + }}, + // Sparse entries: the b-side and dense-operand surfaces ride + // the accessor walks; the value dtypes are the core + // constructors' own gate, pinned below. + {"SpSolve b", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + p, perr := SpSolve(s, probe([]float64{1, 2, 3}, 3), 1e-10, 50) + b, berr := SpSolve(s, base([]float64{1, 2, 3}, 3), 1e-10, 50) + cenArrays(t, "SpSolve", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SpSolveBiCGSTAB b", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + p, perr := SpSolveBiCGSTAB(s, probe([]float64{1, 2, 3}, 3), 1e-10, 50) + b, berr := SpSolveBiCGSTAB(s, base([]float64{1, 2, 3}, 3), 1e-10, 50) + cenArrays(t, "SpSolveBiCGSTAB", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SpLSQR b", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + p, _, perr := SpLSQR(s, probe([]float64{1, 2, 3}, 3), 1e-10, 50, 1e8) + b, _, berr := SpLSQR(s, base([]float64{1, 2, 3}, 3), 1e-10, 50, 1e8) + cenArrays(t, "SpLSQR", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SpLSMR b", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + p, _, perr := SpLSMR(s, probe([]float64{1, 2, 3}, 3), 1e-10, 50, 1e8) + b, _, berr := SpLSMR(s, base([]float64{1, 2, 3}, 3), 1e-10, 50, 1e8) + cenArrays(t, "SpLSMR", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SpExpApply v", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + p, perr := SpExpApply(s, probe([]float64{1, 2, 3}, 3), 3) + b, berr := SpExpApply(s, base([]float64{1, 2, 3}, 3), 3) + cenArrays(t, "SpExpApply", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SparseLU Solve b", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + f, ferr := NewSparseLU(s) + if ferr != nil { + t.Fatalf("NewSparseLU: %v", ferr) + } + p, perr := f.Solve(probe([]float64{1, 2, 3}, 3)) + b, berr := f.Solve(base([]float64{1, 2, 3}, 3)) + cenArrays(t, "SparseLU.Solve", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SparseCholesky Solve b", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + f, ferr := NewSparseCholesky(s, SparseOrderingNatural) + if ferr != nil { + t.Fatalf("NewSparseCholesky: %v", ferr) + } + p, perr := f.Solve(probe([]float64{1, 2, 3}, 3)) + b, berr := f.Solve(base([]float64{1, 2, 3}, 3)) + cenArrays(t, "SparseCholesky.Solve", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SparseILU precond b", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + ilu, ierr := NewSparseILU(s) + if ierr != nil { + t.Fatalf("NewSparseILU: %v", ierr) + } + p, perr := SpSolve(s, probe([]float64{1, 2, 3}, 3), 1e-10, 50, ilu) + b, berr := SpSolve(s, base([]float64{1, 2, 3}, 3), 1e-10, 50, ilu) + cenArrays(t, "SpSolve+ILU", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SpEigen Int-valued sparse computes", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt != core.Int { + // Narrow sparse values are the core constructors' + // refusal; only Int values reach this entry. + return + } + pi := cenSparse(t, probe, []float64{4, 5, 6}, 3) + bf := cenSparse(t, base, []float64{4, 5, 6}, 3) + pv, _, perr := SpEigen(pi, 2, core.NewGenerator(7)) + bv, _, berr := SpEigen(bf, 2, core.NewGenerator(7)) + cenArrays(t, "SpEigen", dt, []*core.Array{pv}, perr, []*core.Array{bv}, berr) + }}, + {"SpEigenGeneral Int-valued sparse computes", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt != core.Int { + // Narrow sparse values are the core constructors' + // refusal, pinned in the root census; only Int values + // reach this entry. + return + } + pi := cenSparse(t, probe, []float64{4, 5, 6}, 3) + bf := cenSparse(t, base, []float64{4, 5, 6}, 3) + pv, _, perr := SpEigenGeneral(pi, 2, core.NewGenerator(7)) + bv, _, berr := SpEigenGeneral(bf, 2, core.NewGenerator(7)) + cenArrays(t, "SpEigenGeneral", dt, []*core.Array{pv}, perr, []*core.Array{bv}, berr) + }}, + {"SpSolveComplexCG refuses real values", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + _, perr := SpSolveComplexCG(s, probe([]float64{1, 2, 3}, 3), 1e-10, 10) + cenWantErr(t, "SpSolveComplexCG real values", perr, "SpSolveComplexCG", "complex128") + }}, + {"SpSolveComplexCG refuses real b on complex values", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + ind, err := core.FromInts([]int64{0, 0, 1, 1, 2, 2}, 3, 2) + if err != nil { + t.Fatal(err) + } + cv, err := core.FromComplexes([]complex128{4, 5, 6}, 3) + if err != nil { + t.Fatal(err) + } + s, err := core.NewSparseCOO(ind, cv, []int{3, 3}) + if err != nil { + t.Fatal(err) + } + _, perr := SpSolveComplexCG(s, probe([]float64{1, 2, 3}, 3), 1e-10, 10) + cenWantErr(t, "SpSolveComplexCG b/"+dt.String(), perr, "b must be complex128", dt.String()) + }}, + {"SpEigenGeneralComplex refuses real values", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + _, _, perr := SpEigenGeneralComplex(s, 2, core.NewGenerator(7)) + cenWantErr(t, "SpEigenGeneralComplex", perr, "complex128") + }}, + {"CSRFromCOO Int values compute", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt != core.Int { + return + } + pi := cenSparse(t, probe, []float64{4, 5, 6}, 3) + bf := cenSparse(t, base, []float64{4, 5, 6}, 3) + p, perr := CSRFromCOO(pi) + b, berr := CSRFromCOO(bf) + if perr != nil || berr != nil { + t.Fatalf("CSRFromCOO: probe %v, base %v", perr, berr) + } + if len(p.Values) != len(b.Values) { + t.Fatalf("CSR values length %d, want %d", len(p.Values), len(b.Values)) + } + for i := range p.Values { + if p.Values[i] != b.Values[i] { + t.Fatalf("CSR value %d = %v, want %v", i, p.Values[i], b.Values[i]) + } + } + }}, + {"CSCFromCOO Int values compute", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt != core.Int { + return + } + pi := cenSparse(t, probe, []float64{4, 5, 6}, 3) + bf := cenSparse(t, base, []float64{4, 5, 6}, 3) + p, perr := CSCFromCOO(pi) + b, berr := CSCFromCOO(bf) + if perr != nil || berr != nil { + t.Fatalf("CSCFromCOO: probe %v, base %v", perr, berr) + } + for i := range p.Values { + if p.Values[i] != b.Values[i] { + t.Fatalf("CSC value %d = %v, want %v", i, p.Values[i], b.Values[i]) + } + } + }}, + {"NewSparseCOO narrow values refused by core", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt == core.Int || dt == core.Float { + return // accepted value dtypes, covered above + } + ind, err := core.FromInts([]int64{0, 0}, 1, 2) + if err != nil { + t.Fatal(err) + } + _, serr := core.NewSparseCOO(ind, probe([]float64{4}, 1), []int{2, 2}) + cenWantErr(t, "NewSparseCOO/"+dt.String(), serr, "NewSparseCOO", dt.String(), "convert with Astype") + }}, + // The generalised eigensolver widens both operand matrices + // through the accessor walks: narrow follows Int bit for bit. + {"EigenGeneralised", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + v, e, perr := EigenGeneralised(probe([]float64{4, 1, 1, 3}, 2, 2), probe([]float64{1, 0, 0, 1}, 2, 2)) + bv, be, berr := EigenGeneralised(base([]float64{4, 1, 1, 3}, 2, 2), base([]float64{1, 0, 0, 1}, 2, 2)) + cenArrays(t, "EigenGeneralised", dt, []*core.Array{v, e}, perr, []*core.Array{bv, be}, berr) + }}, + // The complex sparse solvers keep their standing dtype gates, + // wording verbatim from sparsecomplex.go: the values gate fires + // before b on real-valued sparse, and b is refused by name on + // complex values, narrow widths following Int into the same + // refusal. + {"SpSolveComplexBiCGSTAB values gate", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + s := cenSparse(t, base, []float64{4, 5, 6}, 3) + _, perr := SpSolveComplexBiCGSTAB(s, probe([]float64{1, 2, 3}, 3), 1e-10, 50) + cenWantErr(t, "SpSolveComplexBiCGSTAB/"+dt.String(), perr, + "SpSolveComplexBiCGSTAB", "needs complex128 values", "float") + _, berr := SpSolveComplexBiCGSTAB(s, base([]float64{1, 2, 3}, 3), 1e-10, 50) + cenWantErr(t, "SpSolveComplexBiCGSTAB float baseline", berr, + "SpSolveComplexBiCGSTAB", "needs complex128 values", "float") + }}, + {"SpSolveComplexBiCGSTAB b gate on complex values", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + ind, err := core.FromInts([]int64{0, 0, 1, 1, 2, 2}, 3, 2) + if err != nil { + t.Fatal(err) + } + cv, err := core.FromComplexes([]complex128{4, 5, 6}, 3) + if err != nil { + t.Fatal(err) + } + s, err := core.NewSparseCOO(ind, cv, []int{3, 3}) + if err != nil { + t.Fatal(err) + } + _, perr := SpSolveComplexBiCGSTAB(s, probe([]float64{1, 2, 3}, 3), 1e-10, 50) + cenWantErr(t, "SpSolveComplexBiCGSTAB b/"+dt.String(), perr, + "SpSolveComplexBiCGSTAB", "b must be complex128", dt.String()) + }}, + {"SpEigenComplex value gate", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + // Int-valued sparse reaches the entry and hits its values + // gate; the narrow widths and bool ride the core sparse + // constructor refusal before the entry is ever reached. + // Both refusals are loud and keep their recorded wording. + ps, serr := cenSparseErr(probe, []float64{4, 5, 6}, 3) + if serr != nil { + cenWantErr(t, "SpEigenComplex/"+dt.String(), serr, + "NewSparseCOO", dt.String(), "convert with Astype") + } else { + _, _, perr := SpEigenComplex(ps, 2, core.NewGenerator(7)) + cenWantErr(t, "SpEigenComplex/"+dt.String(), perr, + "SpEigenComplex", "needs complex128 values", dt.String()) + } + bs, berr := cenSparseErr(base, []float64{4, 5, 6}, 3) + if berr != nil { + t.Fatalf("float baseline sparse: %v", berr) + } + _, _, fberr := SpEigenComplex(bs, 2, core.NewGenerator(7)) + cenWantErr(t, "SpEigenComplex float baseline", fberr, + "SpEigenComplex", "needs complex128 values", "float") + }}, + // The pipeline's elementwise surface: Bool rides the core + // arithmetic refusal through Result, the narrow widths and Int + // keep the operand dtype and compute the widened values. + {"Pipe elementwise keeps operand dtype", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + for _, row := range []struct { + name string + run func(mk cenMaker) (*core.Array, error) + }{ + {"Sub", func(mk cenMaker) (*core.Array, error) { + return Pipe(mk([]float64{3, 4}, 2)).Sub(mk([]float64{1, 2}, 2)).Result() + }}, + {"Mul", func(mk cenMaker) (*core.Array, error) { + return Pipe(mk([]float64{3, 4}, 2)).Mul(mk([]float64{2, 5}, 2)).Result() + }}, + {"Maximum", func(mk cenMaker) (*core.Array, error) { + return Pipe(mk([]float64{3, 1}, 2)).Maximum(mk([]float64{2, 5}, 2)).Result() + }}, + {"Minimum", func(mk cenMaker) (*core.Array, error) { + return Pipe(mk([]float64{3, 1}, 2)).Minimum(mk([]float64{2, 5}, 2)).Result() + }}, + } { + if dt == core.Bool { + _, err := row.run(probe) + cenWantErr(t, "Pipe."+row.name+" bool", err, row.name, "bool arrays have no arithmetic") + continue + } + pp, perr := row.run(probe) + bp, berr := row.run(base) + if perr != nil || berr != nil { + t.Fatalf("Pipe.%s: probe %v, base %v", row.name, perr, berr) + } + if pp.Dtype() != dt { + t.Fatalf("Pipe.%s(%s) dtype = %s, want %s", row.name, dt, pp.Dtype(), dt) + } + pv, bv := cenWiden(t, pp), cenWiden(t, bp) + for i := range pv { + if pv[i] != bv[i] { + t.Fatalf("Pipe.%s(%s) element %d = %v, want %v", row.name, dt, i, pv[i], bv[i]) + } + } + } + }}, + // Div answers the float baseline exactly on the integer class + // (the quotient is not representable in the operand dtype), and + // Bool keeps the arithmetic refusal. + {"Pipe Div", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt == core.Bool { + _, err := Pipe(probe([]float64{6, 8}, 2)).Div(probe([]float64{3, 4}, 2)).Result() + cenWantErr(t, "Pipe.Div bool", err, "Div", "bool arrays have no arithmetic") + return + } + pp, perr := Pipe(probe([]float64{6, 8}, 2)).Div(probe([]float64{3, 4}, 2)).Result() + bp, berr := Pipe(base([]float64{6, 8}, 2)).Div(base([]float64{3, 4}, 2)).Result() + cenArrays(t, "Pipe.Div", dt, []*core.Array{pp}, perr, []*core.Array{bp}, berr) + }}, + // MatMul2D through the pipeline rides the core narrow-dtype + // deferral on Bool and the narrow widths; Int computes in its + // own dtype, its exact products equal to the float baseline's. + {"Pipe MatMul2D", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + if dt != core.Int { + _, err := Pipe(probe([]float64{1, 2, 3, 4}, 2, 2)).MatMul2D(probe([]float64{1, 0, 0, 1}, 2, 2)).Result() + cenWantErr(t, "Pipe.MatMul2D/"+dt.String(), err, "MatMul", dt.String(), "convert with Astype") + return + } + pp, perr := Pipe(probe([]float64{1, 2, 3, 4}, 2, 2)).MatMul2D(probe([]float64{5, 6, 7, 8}, 2, 2)).Result() + bp, berr := Pipe(base([]float64{1, 2, 3, 4}, 2, 2)).MatMul2D(base([]float64{5, 6, 7, 8}, 2, 2)).Result() + if perr != nil || berr != nil { + t.Fatalf("Pipe.MatMul2D int: probe %v, base %v", perr, berr) + } + if pp.Dtype() != core.Int { + t.Fatalf("Pipe.MatMul2D(int) dtype = %s, want int", pp.Dtype()) + } + pv, bv := cenWiden(t, pp), cenWiden(t, bp) + for i := range pv { + if pv[i] != bv[i] { + t.Fatalf("Pipe.MatMul2D(int) element %d = %v, want %v", i, pv[i], bv[i]) + } + } + }}, + // Solve through the pipeline widens through the accessor walks: + // narrow follows Int bit for bit. + {"Pipe Solve", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + pp, perr := Pipe(probe([]float64{4, 1, 1, 3}, 2, 2)).Solve(probe([]float64{1, 2}, 2)).Result() + bp, berr := Pipe(base([]float64{4, 1, 1, 3}, 2, 2)).Solve(base([]float64{1, 2}, 2)).Result() + cenArrays(t, "Pipe.Solve", dt, []*core.Array{pp}, perr, []*core.Array{bp}, berr) + }}, + // The sparse format methods read their dense operand through + // the accessor walks and answer float64: identical to the + // baseline whatever dtype the operand carries. + {"SparseCSR.MatVec operand", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + cs, err := CSRFromCOO(cenSparse(t, base, []float64{4, 5, 6}, 3)) + if err != nil { + t.Fatal(err) + } + p, perr := cs.MatVec(probe([]float64{1, 2, 3}, 3)) + b, berr := cs.MatVec(base([]float64{1, 2, 3}, 3)) + cenArrays(t, "SparseCSR.MatVec", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SparseCSR.MatMulDense operand", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + cs, err := CSRFromCOO(cenSparse(t, base, []float64{4, 5, 6}, 3)) + if err != nil { + t.Fatal(err) + } + p, perr := cs.MatMulDense(probe([]float64{1, 0, 0, 1, 1, 0}, 3, 2)) + b, berr := cs.MatMulDense(base([]float64{1, 0, 0, 1, 1, 0}, 3, 2)) + cenArrays(t, "SparseCSR.MatMulDense", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SparseCSC.MatVec operand", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + cs, err := CSCFromCOO(cenSparse(t, base, []float64{4, 5, 6}, 3)) + if err != nil { + t.Fatal(err) + } + p, perr := cs.MatVec(probe([]float64{1, 2, 3}, 3)) + b, berr := cs.MatVec(base([]float64{1, 2, 3}, 3)) + cenArrays(t, "SparseCSC.MatVec", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + // The rank-1 modification surface reads its update vector + // through the accessor walks: every probe dtype is accepted on + // a factor whose pattern holds the modified entry. + {"SparseCholesky.Update vector", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + fp, perr := NewSparseCholesky(cenSPDFull(), SparseOrderingNatural) + if perr != nil { + t.Fatal(perr) + } + fb, berr := NewSparseCholesky(cenSPDFull(), SparseOrderingNatural) + if berr != nil { + t.Fatal(berr) + } + if uerr := fp.Update(probe([]float64{1, 1}, 2)); uerr != nil { + t.Fatalf("SparseCholesky.Update(%s): %v", dt, uerr) + } + if uerr := fb.Update(base([]float64{1, 1}, 2)); uerr != nil { + t.Fatalf("SparseCholesky.Update float baseline: %v", uerr) + } + }}, + {"SparseCholesky.Downdate vector", func(t *testing.T, probe, base cenMaker, dt core.Dtype) { + fp, perr := NewSparseCholesky(cenSPDFull(), SparseOrderingNatural) + if perr != nil { + t.Fatal(perr) + } + fb, berr := NewSparseCholesky(cenSPDFull(), SparseOrderingNatural) + if berr != nil { + t.Fatal(berr) + } + if uerr := fp.Downdate(probe([]float64{1, 1}, 2)); uerr != nil { + t.Fatalf("SparseCholesky.Downdate(%s): %v", dt, uerr) + } + if uerr := fb.Downdate(base([]float64{1, 1}, 2)); uerr != nil { + t.Fatalf("SparseCholesky.Downdate float baseline: %v", uerr) + } + }}, + } + for _, row := range rows { + for _, dt := range cenDtypes { + t.Run(row.name+"/"+dt.String(), func(t *testing.T) { + probe, base := cenMakers(t, dt) + row.run(t, probe, base, dt) + }) + } + } +} + +// TestDtypesCensusLinalgFloat16NarrowGate pins the Float16 treatment +// GMRES carries: refused by name alongside the narrow widths, the +// wording the standing gate has always produced. +func TestDtypesCensusLinalgFloat16NarrowGate(t *testing.T) { + f16, err := core.Astype(mustCenFloats(t, 2, 4), core.Float16) + if err != nil { + t.Fatal(err) + } + op := func(v *core.Array) (*core.Array, error) { return core.Add(v, v) } + _, gerr := GMRES(op, f16, 2, 5, 1e-10) + cenWantErr(t, "GMRES float16 b", gerr, "GMRES", "b must be a real dtype", "float16") +} + +func mustCenFloats(t *testing.T, vals ...float64) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + return a +} diff --git a/linalg/eigenreal.go b/linalg/eigenreal.go new file mode 100644 index 0000000..1e0dcbf --- /dev/null +++ b/linalg/eigenreal.go @@ -0,0 +1,446 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "slices" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The general eigenproblem. Eigen answers real symmetric matrices and +// EigenComplex Hermitian ones, where structure pays for itself. A +// general square matrix has no structure to exploit, so EigenGeneral +// goes through the complex Schur form: the input is promoted to +// complex128, reduced to upper Hessenberg form by complex Householder +// reflectors, and driven to upper triangular form by the shifted QR +// iteration, one Wilkinson-shifted sweep of Givens rotations per step, +// deflating whenever a subdiagonal entry falls under a scale-relative +// tolerance. Working over the complex plane is what keeps the +// iteration single-shift: conjugate eigenvalue pairs are ordinary +// points there, where a real iteration would need the double-shift +// bulge chase to reach them. +// +// Eigenvectors come from the triangular Schur factor by back +// substitution and one multiplication with the accumulated unitary +// similarity, which keeps them orthonormal to rounding level even +// when the eigenvalues themselves are ill conditioned. + +// EigenGeneral returns the eigenvalues and eigenvectors of a square +// matrix of any dtype; real and integer inputs are promoted to +// complex128. Values is a complex128 vector sorted descending by +// magnitude, bitwise ties broken by descending real part and then +// descending imaginary part; conjugate pairs share a magnitude only +// to rounding, so their relative order follows the rounding noise +// rather than the tiebreak. Vectors is a complex128 (n, n) array +// whose column j is the unit eigenvector for values[j]. Symmetric +// real matrices get a faster answer from Eigen and Hermitian matrices +// from EigenComplex. +func EigenGeneral(a *core.Array) (values, vectors *core.Array, err error) { + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, nil, base.Errf("EigenGeneral: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + if n == 0 { + return nil, nil, base.Errf("EigenGeneral: zero-sized matrix, got shape %s", base.ShapeText(a.Shape())) + } + h := make([]complex128, n*n) + switch { + case a.Dtype() == core.Complex && !a.Strided() && len(a.RawComplexes()) == n*n: + copy(h, a.RawComplexes()) + case a.Dtype() == core.Float && !a.Strided() && len(a.RawFloats()) == n*n: + for i := range n * n { + h[i] = complex(a.RawFloats()[i], 0) + } + default: + for i := range n * n { + h[i] = a.ComplexAt(i) + } + } + // Outside the safe window the reflector norms, the norm the shift + // floors are taken against and the squared magnitudes the sweeps form + // all leave the normal range: a tiny matrix collapses to zeros and a + // huge one has vMax*vMax overflow, which zeroes beta so the + // Hessenberg reduction is skipped and the QR iteration then fails on + // a perfectly valid matrix. The matrix is moved into the window and + // the eigenvalues, which carry the scale, are moved back on return; + // the Schur vectors are scale-free. + ws := windowScale(maxMagComplex(h)) + if ws != 1 { + scaleComplexes(h, ws) + } + scale := 0.0 + for _, z := range h { + if m := base.AbsComplex(z); m > scale { + scale = m + } + } + // q accumulates the unitary similarity: the input satisfies + // A = q·H·qᴴ at every stage. + q := eyeComplex(n) + hessenbergComplex(h, q, n) + if err := hessenbergQr(h, q, n, scale); err != nil { + return nil, nil, err + } + + vals := make([]complex128, n) + for i := range n { + vals[i] = h[i*n+i] + } + idx := sortMagDescIndices(vals) + outVals := make([]complex128, n) + for j := range n { + outVals[j] = vals[idx[j]] + } + // Eigenvector j depends on the position of λ_j inside the Schur + // factor, so they are built in place order and permuted afterwards. + inPlace := make([]complex128, n*n) + for j := range n { + v := schurEigenVector(h, q, n, j, scale) + for i := range n { + inPlace[i*n+j] = v[i] + } + } + outVecs := make([]complex128, n*n) + for j := range n { + for i := range n { + outVecs[i*n+j] = inPlace[i*n+idx[j]] + } + } + // Only the eigenvalues carry the matrix's scale. + if ws != 1 { + unscaleComplexes(outVals, ws) + } + valuesArr := core.New(core.Complex, []int{n}...) + copy(valuesArr.RawComplexes(), outVals) + vecArr := core.New(core.Complex, []int{n, n}...) + copy(vecArr.RawComplexes(), outVecs) + return valuesArr, vecArr, nil +} + +// hessenbergComplex reduces h to upper Hessenberg form in place by +// complex Householder reflectors, accumulating the unitary similarity +// into q so that the original matrix satisfies A = q·H·qᴴ. +func hessenbergComplex(h, q []complex128, n int) { + for k := range n - 2 { + // The reflector zeroes column k below its subdiagonal; columns + // already in form are skipped. Squared magnitudes are summed + // relative to the column's largest entry so a column with + // entries near 1e154 cannot overflow on the way to its norm. + scale := 0.0 + for i := k + 1; i < n; i++ { + if m := base.AbsComplex(h[i*n+k]); m > scale { + scale = m + } + } + below := 0.0 + if scale > 0 { + for i := k + 2; i < n; i++ { + m := base.AbsComplex(h[i*n+k]) / scale + below += m * m + } + } + if below == 0 { + continue + } + x0 := h[(k+1)*n+k] + norm := scale * math.Hypot(base.AbsComplex(x0)/scale, math.Sqrt(below)) + // alpha = −sign(x0)·‖x‖ lands the reflection on the far side + // of x0, away from the cancellation zone. + phase := complex(1, 0) + if m := base.AbsComplex(x0); m > 0 { + phase = x0 / complex(m, 0) + } + alpha := -phase * complex(norm, 0) + v := make([]complex128, n-k-1) + v[0] = x0 - alpha + for i := 1; i < n-k-1; i++ { + v[i] = h[(k+1+i)*n+k] + } + vMax := 0.0 + for _, z := range v { + if m := base.AbsComplex(z); m > vMax { + vMax = m + } + } + vNorm2 := 0.0 + if vMax > 0 { + for _, z := range v { + re := real(z) / vMax + im := imag(z) / vMax + vNorm2 += re*re + im*im + } + } + // beta = 2/(vᴴv) needs vMax*vMax to stay in the normal range. A + // column whose entries are extreme in the matrix's own units + // (far below it, say) overflows that square to +Inf, which makes + // beta +0, or underflows it to 0, which makes beta +Inf and the + // update NaN. Both are silent corruption of an otherwise valid + // reduction, so the reflector is re-expressed in window units + // with the scaling compensated in beta, exactly as + // householderVectorInto does for the real reflectors. + betaR := 2 / (vMax * vMax * vNorm2) + if betaR <= 0 || math.IsInf(betaR, 0) { + wsV := windowScale(vMax) + for i := range v { + v[i] = complex(real(v[i])*wsV, imag(v[i])*wsV) + } + w := vMax * wsV + betaR = 2 / (w * w * vNorm2) + } + beta := complex(betaR, 0) + + // H = P·H with P = I − β·v·vᴴ, rows k+1..n−1, columns k..n−1. + for j := k; j < n; j++ { + s := complex(0, 0) + for i := range v { + s += cmplx.Conj(v[i]) * h[(k+1+i)*n+j] + } + s *= beta + for i := range v { + h[(k+1+i)*n+j] -= s * v[i] + } + } + // H = H·P, every row, columns k+1..n−1. + for i := range n { + s := complex(0, 0) + for j := range v { + s += h[i*n+k+1+j] * v[j] + } + s *= beta + for j := range v { + h[i*n+k+1+j] -= s * cmplx.Conj(v[j]) + } + } + // q = q·P keeps the similarity: A = q·H·qᴴ. + for i := range n { + s := complex(0, 0) + for j := range v { + s += q[i*n+k+1+j] * v[j] + } + s *= beta + for j := range v { + q[i*n+k+1+j] -= s * cmplx.Conj(v[j]) + } + } + // Column k below the subdiagonal is zero by construction; make + // it exactly so instead of carrying rounding residue. + for i := k + 2; i < n; i++ { + h[i*n+k] = 0 + } + } +} + +// hessenbergQr drives the shifted QR iteration on an upper Hessenberg +// matrix until every subdiagonal entry deflates, accumulating the +// Schur similarity into q. On return h is upper triangular: the Schur +// form of the original matrix, with its eigenvalues on the diagonal. +func hessenbergQr(h, q []complex128, n int, scale float64) error { + // Purely relative to the matrix magnitude: an absolute floor would + // treat a legitimate tiny-scale matrix (a 1e-20 rotation, say) as + // one big deflated block and return zeros. + floor := base.EpsF * scale + negligible := func(i int) bool { + local := base.EpsF * (base.AbsComplex(h[i*n+i]) + base.AbsComplex(h[(i-1)*n+i-1])) + return base.AbsComplex(h[i*n+i-1]) <= math.Max(local, floor) + } + // Rotation scratch reused across sweeps; each sweep uses the first + // hi−lo entries of each. + cs := make([]complex128, n) + sn := make([]complex128, n) + hi := n - 1 + iter := 0 + for hi > 0 { + for hi > 0 && negligible(hi) { + h[hi*n+hi-1] = 0 + hi-- + iter = 0 + } + if hi == 0 { + return nil + } + lo := hi + for lo > 0 && !negligible(lo) { + lo-- + } + if lo > 0 { + // A negligible entry splits the matrix; the block above it + // is settled in later rounds. + h[lo*n+lo-1] = 0 + } + iter++ + if iter > 100 { + return base.Errf("EigenGeneral: QR iteration failed to converge at row %d", hi) + } + var mu complex128 + if iter%10 == 0 { + // An exceptional shift breaks the rare cyclic pattern the + // Wilkinson shift can settle into. + mu = h[hi*n+hi] + complex(0.75*base.AbsComplex(h[hi*n+hi-1]), 0) + } else { + mu = wilkinsonShift(h, n, hi) + } + qrSweepComplex(h, q, n, lo, hi, mu, cs, sn) + } + return nil +} + +// wilkinsonShift returns the eigenvalue of the trailing 2×2 block +// closest to its bottom-right entry: the shift that deflates the +// subdiagonal under it fastest. +func wilkinsonShift(h []complex128, n, i int) complex128 { + a := h[(i-1)*n+i-1] + b := h[(i-1)*n+i] + c := h[i*n+i-1] + d := h[i*n+i] + delta := (a - d) / 2 + disc := cmplx.Sqrt(delta*delta + b*c) + low := d + delta - disc + high := d + delta + disc + if base.AbsComplex(low-d) <= base.AbsComplex(high-d) { + return low + } + return high +} + +// qrSweepComplex performs one explicit single-shift QR step over the +// active block [lo, hi]: H = G·(H − μI)·Gᴴ + μI for the sequence of +// Givens rotations G that triangularises H − μI, with every rotation +// folded into q so the similarity stays exact. Each rotation satisfies +// G·[a; b] = [r; 0] with r real, which is what zeroes the subdiagonal +// one entry per step. cs and sn are the caller's scratch for the +// rotations; each sweep fully overwrites its first hi−lo entries. +func qrSweepComplex(h, q []complex128, n, lo, hi int, mu complex128, cs, sn []complex128) { + for i := lo; i <= hi; i++ { + h[i*n+i] -= mu + } + for j := lo; j < hi; j++ { + a, b := h[j*n+j], h[(j+1)*n+j] + r := math.Hypot(base.AbsComplex(a), base.AbsComplex(b)) + var c, s complex128 + if r == 0 { + c, s = 1, 0 + } else { + c, s = a/complex(r, 0), b/complex(r, 0) + } + cs[j-lo], sn[j-lo] = c, s + // Left: rows j and j+1 over columns j..n−1. + for k := j; k < n; k++ { + x, y := h[j*n+k], h[(j+1)*n+k] + h[j*n+k] = cmplx.Conj(c)*x + cmplx.Conj(s)*y + h[(j+1)*n+k] = -s*x + c*y + } + h[(j+1)*n+j] = 0 + } + for j := lo; j < hi; j++ { + c, s := cs[j-lo], sn[j-lo] + // Right: columns j and j+1 down to row j+1, the deepest row + // the triangular factor can reach in either column. + for i := 0; i <= j+1; i++ { + x, y := h[i*n+j], h[i*n+j+1] + h[i*n+j] = x*c + y*s + h[i*n+j+1] = -x*cmplx.Conj(s) + y*cmplx.Conj(c) + } + for i := range n { + x, y := q[i*n+j], q[i*n+j+1] + q[i*n+j] = x*c + y*s + q[i*n+j+1] = -x*cmplx.Conj(s) + y*cmplx.Conj(c) + } + } + for i := lo; i <= hi; i++ { + h[i*n+i] += mu + } +} + +// schurEigenVector builds the unit eigenvector for the diagonal entry +// j of a triangular Schur factor t with its similarity q (the input +// satisfies A = q·t·qᴴ): back substitution on the leading block gives +// the coordinates in Schur space, one multiplication with q lifts +// them back. A nearly multiple eigenvalue perturbs the denominator +// off exact zero, the standard defence against dividing by the +// spectrum's own degeneracy. +func schurEigenVector(t, q []complex128, n, j int, scale float64) []complex128 { + x := make([]complex128, n) + x[j] = 1 + lam := t[j*n+j] + // Relative to the matrix magnitude for the same reason as the QR + // floor above: tiny-scale spectra deserve their eigenvectors too. + denFloor := base.EpsF * scale + for i := j - 1; i >= 0; i-- { + s := -t[i*n+j] + for k := i + 1; k < j; k++ { + s -= t[i*n+k] * x[k] + } + d := t[i*n+i] - lam + if base.AbsComplex(d) < denFloor { + d = complex(denFloor, 0) + } + if base.AbsComplex(d) == 0 { + // The zero matrix: every denominator vanishes, and the + // coordinate is unconstrained. Zero keeps the vector finite + // where 0/0 would poison it with NaN. + x[i] = 0 + continue + } + x[i] = s / d + } + v := make([]complex128, n) + for i := range n { + acc := complex(0, 0) + for k := range j + 1 { + acc += q[i*n+k] * x[k] + } + v[i] = acc + } + norm := 0.0 + for _, z := range v { + norm += real(z)*real(z) + imag(z)*imag(z) + } + norm = math.Sqrt(norm) + if norm > 0 { + for i := range v { + v[i] /= complex(norm, 0) + } + } + return v +} + +// sortMagDescIndices sorts indices so the values come descending by +// magnitude, bitwise ties broken by descending real part and then +// descending imaginary part. The sort is stable, so ties beyond the +// comparator keep their original order exactly as before. +func sortMagDescIndices(vals []complex128) []int { + idx := make([]int, len(vals)) + for i := range idx { + idx[i] = i + } + slices.SortStableFunc(idx, func(a, b int) int { + ma, mb := base.AbsComplex(vals[a]), base.AbsComplex(vals[b]) + if ma != mb { + if ma > mb { + return -1 + } + return 1 + } + ra, rb := real(vals[a]), real(vals[b]) + if ra != rb { + if ra > rb { + return -1 + } + return 1 + } + ia, ib := imag(vals[a]), imag(vals[b]) + switch { + case ia > ib: + return -1 + case ia < ib: + return 1 + default: + return 0 + } + }) + return idx +} diff --git a/linalg/eigenreal_test.go b/linalg/eigenreal_test.go new file mode 100644 index 0000000..37a9a4a --- /dev/null +++ b/linalg/eigenreal_test.go @@ -0,0 +1,307 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// eigenResidual returns ‖A·v_j − λ_j·v_j‖∞ for every eigenpair j of a +// square matrix a. +func eigenResidual(a *core.Array, values, vectors *core.Array) []float64 { + n := a.Shape()[0] + out := make([]float64, n) + for j := range n { + lam := values.ComplexAt(j) + worst := 0.0 + for i := range n { + acc := complex(0, 0) + for k := range n { + acc += a.ComplexAt(i*n+k) * vectors.ComplexAt(k*n+j) + } + acc -= lam * vectors.ComplexAt(i*n+j) + if m := absComplex(acc); m > worst { + worst = m + } + } + out[j] = worst + } + return out +} + +// checkEigenPairs asserts every eigenpair satisfies A·v = λ·v to +// rounding level and every vector has unit norm. +func checkEigenPairs(t *testing.T, a *core.Array, values, vectors *core.Array, tol float64) { + t.Helper() + n := a.Shape()[0] + for _, r := range eigenResidual(a, values, vectors) { + if r > tol { + t.Fatalf("eigenpair residual %g, want ≤ %g", r, tol) + } + } + for j := range n { + norm := 0.0 + for i := range n { + z := vectors.ComplexAt(i*n + j) + norm += real(z)*real(z) + imag(z)*imag(z) + } + if math.Abs(norm-1) > 1e-12 { + t.Fatalf("vector %d has norm² %g, want 1", j, norm) + } + } +} + +// matchComplexSpectrum asserts every expected value appears in the +// computed spectrum and the magnitudes come out descending; within a +// magnitude tie (conjugate pairs, ±λ) the computed values are equal +// only to rounding, so the order inside the tie is not asserted. +func matchComplexSpectrum(t *testing.T, values *core.Array, want []complex128, tol float64) { + t.Helper() + used := make([]bool, len(want)) + for i := range values.Len() { + got := values.ComplexAt(i) + found := false + for k, w := range want { + if !used[k] && absComplex(got-w) <= tol { + used[k] = true + found = true + break + } + } + if !found { + t.Fatalf("value[%d] = %v has no expected match within %g", i, got, tol) + } + if i > 0 && absComplex(values.ComplexAt(i-1)) < absComplex(got) { + t.Fatalf("magnitude order broken at %d: |%v| < |%v|", i, values.ComplexAt(i-1), got) + } + } +} + +// TestEigenGeneralPauli checks the σx matrix: the eigenvalues are ±1. +func TestEigenGeneralPauli(t *testing.T) { + a := mustFloats(t, []float64{0, 1, 1, 0}, 2, 2) + values, vectors, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + matchComplexSpectrum(t, values, []complex128{1, -1}, 1e-12) + checkEigenPairs(t, a, values, vectors, 1e-14) +} + +// TestEigenGeneralRotation checks the rotation generator [[0, −θ], +// [θ, 0]]: the eigenvalues are the purely imaginary pair ±iθ. +func TestEigenGeneralRotation(t *testing.T) { + const theta = 0.7 + a := mustFloats(t, []float64{0, -theta, theta, 0}, 2, 2) + values, vectors, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + matchComplexSpectrum(t, values, []complex128{complex(0, theta), complex(0, -theta)}, 1e-12) + checkEigenPairs(t, a, values, vectors, 1e-14) +} + +// TestEigenGeneralKnownSpectrum checks a 6×6 real nonsymmetric matrix +// assembled as V·D·V⁻¹, whose spectrum is known by construction: +// 5, −3, the conjugate pair 2 ± 1.5i and the conjugate pair +// 0.5 ± 4i. The expected order follows the descending-magnitude sort +// with the documented tiebreaks: 5, 0.5+4i, 0.5−4i, −3, 2+1.5i, +// 2−1.5i. The trace and determinant cross-check the spectrum through +// two independent invariants. +func TestEigenGeneralKnownSpectrum(t *testing.T) { + d := mustFloats(t, []float64{ + 5, 0, 0, 0, 0, 0, + 0, -3, 0, 0, 0, 0, + 0, 0, 2, -1.5, 0, 0, + 0, 0, 1.5, 2, 0, 0, + 0, 0, 0, 0, 0.5, -4, + 0, 0, 0, 0, 4, 0.5, + }, 6, 6) + v := mustFloats(t, []float64{ + 1, 0.2, -0.2, 0, 0.2, 0, + 0.2, 1, 0, 0.2, -0.2, 0.2, + -0.2, 0, 1, 0.2, 0, -0.2, + 0, 0.2, 0.2, 1, -0.2, 0, + 0.2, -0.2, 0, -0.2, 1, 0.2, + 0, 0.2, -0.2, 0, 0.2, 1, + }, 6, 6) + vInv, err := Inv(v) + if err != nil { + t.Fatalf("Inv: %v", err) + } + dv, err := core.MatMul2D(d, vInv) + if err != nil { + t.Fatalf("MatMul2D: %v", err) + } + a, err := core.MatMul2D(v, dv) + if err != nil { + t.Fatalf("MatMul2D: %v", err) + } + values, vectors, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + want := []complex128{ + 5, + complex(0.5, 4), complex(0.5, -4), + -3, + complex(2, 1.5), complex(2, -1.5), + } + matchComplexSpectrum(t, values, want, 1e-8) + checkEigenPairs(t, a, values, vectors, 1e-9) + + sum := complex(0, 0) + prod := complex(1, 0) + for i := range 6 { + sum += values.ComplexAt(i) + prod *= values.ComplexAt(i) + } + if absComplex(sum-complex(7, 0)) > 1e-8 { + t.Fatalf("Σλ = %v, want 7", sum) + } + det, err := Det(a) + if err != nil { + t.Fatalf("Det: %v", err) + } + // det = 5·(−3)·(2²+1.5²)·(0.5²+4²) = −1523.4375. + if absComplex(prod-complex(det, 0)) > 1e-6*math.Abs(det) { + t.Fatalf("Πλ = %v, det = %g", prod, det) + } +} + +// TestEigenGeneralHermitian checks a complex input: [[2, i], [−i, 2]] +// has eigenvalues 1 and 3 (trace 4, determinant 3). +func TestEigenGeneralHermitian(t *testing.T) { + a := mustComplexes(t, []complex128{ + 2, complex(0, 1), + complex(0, -1), 2, + }, 2, 2) + values, vectors, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + want := []complex128{3, 1} + for i := range 2 { + if absComplex(values.ComplexAt(i)-want[i]) > 1e-12 { + t.Fatalf("value[%d] = %v, want %v", i, values.ComplexAt(i), want[i]) + } + } + checkEigenPairs(t, a, values, vectors, 1e-14) +} + +// TestEigenGeneralSymmetricCrossCheck pins the general path against +// the dedicated symmetric solver on the same matrix. +func TestEigenGeneralSymmetricCrossCheck(t *testing.T) { + vals := []float64{ + 4, 1, 0, + 1, 3, 2, + 0, 2, 5, + } + a := mustFloats(t, vals, 3, 3) + gen, _, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + sym, _, err := Eigen(a) + if err != nil { + t.Fatalf("Eigen: %v", err) + } + for i := range 3 { + best := math.MaxFloat64 + for k := range 3 { + if d := absComplex(gen.ComplexAt(i) - complex(sym.FloatAt(k), 0)); d < best { + best = d + } + } + if best > 1e-12 { + t.Fatalf("general value[%d] = %v has no symmetric match within 1e-12", i, gen.ComplexAt(i)) + } + } +} + +// TestEigenGeneralCyclicPermutation checks the matrix that made the +// shifted QR iteration famous: the cyclic permutation, where the +// spectrum is the full set of n-th roots of unity and the plain +// Wilkinson shift cycles until the exceptional shift breaks the +// symmetry. +func TestEigenGeneralCyclicPermutation(t *testing.T) { + const n = 6 + vals := make([]float64, n*n) + for i := range n { + vals[i*n+(i+1)%n] = 1 + } + a := mustFloats(t, vals, n, n) + values, vectors, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + want := make([]complex128, n) + for k := range n { + want[k] = cmplx.Exp(complex(0, 2*math.Pi*float64(k)/float64(n))) + } + matchComplexSpectrum(t, values, want, 1e-8) + checkEigenPairs(t, a, values, vectors, 1e-9) +} + +// TestEigenGeneralSingleCell covers the trivial 1×1 case. +func TestEigenGeneralSingleCell(t *testing.T) { + a := mustFloats(t, []float64{2.5}, 1, 1) + values, vectors, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + if absComplex(values.ComplexAt(0)-complex(2.5, 0)) > 1e-15 { + t.Fatalf("value = %v, want 2.5", values.ComplexAt(0)) + } + if absComplex(vectors.ComplexAt(0)-complex(1, 0)) > 1e-15 { + t.Fatalf("vector = %v, want 1", vectors.ComplexAt(0)) + } +} + +// TestEigenGeneralTinyRotation pins the purely relative deflation +// floors: a rotation scaled by 1e-20 keeps its complex conjugate +// spectrum instead of deflating to the real diagonal (the old +// max(1, scale) floor treated the whole matrix as rounding noise). +func TestEigenGeneralTinyRotation(t *testing.T) { + const scale = 1e-20 + theta := math.Pi / 5 + c, s := math.Cos(theta), math.Sin(theta) + a := mustFloats(t, []float64{ + scale * c, -scale * s, + scale * s, scale * c, + }, 2, 2) + values, vectors, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + want := []complex128{ + scale * cmplx.Exp(complex(0, theta)), + scale * cmplx.Exp(complex(0, -theta)), + } + for j := range 2 { + got := values.ComplexAt(j) + best := math.Inf(1) + for _, w := range want { + if m := absComplex(got - w); m < best { + best = m + } + } + if best > 1e-26 { + t.Fatalf("value[%d] = %v has no expected match within %g", j, got, 1e-26) + } + } + checkEigenPairs(t, a, values, vectors, 1e-25) +} + +func TestEigenGeneralErrors(t *testing.T) { + if _, _, err := EigenGeneral(mustFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)); err == nil { + t.Fatal("non-square matrix: want an error") + } + if _, _, err := EigenGeneral(mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2)); err == nil { + t.Fatal("3-D input: want an error") + } +} diff --git a/linalg/example_test.go b/linalg/example_test.go new file mode 100644 index 0000000..da8bae4 --- /dev/null +++ b/linalg/example_test.go @@ -0,0 +1,203 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg_test + +// Runnable examples for the flagship workflows of the package. `go +// test` executes them against the Output comments below, so the +// documentation cannot drift from the code. + +import ( + "fmt" + "log" + "math" + + "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/linalg" +) + +// A dense solve, and the same matrix factored once for reuse. +func ExampleSolve() { + a, err := tensor.FromFloats([]float64{4, 1, 1, 3}, 2, 2) + if err != nil { + log.Fatal(err) + } + b, err := tensor.FromFloats([]float64{1, 2}, 2) + if err != nil { + log.Fatal(err) + } + x, err := linalg.Solve(a, b) + if err != nil { + log.Fatal(err) + } + fmt.Printf("x = %.4f %.4f\n", x.FloatAt(0), x.FloatAt(1)) + + // The same matrix through Cholesky: a = L·Lᵀ, one factor for any + // number of right-hand sides. + l, err := linalg.Cholesky(a) + if err != nil { + log.Fatal(err) + } + fmt.Printf("L = [%.1f %.1f; %.1f %.4f]\n", + l.FloatAt(0), l.FloatAt(1), l.FloatAt(2), l.FloatAt(3)) + // Output: + // x = 0.0909 0.6364 + // L = [2.0 0.0; 0.5 1.6583] +} + +// A symmetric eigenproblem: ascending eigenvalues and orthonormal +// eigenvector columns. +func ExampleEigen() { + a, err := tensor.FromFloats([]float64{1, 2, 2, 1}, 2, 2) + if err != nil { + log.Fatal(err) + } + values, vectors, err := linalg.Eigen(a) + if err != nil { + log.Fatal(err) + } + fmt.Printf("values = %.4f %.4f\n", values.FloatAt(0), values.FloatAt(1)) + + // The sign of an eigenvector is arbitrary, so the magnitudes are + // what a stable example can print. Here |V[i][j]| = 1/√2. + fmt.Printf("|V| = [%.4f %.4f; %.4f %.4f]\n", + math.Abs(vectors.FloatAt(0)), math.Abs(vectors.FloatAt(1)), + math.Abs(vectors.FloatAt(2)), math.Abs(vectors.FloatAt(3))) + // Output: + // values = -1.0000 3.0000 + // |V| = [0.7071 0.7071; 0.7071 0.7071] +} + +// A matrix function: the exponential of a rotation generator is the +// rotation itself. +func ExampleMatrixExp() { + // A = [[0, -1], [1, 0]] generates a quarter-turn rotation rate, so + // exp(A) = [[cos 1, -sin 1], [sin 1, cos 1]]. + a, err := tensor.FromFloats([]float64{0, -1, 1, 0}, 2, 2) + if err != nil { + log.Fatal(err) + } + e, err := linalg.MatrixExp(a) + if err != nil { + log.Fatal(err) + } + fmt.Printf("%.4f %.4f %.4f %.4f\n", + e.FloatAt(0), e.FloatAt(1), e.FloatAt(2), e.FloatAt(3)) + // Output: 0.5403 -0.8415 0.8415 0.5403 +} + +// A sparse symmetric positive definite solve by conjugate gradient. +func ExampleSpSolve() { + a := laplacian1D() + // The right-hand side A·1, so the exact solution is the vector of + // ones and the printed digits show the residual the iteration left. + b, err := tensor.FromFloats([]float64{1, 0, 0, 1}, 4) + if err != nil { + log.Fatal(err) + } + x, err := linalg.SpSolve(a, b, 0, 0) + if err != nil { + log.Fatal(err) + } + fmt.Printf("%.4f %.4f %.4f %.4f\n", + x.FloatAt(0), x.FloatAt(1), x.FloatAt(2), x.FloatAt(3)) + // Output: 1.0000 1.0000 1.0000 1.0000 +} + +// A sparse direct factorisation under a fill-reducing ordering. +func ExampleNewSparseCholesky() { + a := laplacian1D() + b, err := tensor.FromFloats([]float64{1, 0, 0, 1}, 4) + if err != nil { + log.Fatal(err) + } + f, err := linalg.NewSparseCholesky(a, linalg.SparseOrderingReverseCuthillMcKee) + if err != nil { + log.Fatal(err) + } + x, err := f.Solve(b) + if err != nil { + log.Fatal(err) + } + fmt.Printf("%.4f %.4f %.4f %.4f\n", + x.FloatAt(0), x.FloatAt(1), x.FloatAt(2), x.FloatAt(3)) + + // The factor of a tridiagonal matrix keeps the band: 4 diagonal + // entries and 3 subdiagonal ones. + fmt.Println("factor non-zeros:", f.NNZ()) + // Output: + // 1.0000 1.0000 1.0000 1.0000 + // factor non-zeros: 7 +} + +// A natural cubic spline through samples of a straight line. +func ExampleNewCubicSpline() { + xs, err := tensor.FromFloats([]float64{0, 1, 2}, 3) + if err != nil { + log.Fatal(err) + } + ys, err := tensor.FromFloats([]float64{1, 3, 5}, 3) + if err != nil { + log.Fatal(err) + } + s, err := linalg.NewCubicSpline(xs, ys) + if err != nil { + log.Fatal(err) + } + fmt.Printf("%.4f %.4f\n", s.At(0.5), s.At(1.5)) + + // Extrapolation has no boundary condition here, so it is undefined. + fmt.Printf("%.4f\n", s.At(2.5)) + // Output: + // 2.0000 4.0000 + // NaN +} + +// A sequence of element-wise steps read as one expression. +func ExamplePipe() { + a, err := tensor.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + log.Fatal(err) + } + out, err := linalg.Pipe(a).AddF(1).Sqrt().MulF(2).Result() + if err != nil { + log.Fatal(err) + } + fmt.Printf("%.4f %.4f %.4f\n", out.FloatAt(0), out.FloatAt(1), out.FloatAt(2)) + + // A step that fails is recorded, and the steps after it are no-ops: + // adding a length-2 array to a length-3 one is a shape error. + short, err := tensor.FromFloats([]float64{1, 2}, 2) + if err != nil { + log.Fatal(err) + } + _, err = linalg.Pipe(a).Add(short).MulF(2).Result() + fmt.Println("error:", err != nil) + // Output: + // 2.8284 3.4641 4.0000 + // error: true +} + +// laplacian1D builds the 4×4 second-difference matrix with -1 on both +// off-diagonals and 2 on the diagonal, as a coordinate matrix. +func laplacian1D() *tensor.SparseCOO { + indices, err := tensor.FromInts([]int64{ + 0, 0, 1, 0, 0, 1, 1, 1, + 2, 1, 1, 2, 2, 2, 3, 2, + 2, 3, 3, 3, + }, 10, 2) + if err != nil { + panic(err) + } + values, err := tensor.FromFloats([]float64{ + 2, -1, -1, 2, -1, -1, 2, -1, -1, 2, + }, 10) + if err != nil { + panic(err) + } + a, err := tensor.NewSparseCOO(indices, values, []int{4, 4}) + if err != nil { + panic(err) + } + return a +} diff --git a/linalg/expm.go b/linalg/expm.go new file mode 100644 index 0000000..7df3703 --- /dev/null +++ b/linalg/expm.go @@ -0,0 +1,328 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// Matrix exponential, evaluated by scaling and squaring over a +// diagonal Pade approximant (Higham, 2005). That is the construction +// LAPACK's *ex* routines and every mainstream runtime use, because it +// is backward stable and needs nothing but matrix products and one LU +// solve, both of which the library already has. +// +// The kernel is generic over the same scalar family the LU kernel runs +// on, so the complex path stays in complex128 end to end instead of +// splitting off into a parallel implementation. + +// expmTheta holds the 1-norm bound below which each Pade degree keeps +// its approximation error inside double precision (Higham 2005, +// Table 10.1). +var expmTheta = map[int]float64{ + 3: 1.495585217958292e-2, + 5: 2.539398330063230e-1, + 7: 9.504178996162932e-1, + 9: 2.097847961257068e+0, + 13: 5.371920351148152e+0, +} + +// expmCoeffs returns the numerator coefficients b_j of the degree-m +// diagonal Pade approximant, as the integers of Higham's Table 10.1. +// The common factor across the row cancels in (V-U)^{-1}(V+U), so the +// unnormalised values are used directly, exactly as the reference +// implementations do. +func expmCoeffs(deg int) []float64 { + switch deg { + case 3: + return []float64{120, 60, 12, 1} + case 5: + return []float64{30240, 15120, 3360, 420, 30, 1} + case 7: + return []float64{17297280, 8648640, 1995840, 277200, 25200, 1512, 56, 1} + case 9: + return []float64{17643225600, 8821612800, 2075673600, 302702400, + 30270240, 2162160, 110880, 3960, 90, 1} + default: + return []float64{64764752532480000, 32382376266240000, 7771770303897600, + 1187353796428800, 129060195264000, 10559470521600, 670442572800, + 33522128640, 1323241920, 40840800, 960960, 16380, 182, 1} + } +} + +// padeDegree picks the approximant degree and the number of squarings +// for a given 1-norm: the smallest degree whose bound covers the norm, +// else degree 13 with the matrix scaled down by a power of two. +func padeDegree(nA float64) (deg, s int) { + for _, d := range []int{3, 5, 7, 9, 13} { + if nA <= expmTheta[d] { + return d, 0 + } + } + s = max(int(math.Ceil(math.Log2(nA/expmTheta[13]))), 0) + return 13, s +} + +// MatrixExp returns the matrix exponential exp(A) of a square matrix. +// Ints and float32 promote to float64; a complex matrix answers in +// complex128. A zero matrix gives the identity, and exp(A) for +// diagonal A is the exponential of the diagonal. +func MatrixExp(a *core.Array) (*core.Array, error) { + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, base.Errf("MatrixExp: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + if n == 0 { + return nil, base.Errf("MatrixExp: zero-sized matrix, got shape %s", base.ShapeText(a.Shape())) + } + if a.Dtype() == core.Complex { + m, err := squareComplexMatrix(a, "MatrixExp") + if err != nil { + return nil, err + } + r, err := expmKernel(m) + if err != nil { + return nil, err + } + flat := make([]complex128, n*n) + for i := range n { + copy(flat[i*n:(i+1)*n], r[i]) + } + return core.FromComplexes(flat, n, n) + } + m, err := squareFloatMatrix(a, "MatrixExp") + if err != nil { + return nil, err + } + r, err := expmKernel(m) + if err != nil { + return nil, err + } + flat := make([]float64, n*n) + for i := range n { + copy(flat[i*n:(i+1)*n], r[i]) + } + return floatsToArray(flat, []int{n, n}), nil +} + +// expmKernel computes exp(a) by scaling and squaring over a Pade +// approximant. The approximant is r(A) = (V-U)^{-1}(V+U) where U holds +// the odd part and V the even part of the numerator polynomial, both +// evaluated by Horner in A^2; s squarings then undo the scaling. +func expmKernel[T scalar](a [][]T) ([][]T, error) { + n := len(a) + nA := matrixNorm1(a) + if math.IsNaN(nA) || math.IsInf(nA, 0) { + // The theta ladder divides by the norm and int(Ceil(...)) of a + // non-finite value is implementation-defined; the Padé solve on + // Inf data would return an all-NaN matrix with a nil error. + return nil, base.Errf("MatrixExp: the matrix holds a non-finite entry (norm %g)", nA) + } + deg, s := padeDegree(nA) + scaled := a + if s > 0 { + scaled = matScaleT(a, math.Ldexp(1, -s)) + } + b := expmCoeffs(deg) + a2 := matMulT(scaled, scaled) + // The Horner sweeps alternate between two buffers instead of landing + // a fresh matrix per coefficient: the product reads both operands + // and writes the buffer the previous step left free, so the entries + // the next multiply reads are the ones this one wrote. + v, vFree := eyeScaled[T](n, b[deg-1]), newSquare[T](n) + uIn, uFree := eyeScaled[T](n, b[deg]), newSquare[T](n) + for j := deg - 3; j >= 0; j -= 2 { + matMulAddEyeInto(vFree, a2, v, b[j]) + v, vFree = vFree, v + } + for j := deg - 2; j >= 1; j -= 2 { + matMulAddEyeInto(uFree, a2, uIn, b[j]) + uIn, uFree = uFree, uIn + } + u := matMulT(scaled, uIn) + r, err := solveRight("MatrixExp", matSubT(v, u), matAddT(v, u)) + if err != nil { + return nil, err + } + if s > 0 { + sq := newSquare[T](n) + for range s { + matMulTInto(sq, r, r) + r, sq = sq, r + } + } + return r, nil +} + +// newSquare allocates an n×n matrix in the kernel's element type as n +// row views over one flat backing slice: two allocations instead of +// n+1. The rows stay disjoint, so in-place row swaps remain valid. +func newSquare[T scalar](n int) [][]T { + back := make([]T, n*n) + rows := make([][]T, n) + for i := range n { + rows[i] = back[i*n : (i+1)*n] + } + return rows +} + +// realT widens a float64 into the kernel's element type, giving a zero +// imaginary part when the kernel runs on complex128. +func realT[T scalar](f float64) T { + var zero T + if _, ok := any(zero).(complex128); ok { + return any(complex(f, 0)).(T) + } + return any(f).(T) +} + +// eyeScaled returns c times the n×n identity. +func eyeScaled[T scalar](n int, c float64) [][]T { + out := newSquare[T](n) + cc := realT[T](c) + for i := range n { + out[i][i] = cc + } + return out +} + +// matMulAddEyeInto writes a·m + c·I into out, which must be a square +// matrix of the same order aliasing neither operand. +func matMulAddEyeInto[T scalar](out, a, m [][]T, c float64) { + matMulTInto(out, a, m) + cc := realT[T](c) + for i := range out { + out[i][i] += cc + } +} + +// matMulT multiplies two square matrices in the kernel's element type +// and returns a fresh matrix. +func matMulT[T scalar](a, b [][]T) [][]T { + out := newSquare[T](len(a)) + matMulTInto(out, a, b) + return out +} + +// matMulTInto writes the product a·b into out, which must alias neither +// operand. out arrives holding an earlier product, so it is cleared +// first: a fresh matrix is zero by construction and the accumulation +// below depends on that. The loop order keeps the inner product over p +// contiguous in both operands, matching the dense kernel's cache +// behaviour, and the zero multiplier is skipped the way the product has +// always skipped it, so a zero times a non-finite entry contributes +// nothing rather than a NaN. +func matMulTInto[T scalar](out, a, b [][]T) { + n := len(a) + var zero T + for i := range n { + oi := out[i] + clear(oi) + ai := a[i] + for p := range n { + aip := ai[p] + if aip == zero { + continue + } + bp := b[p] + for j := range n { + oi[j] += aip * bp[j] + } + } + } +} + +// matAddT returns the element-wise sum of two square matrices. +func matAddT[T scalar](a, b [][]T) [][]T { + n := len(a) + out := newSquare[T](n) + for i := range n { + oi, ai, bi := out[i], a[i], b[i] + for j := range n { + oi[j] = ai[j] + bi[j] + } + } + return out +} + +// matSubT returns the element-wise difference of two square matrices. +func matSubT[T scalar](a, b [][]T) [][]T { + n := len(a) + out := newSquare[T](n) + for i := range n { + oi, ai, bi := out[i], a[i], b[i] + for j := range n { + oi[j] = ai[j] - bi[j] + } + } + return out +} + +// matScaleT multiplies every element by the real scalar f. +func matScaleT[T scalar](a [][]T, f float64) [][]T { + n := len(a) + fT := realT[T](f) + out := newSquare[T](n) + for i := range n { + oi, ai := out[i], a[i] + for j := range n { + oi[j] = ai[j] * fT + } + } + return out +} + +// matrixNorm1 returns the maximum absolute column sum, the norm the +// scaling step is measured against. +func matrixNorm1[T scalar](m [][]T) float64 { + n := len(m) + maxCol := 0.0 + for j := range n { + sum := 0.0 + for i := range n { + sum += absOf(m[i][j]) + } + // A NaN column sum can never win the strict max below, so it + // propagates explicitly: the caller reads the norm as the + // finiteness screen. + if math.IsNaN(sum) { + return math.NaN() + } + if sum > maxCol { + maxCol = sum + } + } + return maxCol +} + +// solveRight solves a·x = b for row-major square matrices and returns +// x row-major. a is consumed by the factorisation, so callers hand +// over a freshly built matrix; the column-oriented solveSystem is +// reused underneath rather than a second LU being written. +func solveRight[T scalar](name string, a, b [][]T) ([][]T, error) { + n := len(a) + colBack := make([]T, n*n) + cols := make([][]T, n) + for j := range n { + col := colBack[j*n : (j+1)*n] + for i := range n { + col[i] = b[i][j] + } + cols[j] = col + } + sol, err := base.SolveSystem(name, a, cols) + if err != nil { + return nil, err + } + x := newSquare[T](n) + for i := range n { + for j := range n { + x[i][j] = sol[j][i] + } + } + return x, nil +} diff --git a/linalg/expm_test.go b/linalg/expm_test.go new file mode 100644 index 0000000..3ed67e5 --- /dev/null +++ b/linalg/expm_test.go @@ -0,0 +1,323 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// approxEqual reports whether two real arrays agree elementwise within +// an absolute and a relative tolerance. +func approxEqual(t *testing.T, got, want *core.Array, atol, rtol float64, label string) { + t.Helper() + if got.NDim() != want.NDim() { + t.Fatalf("%s: rank %d, want %d", label, got.NDim(), want.NDim()) + } + for d := range got.NDim() { + if got.Shape()[d] != want.Shape()[d] { + t.Fatalf("%s: shape %s, want %s", label, base.ShapeText(got.Shape()), base.ShapeText(want.Shape())) + } + } + for i := range got.Len() { + g, w := got.FloatAt(i), want.FloatAt(i) + diff := math.Abs(g - w) + if diff > atol+rtol*math.Abs(w) { + t.Fatalf("%s: element %d = %.12g, want %.12g (|diff| %.3g)", + label, i, g, w, diff) + } + } +} + +// TestMatrixExpClosedForms checks exp(A) against forms with a known +// analytic answer: the identity, a nilpotent matrix whose series +// terminates, a diagonal matrix, and the rotation generator whose +// exponential is a rotation. +func TestMatrixExpClosedForms(t *testing.T) { + cases := []struct { + name string + vals []float64 + n int + want []float64 + }{ + { + name: "zero_gives_identity", + vals: []float64{0, 0, 0, 0}, + n: 2, + want: []float64{1, 0, 0, 1}, + }, + { + name: "scalar_multiple_of_identity", + vals: []float64{2, 0, 0, 2}, + n: 2, + want: []float64{math.E * math.E, 0, 0, math.E * math.E}, + }, + { + name: "nilpotent_series_terminates", + // A = [[0,1],[0,0]] squares to zero, so exp(A) = I + A. + vals: []float64{0, 1, 0, 0}, + n: 2, + want: []float64{1, 1, 0, 1}, + }, + { + name: "diagonal", + vals: []float64{1, 0, 0, 0, 2, 0, 0, 0, 3}, + n: 3, + want: []float64{math.Exp(1), 0, 0, 0, math.Exp(2), 0, 0, 0, math.Exp(3)}, + }, + { + name: "jordan_block", + // A = [[1,1],[0,1]] gives exp(A) = e·[[1,1],[0,1]]. + vals: []float64{1, 1, 0, 1}, + n: 2, + want: []float64{math.E, math.E, 0, math.E}, + }, + { + name: "rotation_generator", + // exp([[0,-t],[t,0]]) = [[cos t, -sin t],[sin t, cos t]]. + vals: []float64{0, -0.7, 0.7, 0}, + n: 2, + want: []float64{math.Cos(0.7), -math.Sin(0.7), math.Sin(0.7), math.Cos(0.7)}, + }, + { + name: "large_norm_needs_squarings", + // A norm above theta_13 forces the scaling path. + vals: []float64{30, 0, 0, -30}, + n: 2, + want: []float64{math.Exp(30), 0, 0, math.Exp(-30)}, + }, + } + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + a, err := core.FromFloats(tt.vals, tt.n, tt.n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := MatrixExp(a) + if err != nil { + t.Fatalf("MatrixExp: %v", err) + } + want, err := core.FromFloats(tt.want, tt.n, tt.n) + if err != nil { + t.Fatalf("FromFloats want: %v", err) + } + approxEqual(t, got, want, 1e-12, 1e-10, tt.name) + }) + } +} + +// TestMatrixExpSeries cross-checks exp(A) against the Taylor series +// sum A^k/k! on a matrix small enough for the series to converge +// tightly, covering a non-normal input the closed forms do not. +func TestMatrixExpSeries(t *testing.T) { + vals := []float64{ + 0.5, 0.3, -0.1, + 0.2, -0.4, 0.6, + 0.1, 0.7, 0.2, + } + const n = 3 + a, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := MatrixExp(a) + if err != nil { + t.Fatalf("MatrixExp: %v", err) + } + + // Series by repeated multiplication with the running term. + term := make([]float64, n*n) + for i := range n { + term[i*n+i] = 1 + } + sum := make([]float64, n*n) + copy(sum, term) + am := make([]float64, n*n) + copy(am, vals) + for k := 1; k <= 40; k++ { + // term = term · A / k + next := make([]float64, n*n) + for i := range n { + for p := range n { + for j := range n { + next[i*n+j] += term[i*n+p] * am[p*n+j] + } + } + } + for i := range n * n { + term[i] = next[i] / float64(k) + sum[i] += term[i] + } + } + want := floatsToArray(sum, []int{n, n}) + approxEqual(t, got, want, 1e-13, 1e-12, "series") +} + +// TestMatrixExpGroupProperty checks the defining identity +// exp(A)·exp(-A) = I, which exercises both the solve and the squaring +// paths on a matrix with mixed-sign, non-symmetric entries. +func TestMatrixExpGroupProperty(t *testing.T) { + vals := []float64{ + 1.2, -0.7, 0.4, + 0.9, 0.3, -1.1, + -0.5, 0.8, 0.6, + } + const n = 3 + a, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + neg := core.MulF(a, -1) + ea, err := MatrixExp(a) + if err != nil { + t.Fatalf("MatrixExp(a): %v", err) + } + ena, err := MatrixExp(neg) + if err != nil { + t.Fatalf("MatrixExp(-a): %v", err) + } + prod, err := core.MatMul2D(ea, ena) + if err != nil { + t.Fatalf("MatMul2D: %v", err) + } + want, err := core.Zeros(core.Float, n, n) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + for i := range n { + want.RawFloats()[i*n+i] = 1 + } + approxEqual(t, prod, want, 1e-10, 1e-9, "exp(A)exp(-A)") +} + +// TestMatrixExpComplexHermitian checks the complex path against a +// closed form: a 2×2 Hermitian A = m·I + B splits into a scalar part +// and a traceless B with B² = r²·I, so exp(A) = e^m·(cosh r·I + +// (sinh r)/r·B). The real-valued tests cannot reach this path. +func TestMatrixExpComplexHermitian(t *testing.T) { + // A = [[2, 1-i],[1+i, 3]]. + const ( + a11 = 2.0 + a22 = 3.0 + n = 2 + ) + bOff := complex(1, 1) + in := []complex128{ + a11, cmplx.Conj(bOff), + bOff, a22, + } + a, err := core.FromComplexes(in, n, n) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + got, err := MatrixExp(a) + if err != nil { + t.Fatalf("MatrixExp: %v", err) + } + + mid := (a11 + a22) / 2 + r := math.Hypot(a11-mid, cmplx.Abs(bOff)) + em := math.Exp(mid) + ch := math.Cosh(r) + sh := math.Sinh(r) / r + want := make([]complex128, n*n) + for i := range n { + for j := range n { + bij := a.ComplexAt(i*n + j) + if i == j { + bij -= complex(mid, 0) + } + idn := complex(0, 0) + if i == j { + idn = complex(1, 0) + } + want[i*n+j] = complex(em, 0) * (complex(ch, 0)*idn + complex(sh, 0)*bij) + } + } + + for i := range n * n { + g, w := got.ComplexAt(i), want[i] + if diff := cmplx.Abs(g - w); diff > 1e-10*(1+cmplx.Abs(w)) { + t.Fatalf("element %d = %v, want %v (|diff| %.3g)", i, g, w, diff) + } + } +} + +// TestMatrixExpRejectsInvalid pins the error contract: non-square and +// zero-sized inputs are refused rather than panicking. +func TestMatrixExpRejectsInvalid(t *testing.T) { + tall, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + if _, err := MatrixExp(tall); err == nil { + t.Fatal("expected an error for a non-square matrix") + } + empty, err := core.Zeros(core.Float, 0, 0) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + if _, err := MatrixExp(empty); err == nil { + t.Fatal("expected an error for a zero-sized matrix") + } +} + +// TestMatrixExpDtypes checks that int and float32 inputs promote to +// float64 results, matching the promotion the rest of the linalg +// surface applies. +func TestMatrixExpDtypes(t *testing.T) { + ints, err := core.FromInts([]int64{1, 0, 0, 1}, 2, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + gotInt, err := MatrixExp(ints) + if err != nil { + t.Fatalf("MatrixExp(int): %v", err) + } + if gotInt.Dtype() != core.Float { + t.Fatalf("int input gave dtype %s, want %s", gotInt.Dtype(), core.Float) + } + f32, err := core.FromFloat32s([]float32{0, 0, 0, 0}, 2, 2) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + gotF32, err := MatrixExp(f32) + if err != nil { + t.Fatalf("MatrixExp(float32): %v", err) + } + if gotF32.Dtype() != core.Float { + t.Fatalf("float32 input gave dtype %s, want %s", gotF32.Dtype(), core.Float) + } +} + +// TestPadeDegreeSelection pins the degree and squaring choice across +// the whole theta ladder, including the scaled branch. +func TestPadeDegreeSelection(t *testing.T) { + cases := []struct { + nA float64 + deg int + sq int + }{ + {0, 3, 0}, + {1e-3, 3, 0}, + {1e-1, 5, 0}, + {0.5, 7, 0}, + {1.0, 9, 0}, + {2.0, 9, 0}, + {2.097847961257068, 9, 0}, + {2.097847961257069, 13, 0}, + {5.371920351148152, 13, 0}, + {10.743840702296304, 13, 1}, + {1000, 13, 8}, + } + for _, tt := range cases { + deg, sq := padeDegree(tt.nA) + if deg != tt.deg || sq != tt.sq { + t.Fatalf("padeDegree(%g) = (%d, %d), want (%d, %d)", tt.nA, deg, sq, tt.deg, tt.sq) + } + } +} diff --git a/linalg/extreme_scale_pins_test.go b/linalg/extreme_scale_pins_test.go new file mode 100644 index 0000000..2bb6c61 --- /dev/null +++ b/linalg/extreme_scale_pins_test.go @@ -0,0 +1,1056 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "os" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression pins at the extremes of the float64 range: the dense and +// sparse decompositions, solvers and eigensolvers must keep their +// contracts at magnitudes from 1e-300 to 1e300, where a raw square +// overflows or vanishes. + +// scales sweeps the magnitudes the decompositions have to survive: +// ordinary 1 as the control that pins the untouched arithmetic, the two +// edges of the squared-arithmetic window (a raw square leaves the normal +// range at about 1.3e154 and 1.5e-162), and the extremes of the float64 +// range in both directions. +var scales = []float64{1, 1e150, 1e155, 1e200, 1e300, 1e-150, 1e-160, 1e-200, 1e-300} + +// scale multiplies every entry of a flat matrix by s. +func scale(vals []float64, s float64) []float64 { + out := make([]float64, len(vals)) + for i, v := range vals { + out[i] = v * s + } + return out +} + +// scaleC multiplies every complex entry by the real s. +func scaleC(vals []complex128, s float64) []complex128 { + out := make([]complex128, len(vals)) + for i, v := range vals { + out[i] = complex(real(v)*s, imag(v)*s) + } + return out +} + +// iJ is the 4x4 matrix I+J: 2 on the diagonal, 1 elsewhere. Its +// spectrum is closed form, 1 three times and 5 once, which is the +// reference every 1e+-200 eigen test below checks against. +func iJ(n int) []float64 { + out := make([]float64, n*n) + for i := range n { + for j := range n { + if i == j { + out[i*n+j] = 2 + } else { + out[i*n+j] = 1 + } + } + } + return out +} + +// pinDenseSolve solves a·x = b by Gaussian elimination with partial +// pivoting. It is the independent brute-force reference the sparse +// solvers are checked against, deliberately written without any of the +// library's own machinery. +func pinDenseSolve(t *testing.T, a, b []complex128, n int) []complex128 { + t.Helper() + m := append([]complex128(nil), a...) + x := append([]complex128(nil), b...) + for k := range n { + piv := k + for i := k + 1; i < n; i++ { + if cmplx.Abs(m[i*n+k]) > cmplx.Abs(m[piv*n+k]) { + piv = i + } + } + if m[piv*n+k] == 0 { + t.Fatalf("dense reference: singular matrix at column %d", k) + } + if piv != k { + for j := range n { + m[k*n+j], m[piv*n+j] = m[piv*n+j], m[k*n+j] + } + x[k], x[piv] = x[piv], x[k] + } + for i := k + 1; i < n; i++ { + f := m[i*n+k] / m[k*n+k] + for j := k; j < n; j++ { + m[i*n+j] -= f * m[k*n+j] + } + x[i] -= f * x[k] + } + } + for k := n - 1; k >= 0; k-- { + for j := k + 1; j < n; j++ { + x[k] -= m[k*n+j] * x[j] + } + x[k] /= m[k*n+k] + } + return x +} + +// hermitianStencil builds an n×n Hermitian tridiagonal stencil: +// 2 on the diagonal and conjugate mirrored imaginary couplings. It is +// the mirrored-stencil shape the complex sparse solvers are exercised on. +func hermitianStencil(n int) []complex128 { + a := make([]complex128, n*n) + for i := range n { + a[i*n+i] = 2 + if i+1 < n { + a[i*n+i+1] = complex(0, 0.5) + a[(i+1)*n+i] = complex(0, -0.5) + } + } + return a +} + +// cSRFromDense builds a complex SparseCOO from a dense flat matrix, +// dropping the exact zeros the way a caller's assembly would. +func cSRFromDense(t *testing.T, a []complex128, n int) *core.SparseCOO { + t.Helper() + var idx []int64 + var vals []complex128 + for i := range n { + for j := range n { + if a[i*n+j] == 0 { + continue + } + idx = append(idx, int64(i), int64(j)) + vals = append(vals, a[i*n+j]) + } + } + ind, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + val, err := core.FromComplexes(vals, len(vals)) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + coo, err := core.NewSparseCOO(ind, val, []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// TestHouseholderVectorReflectsAtExtremeScale pins the reflector itself +// (report T1/F7): for a non-zero x, beta must be positive and finite at +// every magnitude, not +0 (which callers read as "no reflection") or +// +Inf (whose product with a zero dot is NaN), and H must still map x to +// -sign(x0)·‖x‖·e1. The reflection is applied to x/s, so the check's own +// arithmetic stays in range at every scale. +func TestHouseholderVectorReflectsAtExtremeScale(t *testing.T) { + for _, s := range scales { + for _, x0 := range []float64{s, -s} { + x := []float64{x0, s, s} + dst := make([]float64, 3) + hh := householderVectorInto(dst, x) + if !(hh.beta > 0) || math.IsInf(hh.beta, 0) { + t.Fatalf("scale %g, x0 %g: beta = %v, want a positive finite reflector", s, x0, hh.beta) + } + for i, v := range hh.v { + if math.IsInf(v, 0) || math.IsNaN(v) { + t.Fatalf("scale %g, x0 %g: v[%d] = %v, want finite", s, x0, i, v) + } + } + xs := []float64{x0 / s, 1, 1} + dot := 0.0 + for i := range xs { + dot += hh.v[i] * xs[i] + } + w := hh.beta * dot + sign := 1.0 + if x0 < 0 { + sign = -1 + } + for i := range xs { + want := 0.0 + if i == 0 { + want = -sign * math.Sqrt(3) + } + if got := xs[i] - hh.v[i]*w; math.Abs(got-want) > 1e-13 { + t.Fatalf("scale %g, x0 %g: (H x)[%d] = %g, want %g", s, x0, i, got, want) + } + } + } + // The zero vector still answers the identity reflector: beta == 0 + // is the legitimate "nothing to reflect" signal and must stay + // distinguishable from the overflow the fix removes. + if hh := householderVectorInto(make([]float64, 3), []float64{0, 0, 0}); hh.beta != 0 { + t.Fatalf("scale %g: zero vector beta = %v, want 0", s, hh.beta) + } + } +} + +// TestEigenSpectrumAtExtremeScale pins the closed-form spectrum of +// (I+J)·s: 1,1,1,5 times s, with the eigenvectors orthonormal and the +// residual at rounding level. Before the fix, 1e155 returned the raw +// diagonal {2,2,2,2}e155 (the reflector's beta came out as +0 and was +// skipped) and 1e-200 returned all NaN (beta came out as +Inf). +func TestEigenSpectrumAtExtremeScale(t *testing.T) { + const n = 4 + want := []float64{1, 1, 1, 5} + for _, s := range scales { + a := mustFloats(t, scale(iJ(n), s), n, n) + vals, vecs, err := Eigen(a) + if err != nil { + t.Fatalf("Eigen at scale %g: %v", s, err) + } + for i := range n { + if got := vals.FloatAt(i) / s; math.Abs(got-want[i]) > 1e-12 { + t.Fatalf("scale %g: eigenvalue %d = %g, want %g", s, i, got, want[i]) + } + } + // Residual of the eigenpair on the O(1) shift: ‖(I+J)v - (λ/s)v‖, + // which is the residual of the original problem divided by s. + base := iJ(n) + for k := range n { + lam := vals.FloatAt(k) / s + worst := 0.0 + for i := range n { + acc := 0.0 + for j := range n { + acc += base[i*n+j] * vecs.FloatAt(j*n+k) + } + worst = math.Max(worst, math.Abs(acc-lam*vecs.FloatAt(i*n+k))) + } + if worst > 1e-12 { + t.Fatalf("scale %g: residual of eigenpair %d = %g, want <= 1e-12", s, k, worst) + } + } + // VᵀV = I: a scaling mistake that dropped the factor would still + // leave orthonormal columns, so this is a cheap second net. + for j := range n { + for k := range n { + acc := 0.0 + for i := range n { + acc += vecs.FloatAt(i*n+j) * vecs.FloatAt(i*n+k) + } + want := 0.0 + if j == k { + want = 1 + } + if math.Abs(acc-want) > 1e-12 { + t.Fatalf("scale %g: (VᵀV)[%d,%d] = %g, want %g", s, j, k, acc, want) + } + } + } + } +} + +// TestSVDReconstructsAtExtremeScale pins A = U·Σ·Vᵀ against the +// hand-built 2×2 ([[0,s],[s,0]] has both singular values s) and against +// a dense 4x4 with a closed-form scale-free reconstruction. Before the +// fix, 1e155 returned sigma = (+Inf, +Inf) and the dense case all NaN. +func TestSVDReconstructsAtExtremeScale(t *testing.T) { + for _, s := range scales { + a := mustFloats(t, []float64{0, s, s, 0}, 2, 2) + u, sigma, vt, err := SVD(a) + if err != nil { + t.Fatalf("SVD at scale %g: %v", s, err) + } + for i := range 2 { + if got := sigma.FloatAt(i) / s; math.Abs(got-1) > 1e-13 { + t.Fatalf("scale %g: sigma[%d] = %g, want 1 (in units of s)", s, i, got) + } + } + // [[0,1],[1,0]] = U·(Σ/s)·Vᵀ in scaled units. + for i := range 2 { + for j := range 2 { + acc := 0.0 + for k := range 2 { + acc += u.FloatAt(i*2+k) * (sigma.FloatAt(k) / s) * vt.FloatAt(k*2+j) + } + want := 0.0 + if i != j { + want = 1 + } + if math.Abs(acc-want) > 1e-13 { + t.Fatalf("scale %g: (UΣVᵀ)[%d,%d] = %g, want %g", s, i, j, acc, want) + } + } + } + } + // The dense case the report saw as all NaN at 1e155. + const n = 4 + base := []float64{ + 4, 1, 2, 0.5, + 1, 3, 0.25, 1, + 2, 0.25, 5, 1, + 0.5, 1, 1, 2, + } + for _, s := range []float64{1, 1e155, 1e-155, 1e200} { + a := mustFloats(t, scale(base, s), n, n) + u, sigma, vt, err := SVD(a) + if err != nil { + t.Fatalf("SVD dense at scale %g: %v", s, err) + } + worst := 0.0 + for i := range n { + for j := range n { + // A/s = U·(Σ/s)·Vᵀ. + acc := 0.0 + for k := range n { + acc += u.FloatAt(i*n+k) * (sigma.FloatAt(k) / s) * vt.FloatAt(k*n+j) + } + worst = math.Max(worst, math.Abs(acc-base[i*n+j])) + } + } + if worst > 1e-13 { + t.Fatalf("scale %g: dense reconstruction error %g, want <= 1e-13", s, worst) + } + } +} + +// TestEigenGeneralSpectrumAtExtremeScale pins EigenGeneral on (I+J)·s, +// where the spectrum is closed form and the matrix is far from +// defective. Before the fix, 1e-200 silently returned {2,2,2,2}e-200 and +// 1e160 failed with a spurious "QR iteration failed to converge". +func TestEigenGeneralSpectrumAtExtremeScale(t *testing.T) { + const n = 4 + want := []float64{1, 1, 1, 5} + for _, s := range scales { + a := mustFloats(t, scale(iJ(n), s), n, n) + vals, vecs, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral at scale %g: %v", s, err) + } + // Values come back descending by magnitude: 5 then 1,1,1. + got := make([]float64, n) + for i := range n { + got[i] = cmplx.Abs(vals.ComplexAt(i)) / s + } + if math.Abs(got[0]-5) > 1e-12 { + t.Fatalf("scale %g: |λ| max = %g, want 5", s, got[0]) + } + for i := 1; i < n; i++ { + if math.Abs(got[i]-1) > 1e-12 { + t.Fatalf("scale %g: |λ| %d = %g, want %g", s, i, got[i], want[i]) + } + } + // Residual of every eigenpair on the O(1) shift. + base := iJ(n) + for k := range n { + lam := vals.ComplexAt(k) / complex(s, 0) + worst := 0.0 + for i := range n { + acc := complex(0, 0) + for j := range n { + acc += complex(base[i*n+j], 0) * vecs.ComplexAt(j*n+k) + } + worst = math.Max(worst, cmplx.Abs(acc-lam*vecs.ComplexAt(i*n+k))) + } + if worst > 1e-12 { + t.Fatalf("scale %g: residual of eigenpair %d = %g, want <= 1e-12", s, k, worst) + } + } + } + // A mixed-scale matrix: a huge diagonal with couplings 1e-260 of it. + // The k = 1 reflector's column is then about 1e-200 while the matrix + // max is 1e60, so vMax*vMax underflows to zero, beta comes out as + // +Inf and the update turns the reduction into NaN, even though the + // matrix itself is inside the safe window. The spectrum is that of + // the huge diagonal to rounding. + const big, rel = 1e60, 1e-260 + mixed := []float64{ + big, 0, 0, 0, + 0, big, big * rel, big * rel, + 0, big * rel, big, big * rel, + 0, big * rel, big * rel, big, + } + valsM, _, err := EigenGeneral(mustFloats(t, mixed, n, n)) + if err != nil { + t.Fatalf("EigenGeneral of the mixed-scale matrix: %v", err) + } + for i := range n { + got := valsM.ComplexAt(i) + if math.IsNaN(real(got)) || math.IsNaN(imag(got)) { + t.Fatalf("mixed-scale eigenvalue %d = %v, want a finite value", i, got) + } + if scale := cmplx.Abs(got) / big; math.Abs(scale-1) > 1e-12 { + t.Fatalf("mixed-scale eigenvalue %d = %v, want magnitude %g", i, got, big) + } + } + // The reflector-side face of the same defect, reached with an + // ordinary matrix: a unit diagonal with a column at 1e-160. That + // column is above the magnitude floor, so the reflector is built, but + // its raw squares are subnormal, beta overflows to +Inf and the whole + // reduction comes back NaN where the matrix is perfectly valid. Only + // the reflector's rescaling survives it, so this case pins that + // mechanism on its own. + const tiny = 1e-160 + near := []float64{ + 1, 0, 0, 0, + 0, 1, tiny, tiny, + 0, tiny, 1, tiny, + 0, tiny, tiny, 1, + } + valsN, _, err := EigenGeneral(mustFloats(t, near, n, n)) + if err != nil { + t.Fatalf("EigenGeneral of the subnormal-column matrix: %v", err) + } + for i := range n { + got := valsN.ComplexAt(i) + if math.IsNaN(real(got)) || math.IsNaN(imag(got)) { + t.Fatalf("subnormal-column eigenvalue %d = %v, want a finite value", i, got) + } + if math.Abs(cmplx.Abs(got)-1) > 1e-12 { + t.Fatalf("subnormal-column eigenvalue %d = %v, want magnitude 1", i, got) + } + } +} + +// TestSchurComplexAndMatrixSqrtAtExtremeScale pins the Schur contract +// A = Q·T·Qᴴ with T upper triangular and the diagonal of T the closed-form +// circulant spectrum, plus the defining property of the principal square +// root, r·r = A. Before the fix, SchurComplex, MatrixSqrt and MatrixLog +// all failed with "the Schur iteration did not converge" at 1e-200 and +// 1e160 on perfectly valid input. +func TestSchurComplexAndMatrixSqrtAtExtremeScale(t *testing.T) { + const n = 4 + base := []float64{ + 4, 1, 2, 3, + 3, 4, 1, 2, + 2, 3, 4, 1, + 1, 2, 3, 4, + } + // Eigenvalues of the symmetric circulant with first row [4,1,2,3]: + // 10, 2, 2+2i, 2-2i, computed from the closed form Σ c_j ω^{jk}. + want := []complex128{10, 2, complex(2, 2), complex(2, -2)} + for _, s := range scales { + a := mustFloats(t, scale(base, s), n, n) + tm, q, err := SchurComplex(a) + if err != nil { + t.Fatalf("SchurComplex at scale %g: %v", s, err) + } + // A/s = Q·(T/s)·Qᴴ. + worst := 0.0 + for i := range n { + for j := range n { + acc := complex(0, 0) + for k := range n { + for l := range n { + acc += q.ComplexAt(i*n+k) * (tm.ComplexAt(k*n+l) / complex(s, 0)) * cmplx.Conj(q.ComplexAt(j*n+l)) + } + } + worst = math.Max(worst, cmplx.Abs(acc-complex(base[i*n+j], 0))) + } + } + if worst > 1e-13 { + t.Fatalf("scale %g: Schur reconstruction error %g, want <= 1e-13", s, worst) + } + // T upper triangular, and its diagonal the closed-form spectrum. + for i := 1; i < n; i++ { + for j := 0; j < i; j++ { + if v := cmplx.Abs(tm.ComplexAt(i*n+j)) / s; v > 1e-13 { + t.Fatalf("scale %g: T[%d,%d] = %g, want upper triangular", s, i, j, v) + } + } + } + got := make([]complex128, n) + for i := range n { + got[i] = tm.ComplexAt(i*n+i) / complex(s, 0) + } + if !spectrumMatches(got, want, 1e-12) { + t.Fatalf("scale %g: Schur diagonal %v, want %v", s, got, want) + } + // The principal square root of [[7,10],[15,22]]·s satisfies + // r·r = A; the closed form of the O(1) matrix is [[1,2],[3,4]]², + // so the check is (r/√s)² = [[7,10],[15,22]]. + ms, err := MatrixSqrt(mustFloats(t, scale([]float64{7, 10, 15, 22}, s), 2, 2)) + if err != nil { + t.Fatalf("MatrixSqrt at scale %g: %v", s, err) + } + r := math.Sqrt(s) + rr := []float64{ + ms.FloatAt(0)*ms.FloatAt(0) + ms.FloatAt(1)*ms.FloatAt(2), + ms.FloatAt(0)*ms.FloatAt(1) + ms.FloatAt(1)*ms.FloatAt(3), + ms.FloatAt(2)*ms.FloatAt(0) + ms.FloatAt(3)*ms.FloatAt(2), + ms.FloatAt(2)*ms.FloatAt(1) + ms.FloatAt(3)*ms.FloatAt(3), + } + sq := []float64{7, 10, 15, 22} + for i := range 4 { + if got := rr[i] / (r * r); math.Abs(got-sq[i]) > 1e-12 { + t.Fatalf("scale %g: (r·r)[%d] = %g, want %g", s, i, got, sq[i]) + } + } + } +} + +// spectrumMatches reports whether every want finds a distinct got +// within tol: a spectrum is a set, and the order of near-equal complex +// values is not part of the contract. +func spectrumMatches(got, want []complex128, tol float64) bool { + used := make([]bool, len(got)) + for _, w := range want { + found := false + for i, g := range got { + if !used[i] && cmplx.Abs(g-w) <= tol { + used[i] = true + found = true + break + } + } + if !found { + return false + } + } + return true +} + +// TestSVDComplexReconstructsAtExtremeScale pins A = U·Σ·Vᴴ on the +// hand-built [[0,t],[t,0]], whose singular values are both t, and on +// diag(1e160, 1e-160). Before the fix, t = 1e-200 silently returned +// sigma = (0,0) (the reflector norm underflowed to zero and the +// below-diagonal mass was zeroed with it) and t = 1e160 returned NaN. +func TestSVDComplexReconstructsAtExtremeScale(t *testing.T) { + for _, t0 := range scales { + a := mustFromComplexes(t, []complex128{ + 0, complex(t0, 0), + complex(t0, 0), 0, + }, 2, 2) + u, sigma, vh, err := SVDComplex(a) + if err != nil { + t.Fatalf("SVDComplex at scale %g: %v", t0, err) + } + for i := range 2 { + if got := sigma.FloatAt(i) / t0; math.Abs(got-1) > 1e-13 { + t.Fatalf("scale %g: sigma[%d] = %g, want 1 (in units of t)", t0, i, got) + } + } + for i := range 2 { + for j := range 2 { + acc := complex(0, 0) + for k := range 2 { + acc += u.ComplexAt(i*2+k) * complex(sigma.FloatAt(k)/t0, 0) * vh.ComplexAt(k*2+j) + } + want := complex(0, 0) + if i != j { + want = 1 + } + if d := cmplx.Abs(acc - want); d > 1e-13 { + t.Fatalf("scale %g: (UΣVᴴ)[%d,%d] error %g, want 0", t0, i, j, d) + } + } + } + } + // The mixed-spectrum case the report saw as NaN, NaN. + a := mustFromComplexes(t, []complex128{complex(1e160, 0), 0, 0, complex(1e-160, 0)}, 2, 2) + _, sigma, _, err := SVDComplex(a) + if err != nil { + t.Fatalf("SVDComplex diag(1e160, 1e-160): %v", err) + } + if got := sigma.FloatAt(0) / 1e160; math.Abs(got-1) > 1e-13 { + t.Fatalf("diag(1e160, 1e-160): sigma[0] = %g, want 1e160", sigma.FloatAt(0)) + } + if got := sigma.FloatAt(1) / 1e-160; math.Abs(got-1) > 1e-13 { + t.Fatalf("diag(1e160, 1e-160): sigma[1] = %g, want 1e-160", sigma.FloatAt(1)) + } +} + +// TestEigenComplexSpectrumAtExtremeScale pins the closed-form spectrum of +// the Hermitian [[2s,is],[-is,2s]], which is {s, 3s}. Before the fix the +// Jacobi convergence test's raw sum of squares overflowed to +Inf at +// 1e200 (every sweep looked converged) and underflowed to 0 at 1e-200, +// so both returned the unrotated diagonal {2s, 2s} silently. +func TestEigenComplexSpectrumAtExtremeScale(t *testing.T) { + for _, s := range scales { + a := mustFromComplexes(t, []complex128{ + complex(2*s, 0), complex(0, s), + complex(0, -s), complex(2*s, 0), + }, 2, 2) + vals, vecs, err := EigenComplex(a) + if err != nil { + t.Fatalf("EigenComplex at scale %g: %v", s, err) + } + if got := vals.FloatAt(0) / s; math.Abs(got-1) > 1e-13 { + t.Fatalf("scale %g: eigenvalue 0 = %g, want 1 (in units of s)", s, got) + } + if got := vals.FloatAt(1) / s; math.Abs(got-3) > 1e-13 { + t.Fatalf("scale %g: eigenvalue 1 = %g, want 3 (in units of s)", s, got) + } + // Residual of both eigenpairs on the O(1) shift [[2,i],[-i,2]]. + base := []complex128{2, complex(0, 1), complex(0, -1), 2} + for k := range 2 { + lam := complex(vals.FloatAt(k)/s, 0) + worst := 0.0 + for i := range 2 { + acc := complex(0, 0) + for j := range 2 { + acc += base[i*2+j] * vecs.ComplexAt(j*2+k) + } + worst = math.Max(worst, cmplx.Abs(acc-lam*vecs.ComplexAt(i*2+k))) + } + if worst > 1e-13 { + t.Fatalf("scale %g: residual of eigenpair %d = %g, want <= 1e-13", s, k, worst) + } + } + } +} + +// TestSpEigenComplexAtExtremeScale pins the Hermitian Lanczos on the +// mirrored stencil against the closed-form spectrum of the ordinary +// matrix: 2 + cos(k·π/(n+1)) for k = 1..n, whose two largest values are +// the top Ritz pairs. Before the scaling the recurrence collapsed at +// 1e200 (the projected tridiagonal's squared magnitudes overflowed) and +// returned NaN, and at 1e-200 the values came back as zero. +func TestSpEigenComplexAtExtremeScale(t *testing.T) { + const n = 8 + want := []float64{ + 2 + math.Cos(math.Pi/float64(n+1)), + 2 + math.Cos(2*math.Pi/float64(n+1)), + } + base := hermitianStencil(n) + for _, s := range scales { + coo := cSRFromDense(t, scaleC(base, s), n) + vals, vecs, err := SpEigenComplex(coo, 2, core.NewGenerator(7)) + if err != nil { + t.Fatalf("SpEigenComplex at scale %g: %v", s, err) + } + for i := range 2 { + got := vals.FloatAt(i) + if math.IsNaN(got) { + t.Fatalf("scale %g: Ritz value %d = NaN", s, i) + } + if rel := got / s; math.Abs(rel-want[i]) > 1e-12 { + t.Fatalf("scale %g: Ritz value %d = %g, want %g (in units of s)", s, i, rel, want[i]) + } + } + // The Ritz vectors solve the O(1) stencil to rounding. + for k := range 2 { + lam := complex(vals.FloatAt(k)/s, 0) + worst := 0.0 + for i := range n { + acc := complex(0, 0) + for j := range n { + acc += base[i*n+j] * vecs.ComplexAt(j*2+k) + } + worst = math.Max(worst, cmplx.Abs(acc-lam*vecs.ComplexAt(i*2+k))) + } + if worst > 1e-12 { + t.Fatalf("scale %g: Ritz residual %d = %g, want <= 1e-12", s, k, worst) + } + } + } +} + +// TestSpSolveComplexAtExtremeScale pins both complex sparse solvers +// against a brute-force dense elimination of the mirrored-stencil system, +// at ordinary scale (the control the oracle pins) and at the extremes. +// Before the fix, a huge b made bNorm +Inf so every convergence check +// passed and both solvers returned after one step (residual 0.5s), and a +// tiny b made bNorm 0 so the zero vector was returned as the exact +// solution. +func TestSpSolveComplexAtExtremeScale(t *testing.T) { + const n = 8 + // Reference: the ordinary-scale system solved densely. + base := hermitianStencil(n) + rhs := make([]complex128, n) + for i := range n { + rhs[i] = complex(1, 0.25) + } + want := pinDenseSolve(t, base, rhs, n) + + for _, s := range scales { + coo := cSRFromDense(t, scaleC(base, s), n) + b := mustFromComplexes(t, scaleC(rhs, s), n) + x, err := SpSolveComplexCG(coo, b, 1e-12, 400) + if err != nil { + t.Fatalf("SpSolveComplexCG at scale %g: %v", s, err) + } + // (sA)x = s·b has exactly the ordinary-scale solution, so the + // returned x must equal the dense reference as it stands: the + // factor cancels and must not be applied again. + worst := 0.0 + for i := range n { + worst = math.Max(worst, cmplx.Abs(x.ComplexAt(i)-want[i])) + } + if worst > 1e-9 { + t.Fatalf("scale %g: CG solution differs from the dense reference by %g", s, worst) + } + // The same system through BiCGSTAB (Hermitian input is a valid + // special case of its contract). + coo2 := cSRFromDense(t, scaleC(base, s), n) + xb, err := SpSolveComplexBiCGSTAB(coo2, b, 1e-12, 400) + if err != nil { + t.Fatalf("SpSolveComplexBiCGSTAB at scale %g: %v", s, err) + } + worst = 0.0 + for i := range n { + worst = math.Max(worst, cmplx.Abs(xb.ComplexAt(i)-want[i])) + } + if worst > 1e-9 { + t.Fatalf("scale %g: BiCGSTAB solution differs from the dense reference by %g", s, worst) + } + } + + // The non-Hermitian shape BiCGSTAB exists for, same reference method. + nonHerm := make([]complex128, n*n) + for i := range n { + nonHerm[i*n+i] = complex(2, 1) + if i+1 < n { + nonHerm[i*n+i+1] = complex(1, 0.3) + nonHerm[(i+1)*n+i] = complex(0, -0.25) + } + } + wantNH := pinDenseSolve(t, nonHerm, rhs, n) + for _, s := range scales { + coo := cSRFromDense(t, scaleC(nonHerm, s), n) + b := mustFromComplexes(t, scaleC(rhs, s), n) + x, err := SpSolveComplexBiCGSTAB(coo, b, 1e-12, 400) + if err != nil { + t.Fatalf("non-Hermitian BiCGSTAB at scale %g: %v", s, err) + } + worst := 0.0 + for i := range n { + worst = math.Max(worst, cmplx.Abs(x.ComplexAt(i)-wantNH[i])) + } + if worst > 1e-9 { + t.Fatalf("non-Hermitian scale %g: solution differs from the dense reference by %g", s, worst) + } + } +} + +// TestSpEigenComplexErrorNamesK pins the malformed diagnostic (report +// F14): the k-range error had three verbs and two arguments and rendered +// as "%!s(int=2): k must be in [1, 5], got %!d(MISSING)". +func TestSpEigenComplexErrorNamesK(t *testing.T) { + coo := cSRFromDense(t, hermitianStencil(5), 5) + _, _, err := SpEigenComplex(coo, 6, nil) + if err == nil { + t.Fatal("SpEigenComplex with k > n: want an error") + } + msg := err.Error() + if strings.Contains(msg, "%!") { + t.Fatalf("malformed error message: %q", msg) + } + if !strings.Contains(msg, "SpEigenComplex: k must be in [1, 5], got 6") { + t.Fatalf("error message = %q, want it to name the entry point and the range", msg) + } + if _, _, err := SpEigenComplex(coo, 0, nil); err == nil || strings.Contains(err.Error(), "%!") { + t.Fatalf("k = 0 error = %v, want a well-formed range error", err) + } +} + +// TestSymmetryGuardsAreRelativeAtSmallScale pins the guards of Eigen, +// EigenComplex and the complex sparse Hermitian check on a matrix whose +// asymmetry is a tenth of its own scale: an absolute floor of 1e-15 +// approved such a matrix and the symmetric path then answered for a +// matrix that is neither A nor Aᵀ, while the wording of the refusal is +// unchanged. The exactly symmetric matrix of the same scale must still +// be accepted, so the rule is relative and not simply stricter. +func TestSymmetryGuardsAreRelativeAtSmallScale(t *testing.T) { + const s = 1e-14 + // Relative asymmetry 1e-1. + asym := mustFloats(t, []float64{s, 0, 1e-15, s}, 2, 2) + if _, _, err := Eigen(asym); err == nil { + t.Fatal("Eigen: asymmetric at 10% of scale 1e-14 was accepted, want a refusal") + } else if !strings.Contains(err.Error(), "Eigen: matrix is not symmetric within 1e-12 tolerance") { + t.Fatalf("Eigen refusal = %q, want the unchanged wording", err) + } + // A genuinely symmetric matrix of the same tiny scale still passes + // the guard: if it did not, the test would pin a stricter rule than + // the contract asks for (the tiny p.d. case must reach the solver). + sym := mustFloats(t, []float64{s, 1e-15, 1e-15, s}, 2, 2) + if _, _, err := Eigen(sym); err != nil { + t.Fatalf("Eigen of a symmetric matrix at scale %g: %v", s, err) + } + + nonHerm := mustFromComplexes(t, []complex128{ + complex(s, 0), complex(1e-15, 0), + 0, complex(s, 0), + }, 2, 2) + if _, _, err := EigenComplex(nonHerm); err == nil { + t.Fatal("EigenComplex: non-Hermitian at 10% of scale 1e-14 was accepted, want a refusal") + } else if !strings.Contains(err.Error(), "EigenComplex: matrix is not Hermitian within 1e-12 tolerance") { + t.Fatalf("EigenComplex refusal = %q, want the unchanged wording", err) + } + herm := mustFromComplexes(t, []complex128{ + complex(s, 0), complex(0, 1e-15), + complex(0, -1e-15), complex(s, 0), + }, 2, 2) + if _, _, err := EigenComplex(herm); err != nil { + t.Fatalf("EigenComplex of a Hermitian matrix at scale %g: %v", s, err) + } + + // The sparse complex twin: both mirrored entries are stored, and the + // asymmetry between i·1e-16 and i·1e-17 is about a percent of the + // 1e-14 scale yet below the 1e-15 absolute floor, so the pre-fix guard + // approved it and Lanczos ran on a matrix that is not Hermitian. The + // refusal wording is pinned as well. + coo := cSRFromDense(t, []complex128{ + complex(s, 0), complex(0, 1e-16), + complex(0, 1e-17), complex(s, 0), + }, 2) + if _, _, err := SpEigenComplex(coo, 1, core.NewGenerator(3)); err == nil { + t.Fatal("SpEigenComplex: non-Hermitian at 10% of scale 1e-14 was accepted, want a refusal") + } else if !strings.Contains(err.Error(), "SpEigenComplex: matrix is not Hermitian within 1e-12 tolerance") { + t.Fatalf("SpEigenComplex refusal = %q, want the unchanged wording", err) + } + // The exactly Hermitian matrix of the same scale must still be + // accepted: the rule is relative, not simply stricter. + hermSparse := cSRFromDense(t, []complex128{ + complex(s, 0), complex(0, 1e-16), + complex(0, -1e-16), complex(s, 0), + }, 2) + if _, _, err := SpEigenComplex(hermSparse, 1, core.NewGenerator(3)); err != nil { + t.Fatalf("SpEigenComplex of a Hermitian matrix at scale %g: %v", s, err) + } +} + +// TestNoArrowSymbolsInSource pins the source convention: the arrow +// code points must not appear as assignment notation in comments, +// which the house style forbids. +func TestNoArrowSymbolsInSource(t *testing.T) { + for _, name := range []string{ + "decomp2.go", "decomp3.go", "eigenreal.go", + "schur.go", "csvd.go", "sparsecomplex.go", + } { + src, err := os.ReadFile(name) + if err != nil { + t.Fatalf("read %s: %v", name, err) + } + for i, line := range strings.Split(string(src), "\n") { + for _, r := range []rune{'←', '→', '↑', '↓', '⇒'} { + if strings.ContainsRune(line, r) { + t.Fatalf("%s:%d contains U+%04X: %s", name, i+1, r, line) + } + } + } + } +} + +// stencil builds the n×n tridiagonal 2 / 0.5 stencil. Its spectrum +// is closed form, 2 + cos(k·π/(n+1)) for k = 1..n, which is the reference +// the sparse eigensolvers below are checked against. +func stencil(n int) []float64 { + out := make([]float64, n*n) + for i := range n { + out[i*n+i] = 2 + if i+1 < n { + out[i*n+i+1] = 0.5 + out[(i+1)*n+i] = 0.5 + } + } + return out +} + +// topStencilValues is the two largest eigenvalues of an n×n 2 / 0.5 +// stencil, derived from that closed form. +func topStencilValues(n int) []float64 { + return []float64{ + 2 + math.Cos(math.Pi/float64(n+1)), + 2 + math.Cos(2*math.Pi/float64(n+1)), + } +} + +// realCOO builds a real SparseCOO from a dense flat matrix, +// dropping the exact zeros the way a caller's assembly would. +func realCOO(t *testing.T, a []float64, n int) *core.SparseCOO { + t.Helper() + var idx []int64 + var vals []float64 + for i := range n { + for j := range n { + if a[i*n+j] == 0 { + continue + } + idx = append(idx, int64(i), int64(j)) + vals = append(vals, a[i*n+j]) + } + } + ind, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + val, err := core.FromFloats(vals, len(vals)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + coo, err := core.NewSparseCOO(ind, val, []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// ritzResidual returns max_k ‖(A/s)·u_k − (λ_k/s)·u_k‖ for the k +// Ritz pairs of a sparse matrix held as a COO, with A/s formed from the +// stored values so nothing overflows at either extreme of the scale +// sweep. vals and vecs are the returned values and the (n, k) vectors. +func ritzResidual(a *core.SparseCOO, vals, vecs []complex128, n, k int, s float64) float64 { + idx := a.Indices.RawInts() + nnz := a.Indices.Shape()[0] + worst := 0.0 + for kk := range k { + lam := vals[kk] / complex(s, 0) + for i := range n { + acc := complex(0, 0) + for p := range nnz { + if int(idx[p*2]) != i { + continue + } + v := complex(0, 0) + if a.Values.Dtype() == core.Complex { + v = a.Values.ComplexAt(p) + } else { + v = complex(a.Values.FloatAt(p), 0) + } + acc += (v / complex(s, 0)) * vecs[int(idx[p*2+1])*k+kk] + } + if d := cmplx.Abs(acc - lam*vecs[i*k+kk]); d > worst { + worst = d + } + } + } + return worst +} + +// ritzResidualReal is ritzResidual for a real Ritz pair set. +func ritzResidualReal(a *core.SparseCOO, vals []float64, vecs *core.Array, n, k int, s float64) float64 { + cv := make([]complex128, n*k) + for i := range n { + for j := range k { + cv[i*k+j] = complex(vecs.FloatAt(i*k+j), 0) + } + } + cvals := make([]complex128, k) + for j := range k { + cvals[j] = complex(vals[j], 0) + } + return ritzResidual(a, cvals, cv, n, k, s) +} + +// TestSparseSymmetryGuardIsRelativeAtSmallScale pins the fourth absolute +// floor (report F16, sparseigen.go checkSymmetric, the guard behind +// SpEigen, SpSolve and SpExpApply): an asymmetry that is a percent of a +// 1e-14 matrix, yet below the old 1e-15 floor, must be refused with the +// same wording as its dense and complex siblings, and the exactly +// symmetric matrix of that scale must still be accepted. +func TestSparseSymmetryGuardIsRelativeAtSmallScale(t *testing.T) { + const s = 1e-14 + // Both mirrored entries are stored; the difference is 9e-17, below + // the removed 1e-15 floor and above the purely relative 1e-26. + asym := realCOO(t, []float64{s, 1e-16, 1e-17, s}, 2) + b := mustFloats(t, []float64{s, s}, 2) + if _, _, err := SpEigen(asym, 2, core.NewGenerator(7)); err == nil { + t.Fatal("SpEigen: asymmetric at a percent of scale 1e-14 was accepted, want a refusal") + } else if !strings.Contains(err.Error(), "SpEigen: matrix is not symmetric within 1e-12 tolerance") { + t.Fatalf("SpEigen refusal = %q, want the sibling wording", err) + } + if _, err := SpSolve(asym, b, 1e-12, 100); err == nil { + t.Fatal("SpSolve: asymmetric at a percent of scale 1e-14 was accepted, want a refusal") + } else if !strings.Contains(err.Error(), "SpSolve: matrix is not symmetric within 1e-12 tolerance") { + t.Fatalf("SpSolve refusal = %q, want the sibling wording", err) + } + if _, err := SpExpApply(asym, b, 0); err == nil { + t.Fatal("SpExpApply: asymmetric at a percent of scale 1e-14 was accepted, want a refusal") + } else if !strings.Contains(err.Error(), "SpExpApply: matrix is not symmetric within 1e-12 tolerance") { + t.Fatalf("SpExpApply refusal = %q, want the sibling wording", err) + } + + // The exactly symmetric matrix of the same scale passes all three. + sym := realCOO(t, []float64{s, 1e-16, 1e-16, s}, 2) + if _, _, err := SpEigen(sym, 2, core.NewGenerator(7)); err != nil { + t.Fatalf("SpEigen of a symmetric matrix at scale %g: %v", s, err) + } + if _, err := SpSolve(sym, b, 1e-12, 100); err != nil { + t.Fatalf("SpSolve of a symmetric matrix at scale %g: %v", s, err) + } + if _, err := SpExpApply(sym, b, 0); err != nil { + t.Fatalf("SpExpApply of a symmetric matrix at scale %g: %v", s, err) + } +} + +// TestSpEigenSpectrumAtExtremeScale pins the real Hermitian Lanczos on the +// 2 / 0.5 stencil against the closed-form spectrum 2+cos(k·π/(n+1)). The +// projected tridiagonal's squared accumulation overflows in the symmetric +// sweep above about 1.3e154, which deflates the whole block at once and +// returns the raw Rayleigh quotients as the spectrum. +func TestSpEigenSpectrumAtExtremeScale(t *testing.T) { + const n = 8 + want := topStencilValues(n) + for _, s := range scales { + coo := realCOO(t, scale(stencil(n), s), n) + vals, vecs, err := SpEigen(coo, 2, core.NewGenerator(7)) + if err != nil { + t.Fatalf("SpEigen at scale %g: %v", s, err) + } + for i := range 2 { + got := vals.FloatAt(i) / s + if math.IsNaN(got) || math.IsInf(got, 0) { + t.Fatalf("scale %g: Ritz value %d = %v, want %g", s, i, vals.FloatAt(i), want[i]) + } + if math.Abs(got-want[i]) > 1e-12 { + t.Fatalf("scale %g: Ritz value %d = %g, want %g (in units of s)", s, i, got, want[i]) + } + } + if res := ritzResidualReal(coo, vals.RawFloats(), vecs, n, 2, s); res > 1e-12 { + t.Fatalf("scale %g: Ritz residual %g, want <= 1e-12", s, res) + } + } +} + +// TestSpEigenGeneralSpectrumAtExtremeScale pins both general sparse +// eigensolvers on the same closed-form spectrum. The Arnoldi recurrence +// computes its norms as a raw sum of squares on the unscaled operator, so +// above about 1.3e154 the projected coupling became +Inf and below about +// 1.5e-162 it became zero, which the exhaustion test read as a collapsed +// Krylov block: both extremes returned wrong Ritz values silently. +func TestSpEigenGeneralSpectrumAtExtremeScale(t *testing.T) { + const n = 8 + want := topStencilValues(n) + // The Hermitian stencil for the complex entry: 2 on the diagonal, + // conjugate mirrored imaginary couplings, same closed-form spectrum. + cb := hermitianStencil(n) + // A real symmetric matrix is a valid, though not typical, input to + // the general entry, and its real spectrum makes the reference + // unambiguous. + cScales := []float64{1, 1e150, 1e200, 1e-150, 1e-200} + for _, s := range cScales { + coo := realCOO(t, scale(stencil(n), s), n) + vals, vecs, err := SpEigenGeneral(coo, 2, core.NewGenerator(5)) + if err != nil { + t.Fatalf("SpEigenGeneral at scale %g: %v", s, err) + } + cv := vals.RawComplexes() + for i := range 2 { + got := cmplx.Abs(cv[i]) / s + if math.Abs(got-want[i]) > 1e-12 { + t.Fatalf("SpEigenGeneral scale %g: |λ| %d = %g, want %g", s, i, got, want[i]) + } + } + if res := ritzResidual(coo, cv, vecs.RawComplexes(), n, 2, s); res > 1e-12 { + t.Fatalf("SpEigenGeneral scale %g: Ritz residual %g, want <= 1e-12", s, res) + } + + ccoo := cSRFromDense(t, scaleC(cb, s), n) + cvals, cvecs, err := SpEigenGeneralComplex(ccoo, 2, core.NewGenerator(5)) + if err != nil { + t.Fatalf("SpEigenGeneralComplex at scale %g: %v", s, err) + } + ccv := cvals.RawComplexes() + for i := range 2 { + got := cmplx.Abs(ccv[i]) / s + if math.Abs(got-want[i]) > 1e-12 { + t.Fatalf("SpEigenGeneralComplex scale %g: |λ| %d = %g, want %g", s, i, got, want[i]) + } + } + if res := ritzResidual(ccoo, ccv, cvecs.RawComplexes(), n, 2, s); res > 1e-12 { + t.Fatalf("SpEigenGeneralComplex scale %g: Ritz residual %g, want <= 1e-12", s, res) + } + } +} diff --git a/linalg/fitpoly_test.go b/linalg/fitpoly_test.go new file mode 100644 index 0000000..292ea90 --- /dev/null +++ b/linalg/fitpoly_test.go @@ -0,0 +1,33 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestFitEvaluatePolynomial moved with FitPolynomial from the root +// package: a quadratic fit of y = 2 + x must recover it perfectly. +func TestFitEvaluatePolynomial(t *testing.T) { + x := mustFloats(t, []float64{0, 1, 2}, 3) + y := mustFloats(t, []float64{2, 3, 4}, 3) + coeffs, err := FitPolynomial(x, y, 1) + if err != nil { + t.Fatal(err) + } + if math.Abs(coeffs.FloatAt(0)-2) > 1e-9 || math.Abs(coeffs.FloatAt(1)-1) > 1e-9 { + t.Errorf("poly coeffs: %v", coeffs.RawFloats()) + } + q, _ := core.FromFloats([]float64{5, 7}, 2) + out, err := core.EvaluatePolynomial(coeffs, q) + if err != nil { + t.Fatal(err) + } + if math.Abs(out.FloatAt(0)-7) > 1e-9 || math.Abs(out.FloatAt(1)-9) > 1e-9 { + t.Errorf("poly eval: %v", out.RawFloats()) + } +} diff --git a/linalg/geneigen.go b/linalg/geneigen.go new file mode 100644 index 0000000..bdc3eda --- /dev/null +++ b/linalg/geneigen.go @@ -0,0 +1,131 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The generalised symmetric eigenproblem A·v = λ·B·v, the standard form +// of vibrating-system and covariance questions: the eigenvalues of the +// pencil (A, B) with B symmetric positive definite. The Cholesky route +// reduces it to the ordinary symmetric problem without ever forming +// B⁻¹A, whose asymmetry would square the conditioning. + +// EigenGeneralised solves A·v = λ·B·v for a symmetric a and a +// symmetric positive definite b, both real n×n. b = L·Lᵀ turns the +// pencil into the standard symmetric problem for C = L⁻¹·A·L⁻ᵀ, which +// shares the eigenvalues; its ordinary eigenvectors y transform back as +// v = L⁻ᵀ·y, which lands them B-orthonormal (vᵀ·B·v = 1) for free. +// Values come back ascending in a 1-D array with the eigenvectors as +// the matching columns, the convention Eigen uses. A complex input, a +// size mismatch, or a b that fails its Cholesky factorisation is an +// error; a itself must be symmetric, which is not verified. +func EigenGeneralised(a, b *core.Array) (values, vectors *core.Array, err error) { + const name = "EigenGeneralised" + if a.Dtype() == core.Complex || b.Dtype() == core.Complex { + return nil, nil, base.Errf("%s: complex pencils are not supported", name) + } + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, nil, base.Errf("%s: a must be a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + if b.NDim() != 2 || b.Shape()[0] != b.Shape()[1] { + return nil, nil, base.Errf("%s: b must be a square 2-D matrix, got shape %s", name, base.ShapeText(b.Shape())) + } + n := a.Shape()[0] + if b.Shape()[0] != n { + return nil, nil, base.Errf("%s: size mismatch, a is %d×%d and b is %d×%d", + name, n, n, b.Shape()[0], b.Shape()[1]) + } + if n == 0 { + return nil, nil, base.Errf("%s: zero-sized pencil", name) + } + l, err := Cholesky(b) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + lFlat := denseFloats(l, n, n) + // solveSystem consumes its matrix in place, so each solve gets a + // fresh copy of L's rows as views over one flat backing slice. + freshRows := func() [][]float64 { + back := make([]float64, n*n) + copy(back, lFlat) + rows := make([][]float64, n) + for i := range n { + rows[i] = back[i*n : (i+1)*n] + } + return rows + } + // X = L⁻¹·A, one column of a per right-hand side. + aCols := make([][]float64, n) + for j := range n { + aCols[j] = make([]float64, n) + for i := range n { + aCols[j][i] = a.FloatAt(i*n + j) + } + } + if _, err := base.SolveSystem(name, freshRows(), aCols); err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + // C = X·L⁻ᵀ, gathered by solving L·Z = Xᵀ and transposing. + cMat := make([]float64, n*n) + { + xT := make([][]float64, n) + for j := range n { + xT[j] = make([]float64, n) + for i := range n { + xT[j][i] = aCols[i][j] // column j of Xᵀ is row j of X + } + } + if _, err := base.SolveSystem(name, freshRows(), xT); err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + for i := range n { + for j := range n { + cMat[i*n+j] = xT[j][i] + } + } + } + // Rounding leaves C a hair off symmetric; the eigensolver wants the + // exact form, so take the symmetric part. + for i := range n { + for j := i + 1; j < n; j++ { + m := (cMat[i*n+j] + cMat[j*n+i]) / 2 + cMat[i*n+j] = m + cMat[j*n+i] = m + } + } + cArr := floatsToArray(cMat, []int{n, n}) + values, yArr, err := Eigen(cArr) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + // V = L⁻ᵀ·Y: each eigenvector column solves Lᵀ·v = y. + ltBack := make([]float64, n*n) + lt := make([][]float64, n) + for i := range n { + lt[i] = ltBack[i*n : (i+1)*n] + for j := range n { + lt[i][j] = lFlat[j*n+i] + } + } + yCols := make([][]float64, n) + for j := range n { + yCols[j] = make([]float64, n) + for i := range n { + yCols[j][i] = yArr.FloatAt(i*n + j) + } + } + if _, err := base.SolveSystem(name, lt, yCols); err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + vMat := make([]float64, n*n) + for j := range n { + for i := range n { + vMat[i*n+j] = yCols[j][i] + } + } + return values, floatsToArray(vMat, []int{n, n}), nil +} diff --git a/linalg/geneigen_test.go b/linalg/geneigen_test.go new file mode 100644 index 0000000..2c49442 --- /dev/null +++ b/linalg/geneigen_test.go @@ -0,0 +1,131 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// pencilSample builds a deterministic symmetric a and a symmetric +// positive definite b of size n. +func pencilSample(n int) (*core.Array, *core.Array) { + raw := make([]float64, n*n) + for i := range n { + for j := range n { + raw[i*n+j] = math.Sin(float64(2*i + 3*j + 1)) + } + } + av := make([]float64, n*n) + bv := make([]float64, n*n) + for i := range n { + for j := range n { + av[i*n+j] = raw[i*n+j] + raw[j*n+i] + bv[i*n+j] = math.Cos(float64(3*i+2*j)) + math.Cos(float64(3*j+2*i)) + } + bv[i*n+i] += float64(n) + 1 + } + a, _ := core.FromFloats(av, n, n) + b, _ := core.FromFloats(bv, n, n) + return a, b +} + +// TestEigenGeneralisedDiagonal pins the pencil on a diagonal pair, +// where the eigenvalues are the entrywise ratios and the eigenvectors +// are the scaled unit vectors. +func TestEigenGeneralisedDiagonal(t *testing.T) { + a := mustFloats(t, []float64{1, 0, 0, 4}, 2, 2) + b := mustFloats(t, []float64{1, 0, 0, 2}, 2, 2) + values, vectors, err := EigenGeneralised(a, b) + if err != nil { + t.Fatalf("EigenGeneralised: %v", err) + } + if math.Abs(values.FloatAt(0)-1) > 1e-12 || math.Abs(values.FloatAt(1)-2) > 1e-12 { + t.Fatalf("values = (%.12g, %.12g), want (1, 2)", values.FloatAt(0), values.FloatAt(1)) + } + // First eigenvector: e1 scaled to unit B norm = (1, 0); second: + // e2 with 2·vᵀBv... v = (0, 1/√2) so vᵀBv = (1/2)·2 = 1. + if math.Abs(vectors.FloatAt(0)-1) > 1e-12 || math.Abs(vectors.FloatAt(3)-1/math.Sqrt2) > 1e-12 { + t.Fatalf("vectors = (%.12g, %.12g; %.12g, %.12g), want (1, 0; 0, 1/√2)", + vectors.FloatAt(0), vectors.FloatAt(1), vectors.FloatAt(2), vectors.FloatAt(3)) + } +} + +// TestEigenGeneralisedResidual checks the defining equations on a +// deterministic pencil: A·X = B·X·Λ to rounding, the vectors +// B-orthonormal, the values ascending. +func TestEigenGeneralisedResidual(t *testing.T) { + n := 6 + a, b := pencilSample(n) + values, vectors, err := EigenGeneralised(a, b) + if err != nil { + t.Fatalf("EigenGeneralised: %v", err) + } + scale := 0.0 + for i := range a.Len() { + scale = math.Max(scale, math.Abs(a.FloatAt(i))+math.Abs(b.FloatAt(i))) + } + worst := 0.0 + for j := range n { + lam := values.FloatAt(j) + for i := range n { + // (A·X − λ·B·X) at row i, column j. + ax, bx := 0.0, 0.0 + for k := range n { + ax += a.FloatAt(i*n+k) * vectors.FloatAt(k*n+j) + bx += b.FloatAt(i*n+k) * vectors.FloatAt(k*n+j) + } + worst = math.Max(worst, math.Abs(ax-lam*bx)) + } + } + if worst > 1e-8*scale { + t.Fatalf("pencil residual %.3g exceeds %.3g", worst, 1e-8*scale) + } + // B-orthonormality: Xᵀ·B·X = I. + for j := range n { + for i := j; i < n; i++ { + s := 0.0 + for k := range n { + for l := range n { + s += vectors.FloatAt(k*n+i) * b.FloatAt(k*n+l) * vectors.FloatAt(l*n+j) + } + } + want := 0.0 + if i == j { + want = 1 + } + if math.Abs(s-want) > 1e-8 { + t.Fatalf("XᵀBX[%d][%d] = %.12g, want %.12g", i, j, s, want) + } + } + } + for i := 1; i < n; i++ { + if values.FloatAt(i) < values.FloatAt(i-1) { + t.Fatalf("values not ascending at %d", i) + } + } +} + +// TestEigenGeneralisedErrors pins the validation contract. +func TestEigenGeneralisedErrors(t *testing.T) { + a, b := pencilSample(3) + cx, _ := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + if _, _, err := EigenGeneralised(cx, b); err == nil { + t.Fatal("expected an error for a complex pencil") + } + bad, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if _, _, err := EigenGeneralised(a, bad); err == nil { + t.Fatal("expected an error for mismatched sizes") + } + rect, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 3, 2) + if _, _, err := EigenGeneralised(rect, b); err == nil { + t.Fatal("expected an error for a rectangular a") + } + // b = all-ones is symmetric but singular: the Cholesky must refuse. + singular, _ := core.FromFloats([]float64{1, 1, 1, 1}, 2, 2) + if _, _, err := EigenGeneralised(mustFloats(t, []float64{1, 0, 0, 1}, 2, 2), singular); err == nil { + t.Fatal("expected an error for a singular b") + } +} diff --git a/linalg/gmres.go b/linalg/gmres.go new file mode 100644 index 0000000..57bc26c --- /dev/null +++ b/linalg/gmres.go @@ -0,0 +1,310 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// GMRES returns the solution of A·x = b by restarted GMRES. op maps a +// vector of length n to A·v. restart bounds the Krylov dimension per +// cycle (≤ 0 means n); maxIter bounds the outer cycles (≤ 0 means 20); +// tol is the relative residual target (≤ 0 means 1e-10). Running out +// of cycles with the target unmet is an error naming the best residual +// achieved, never a silent approximation. +// +// The vector op is handed is a read-only buffer the solver owns and +// refills for every call: op must read it within the call and must not +// retain or modify it. An operator that needs to keep its argument has +// to copy it. The array op returns is op's own and the solver only +// reads it, so an operator may hand back a buffer it reuses. +// +// Every new Hessenberg column the Arnoldi process produces is folded +// through the Givens rotations of the previous columns the moment it +// appears, which leaves a triangular projected system behind and +// tracks the least-squares residual incrementally: s starts at β and +// each rotation j updates s[j+1] = −sn[j]·s[j], so |s[j+1]| is the +// residual of the growing fit without ever solving it. That keeps the +// projected problem at the condition of the Hessenberg matrix instead +// of its square, which is what the normal-equations route would pay. +func GMRES(op func(*core.Array) (*core.Array, error), b *core.Array, restart, maxIter int, tol float64) (*core.Array, error) { + if b.NDim() != 1 { + return nil, base.Errf("GMRES: b must be a rank-1 vector, got shape %s", base.ShapeText(b.Shape())) + } + n := b.Len() + if n == 0 { + return nil, base.Errf("GMRES: empty system") + } + if restart <= 0 || restart > n { + restart = n + } + if maxIter <= 0 { + maxIter = 20 + } + if tol <= 0 { + tol = 1e-10 + } + if !isRealDtype(b) { + return nil, base.Errf("GMRES: b must be a real dtype, got %s", b.Dtype()) + } + bv := vectorF64(b, n) + bNorm := norm2F64(bv) + if bNorm == 0 { + return core.Zeros(core.Float, n) + } + + x := make([]float64, n) + m := min(restart, n) + // The whole (m+1)×n basis and the (m+1)×m Hessenberg matrix live in + // one flat buffer each, allocated once per solve: a row is a + // subslice, so the Arnoldi process adds no allocation per step. + basisFlat := make([]float64, (m+1)*n) + basis := make([][]float64, m+1) + for i := range m + 1 { + basis[i] = basisFlat[i*n : (i+1)*n] + } + hessFlat := make([]float64, (m+1)*m) + hess := make([][]float64, m+1) + for i := range m + 1 { + hess[i] = hessFlat[i*m : (i+1)*m] + } + cs := make([]float64, m) + sn := make([]float64, m) + s := make([]float64, m+1) + // The working vectors are reused across columns and cycles: w holds + // the current Arnoldi remainder, r the residual, ax A·x, and y the + // triangular solve (only its first nsteps entries are read per + // cycle). + w := make([]float64, n) + r := make([]float64, n) + axv := make([]float64, n) + y := make([]float64, m) + // The operand op receives is one buffer the solver owns and + // refills before every call, so a solve allocates nothing per + // right-hand side. The contract that makes it safe is stated in the + // doc comment: op must not retain or modify the vector. + opBuf := core.New(core.Float, n) + opVals := opBuf.RawFloats() + applyOp := func(src []float64) (*core.Array, error) { + copy(opVals, src) + return op(opBuf) + } + best := math.Inf(1) + converged := false + // The op call that closes a cycle computes A·x for the x the cycle + // produced, which is exactly what the next cycle's residual needs, + // so the result is reused rather than recomputed for the same x. + // Every call to op still gets a fresh array: the callback may keep + // the vector it is given, so the operand is the one buffer here + // that cannot be recycled between calls. + freshAx := false + // The Arnoldi breakdown floor is relative to the operator's own + // Hessenberg scale, not to ‖b‖: a legitimate system with ‖A‖ ≪ ‖b‖ + // would otherwise collapse every cycle at its first column. This is + // the same purely relative floor the Lanczos code documents. + hScale := 0.0 + for range maxIter { + if !freshAx { + ax, aerr := applyOp(x) + if aerr != nil { + return nil, base.Errf("GMRES: %w", aerr) + } + if !isRealDtype(ax) { + return nil, base.Errf("GMRES: op returned %s, want a real dtype", ax.Dtype()) + } + if ax.Len() != n { + return nil, base.Errf("GMRES: op returned %d elements for a problem of %d", ax.Len(), n) + } + vecF64Into(axv, ax, n) + } + // Residual r = b − A·x. + for i := range n { + r[i] = bv[i] - axv[i] + } + // norm2F64 skips NaN entries, so an all-NaN residual would read + // as a zero norm and as convergence; a non-finite op output is a + // breakdown before that. + if !vecFinite(r) { + return nil, base.Errf("GMRES: non-finite residual from op") + } + beta := norm2F64(r) + if beta < best { + best = beta + } + if beta <= tol*bNorm { + converged = true + break + } + for i := range n { + basis[0][i] = r[i] / beta + } + s[0] = beta + + // Arnoldi, one column at a time, with the rotations folded in + // before the next column is built. + nsteps := 0 + for j := range m { + opv, oerr := applyOp(basis[j]) + if oerr != nil { + return nil, base.Errf("GMRES: %w", oerr) + } + if !isRealDtype(opv) { + return nil, base.Errf("GMRES: op returned %s, want a real dtype", opv.Dtype()) + } + if opv.Len() != n { + return nil, base.Errf("GMRES: op returned %d elements for a problem of %d", opv.Len(), n) + } + vecF64Into(w, opv, n) + if !vecFinite(w) { + return nil, base.Errf("GMRES: non-finite op output at Arnoldi column %d", j) + } + for i := 0; i <= j; i++ { + dot := dotF64(w, basis[i]) + hess[i][j] = dot + hScale = max(hScale, math.Abs(dot)) + for k := range n { + w[k] -= dot * basis[i][k] + } + } + hSub := norm2F64(w) + hess[j+1][j] = hSub + hScale = max(hScale, hSub) + nsteps = j + 1 + + // Fold rotations 0..j−1 into the fresh column. + for k := range j { + h1, h2 := hess[k][j], hess[k+1][j] + hess[k][j] = cs[k]*h1 + sn[k]*h2 + hess[k+1][j] = -sn[k]*h1 + cs[k]*h2 + } + // Rotation j zeroes the subdiagonal entry: the pair + // (hess[j][j], hSub) turns into (denom, 0). + denom := math.Hypot(hess[j][j], hSub) + hScale = max(hScale, denom) + if denom == 0 { + cs[j], sn[j] = 1, 0 + } else { + cs[j], sn[j] = hess[j][j]/denom, hSub/denom + } + hess[j][j] = denom + hess[j+1][j] = 0 + s[j+1] = -sn[j] * s[j] + s[j] *= cs[j] + + // A collapsed direction means the Krylov space is + // exhausted; a converged estimate means the basis already + // spans the answer. Either way the cycle ends here. + if hSub <= float64(n)*base.EpsF*hScale || math.Abs(s[j+1]) <= tol*bNorm { + break + } + next := basis[j+1] + for k := range n { + next[k] = w[k] / hSub + } + } + + // Triangular solve R·y = s[0:nsteps] over the rotated leading + // block, then x = x + V·y. + for i := nsteps - 1; i >= 0; i-- { + t := s[i] + for k := i + 1; k < nsteps; k++ { + t -= hess[i][k] * y[k] + } + if hess[i][i] == 0 { + return nil, base.Errf("GMRES: singular projected system at step %d", i) + } + y[i] = t / hess[i][i] + } + for i := range nsteps { + for k := range n { + x[k] += y[i] * basis[i][k] + } + } + + // Trust only the true residual between cycles. + ax, aerr := applyOp(x) + if aerr != nil { + return nil, base.Errf("GMRES: %w", aerr) + } + if !isRealDtype(ax) { + return nil, base.Errf("GMRES: op returned %s, want a real dtype", ax.Dtype()) + } + vecF64Into(axv, ax, n) + if !vecFinite(axv) { + return nil, base.Errf("GMRES: non-finite op output between cycles") + } + freshAx = true + res := 0.0 + for k := range n { + d := bv[k] - axv[k] + res += d * d + } + residual := math.Sqrt(res) + if residual < best { + best = residual + } + if residual <= tol*bNorm { + converged = true + break + } + } + // Contract consistency with SpSolve and FindRootSystem: running out + // of cycles is an error naming the best residual achieved, never a + // silent approximation. + if !converged { + return nil, base.Errf("GMRES: no convergence in %d cycles, residual %.3g (tolerance %.3g)", + maxIter, best, tol*bNorm) + } + return floatsToArray(x, []int{n}), nil +} + +// norm2F64 returns the L2 norm of a float64 slice. The squares are +// summed relative to the largest magnitude so entries near 1e154 do +// not overflow on the way to the norm. +func norm2F64(a []float64) float64 { + maxAbs := 0.0 + for _, v := range a { + if v := math.Abs(v); v > maxAbs { + maxAbs = v + } + } + if maxAbs == 0 { + return 0 + } + s := 0.0 + for _, v := range a { + v /= maxAbs + s += v * v + } + return maxAbs * math.Sqrt(s) +} + +// ArrayFromFloatsSafe builds an array, copying the values, so callers +// keep ownership of the slice they passed in. +func ArrayFromFloatsSafe(v []float64, n int) *core.Array { + vals := make([]float64, n) + copy(vals, v) + a, _ := core.FloatsFromArray(vals, n) + return a +} + +// vecF64Into fills dst with an array's leading n elements, the +// writing twin of vectorF64 for a destination that is reused. +func vecF64Into(dst []float64, a *core.Array, n int) { + for i := range n { + dst[i] = a.FloatAt(i) + } +} + +// dotF64 returns the dot product of two float64 slices. +func dotF64(a, b []float64) float64 { + s := 0.0 + for i := range a { + s += a[i] * b[i] + } + return s +} diff --git a/linalg/gmres_test.go b/linalg/gmres_test.go new file mode 100644 index 0000000..3aa6e70 --- /dev/null +++ b/linalg/gmres_test.go @@ -0,0 +1,167 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +func TestGMRES(t *testing.T) { + vals := []float64{ + 10, 1, 0, 2, + 1, 12, 3, 0, + 0, 3, 15, 1, + 2, 0, 1, 8, + } + op := func(v *core.Array) (*core.Array, error) { + out := core.New(core.Float, 4) + for i := range 4 { + s := 0.0 + for j := range 4 { + s += vals[i*4+j] * v.FloatAt(j) + } + out.RawFloats()[i] = s + } + return out, nil + } + b := mustFloats(t, []float64{1, 2, 3, 4}, 4) + x, err := GMRES(op, b, 0, 0, 1e-12) + if err != nil { + t.Fatalf("GMRES: %v", err) + } + a := mustFloats(t, vals, 4, 4) + ref, _ := Solve(a, b) + for i := range 4 { + if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-10 { + t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i)) + } + } +} + +func TestGMRESLarger(t *testing.T) { + const n = 50 + vals := make([]float64, n*n) + for i := range n { + vals[i*n+i] = 4 + if i+1 < n { + vals[i*n+i+1] = -1 + vals[(i+1)*n+i] = -1 + } + } + op := func(v *core.Array) (*core.Array, error) { + out := core.New(core.Float, n) + for i := range n { + s := 0.0 + for j := range n { + s += vals[i*n+j] * v.FloatAt(j) + } + out.RawFloats()[i] = s + } + return out, nil + } + bv := make([]float64, n) + for i := range n { + bv[i] = float64(i + 1) + } + b := mustFloats(t, bv, n) + x, err := GMRES(op, b, 0, 0, 1e-12) + if err != nil { + t.Fatalf("GMRES: %v", err) + } + for i := range n { + s := 0.0 + for j := range n { + s += vals[i*n+j] * x.FloatAt(j) + } + if math.Abs(s-bv[i]) > 1e-8 { + t.Fatalf("res[%d] = %e, want < 1e-8", i, math.Abs(s-bv[i])) + } + } +} + +// TestGMRESRestarted forces several outer cycles by capping the +// Krylov dimension below the system size; the accumulated solution +// must still match the dense solve. +func TestGMRESRestarted(t *testing.T) { + vals := []float64{ + 10, 1, 0, 2, + 1, 12, 3, 0, + 0, 3, 15, 1, + 2, 0, 1, 8, + } + op := func(v *core.Array) (*core.Array, error) { + out := core.New(core.Float, 4) + for i := range 4 { + s := 0.0 + for j := range 4 { + s += vals[i*4+j] * v.FloatAt(j) + } + out.RawFloats()[i] = s + } + return out, nil + } + b := mustFloats(t, []float64{1, 2, 3, 4}, 4) + a := mustFloats(t, vals, 4, 4) + ref, err := Solve(a, b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + x, err := GMRES(op, b, 2, 50, 1e-12) + if err != nil { + t.Fatalf("GMRES: %v", err) + } + for i := range 4 { + if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-10 { + t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i)) + } + } +} + +func TestGMRESErrors(t *testing.T) { + b := mustFloats(t, []float64{1, 2}, 2) + if _, err := GMRES(func(*core.Array) (*core.Array, error) { + return nil, base.Errf("operator failed") + }, b, 0, 0, 0); err == nil { + t.Fatal("operator error: want an error") + } + if x, err := GMRES(func(*core.Array) (*core.Array, error) { + return nil, base.Errf("operator failed") + }, mustFloats(t, []float64{}, 0), 0, 0, 0); err == nil { + t.Fatalf("empty system: want an error, got %v", x) + } +} + +// TestGMRESUnconvergedErrors pins the contract shared with SpSolve and +// FindRootSystem: running out of cycles with the tolerance unmet is an +// error naming the residual, never a silent approximation. +func TestGMRESUnconvergedErrors(t *testing.T) { + // A = diag(1, 0.5): a legitimate but slow system for restart-1 + // GMRES, which crawls towards the solution and cannot reach a + // 1e-12 relative residual in 2 cycles. + diag := []float64{1, 0.5} + op := func(v *core.Array) (*core.Array, error) { + out := core.New(core.Float, 2) + for i := range 2 { + out.RawFloats()[i] = diag[i] * v.FloatAt(i) + } + return out, nil + } + if _, err := GMRES(op, mustFloats(t, []float64{1, 1}, 2), 1, 2, 1e-12); err == nil { + t.Fatal("unconverged GMRES: want an error") + } +} + +// TestNorm2F64LargeScale pins the overflow-safe norm: entries near +// 1e154 must not square their way to +Inf. +func TestNorm2F64LargeScale(t *testing.T) { + if got := norm2F64([]float64{1e154, 1e154}); math.IsInf(got, 1) || math.Abs(got-math.Sqrt2*1e154) > 1e-12*math.Sqrt2*1e154 { + t.Fatalf("norm2F64(1e154, 1e154) = %v, want ≈ %g", got, math.Sqrt2*1e154) + } + if got := norm2F64([]float64{0, 0}); got != 0 { + t.Fatalf("norm2F64(zeros) = %v, want 0", got) + } +} diff --git a/linalg/helpers.go b/linalg/helpers.go new file mode 100644 index 0000000..c70e800 --- /dev/null +++ b/linalg/helpers.go @@ -0,0 +1,59 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// scalar is the element type the generic kernels serve. +type scalar interface { + float64 | complex128 +} + +// absComplex returns |z|, the shared magnitude used for tolerances. +func absComplex(z complex128) float64 { return base.AbsComplex(z) } + +func absOf[T scalar](v T) float64 { + switch x := any(v).(type) { + case float64: + return math.Abs(x) + case complex128: + return cmplx.Abs(x) + } + return 0 +} + +// zeros allocates a zeroed array for internally derived shapes. The +// error is part of the signature on purpose: a caller whose shape is +// not provably valid must propagate the failure, not nil-deref. +func zeros(dt core.Dtype, shape []int) (*core.Array, error) { + return core.Zeros(dt, shape...) +} + +// finiteF64 reports whether v is neither NaN nor infinite. +func finiteF64(v float64) bool { return !math.IsNaN(v) && !math.IsInf(v, 0) } + +// vecFinite reports whether every entry of v is finite. +func vecFinite(v []float64) bool { + for _, x := range v { + if !finiteF64(x) { + return false + } + } + return true +} + +// isRealDtype reports whether a holds real (non-complex) numeric data. +func isRealDtype(a *core.Array) bool { + switch a.Dtype() { + case core.Float, core.Float32, core.Int: + return true + } + return false +} diff --git a/linalg/helpers_test.go b/linalg/helpers_test.go new file mode 100644 index 0000000..bd754cf --- /dev/null +++ b/linalg/helpers_test.go @@ -0,0 +1,64 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +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 +} + +// mustFromFloats builds a float array, failing the test on a bad shape. +func mustFromFloats(t *testing.T, vals []float64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, shape...) + if err != nil { + t.Fatalf("FromFloats(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustFromFloat32s builds a float32 array, failing the test on a bad shape. +func mustFromFloat32s(t *testing.T, vals []float32, shape ...int) *core.Array { + t.Helper() + a, err := core.FromFloat32s(vals, shape...) + if err != nil { + t.Fatalf("FromFloat32s(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustFromInts builds an int array, failing the test on a bad shape. +func mustFromInts(t *testing.T, vals []int64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromInts(vals, shape...) + if err != nil { + t.Fatalf("FromInts(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustFromComplexes builds a complex array, failing the test on a bad shape. +func mustFromComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array { + t.Helper() + a, err := core.FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err) + } + return a +} diff --git a/linalg/initialisers_sparse_test.go b/linalg/initialisers_sparse_test.go new file mode 100644 index 0000000..c752584 --- /dev/null +++ b/linalg/initialisers_sparse_test.go @@ -0,0 +1,165 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +func TestTruncatedNormal(t *testing.T) { + g := core.NewGenerator(42) + w := core.TruncatedNormal(g, []int{1000}, 0, 1) + if w.Len() != 1000 { + t.Errorf("TruncatedNormal len: %d", w.Len()) + } + // All values must be in [-2, 2]. + for i := range w.RawFloat32s() { + v := float64(w.RawFloat32s()[i]) + if v < -2 || v > 2 { + t.Errorf("TruncatedNormal [%d]: %v out of [-2, 2]", i, v) + } + } +} + +func TestZerosLikeOnesLike(t *testing.T) { + a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + z := core.ZerosLike(a) + if z.Shape()[0] != 2 || z.Shape()[1] != 2 { + t.Errorf("ZerosLike shape: %v", z.Shape()) + } + for i := range z.RawFloats() { + if z.RawFloats()[i] != 0 { + t.Errorf("ZerosLike [%d]: %v, want 0", i, z.RawFloats()[i]) + } + } + o := core.OnesLike(a) + for i := range o.RawFloats() { + if o.RawFloats()[i] != 1 { + t.Errorf("OnesLike [%d]: %v, want 1", i, o.RawFloats()[i]) + } + } + // Mutating the copy must not affect a. + z.RawFloats()[0] = 99 + if a.RawFloats()[0] == 99 { + t.Error("ZerosLike: mutation leaked into source array") + } +} + +func TestSparseRoundTrip(t *testing.T) { + dense := mustFromFloats(t, []float64{ + 0, 1, 0, + 2, 0, 3, + 0, 0, 4, + }, 3, 3) + s, err := core.SparseFrom(dense) + if err != nil { + t.Fatal(err) + } + if s.NNZ() != 4 { + t.Errorf("NNZ: %d, want 4", s.NNZ()) + } + back, err := s.Dense() + if err != nil { + t.Fatal(err) + } + for i := range 9 { + v, _ := core.FloatAt(back, i/3, i%3) + orig, _ := core.FloatAt(dense, i/3, i%3) + if v != orig { + t.Errorf("sparse round-trip [%d]: got %v, want %v", i, v, orig) + } + } +} + +func TestSparseFromKeepsDtype(t *testing.T) { + // SparseFrom on an int array must keep the int dtype through the + // round trip (the values array is part of the identity). + di := mustFromInts(t, []int64{0, 5, 0, 7}, 2, 2) + si, err := core.SparseFrom(di) + if err != nil { + t.Fatal(err) + } + if si.Values.Dtype() != core.Int { + t.Fatalf("SparseFrom int: values dtype %s, want int", si.Values.Dtype()) + } + back, err := si.Dense() + if err != nil { + t.Fatal(err) + } + if !core.Equal(back, di) { + t.Errorf("int sparse round-trip: got %s, want %s", back, di) + } + + df := mustFromFloat32s(t, []float32{0, 1.5, 0, 2.5}, 2, 2) + sf, err := core.SparseFrom(df) + if err != nil { + t.Fatal(err) + } + if sf.Values.Dtype() != core.Float32 { + t.Fatalf("SparseFrom float32: values dtype %s, want float32", sf.Values.Dtype()) + } +} + +func TestSparseMul(t *testing.T) { + dense := mustFromFloats(t, []float64{ + 0, 1, 0, + 2, 0, 3, + 0, 0, 4, + }, 3, 3) + s, _ := core.SparseFrom(dense) + mul, err := core.FromFloats([]float64{10, 20, 30, 40, 50, 60, 70, 80, 90}, 3, 3) + if err != nil { + t.Fatal(err) + } + out, err := core.SpMul(s, mul) + if err != nil { + t.Fatal(err) + } + // Only non-zero positions are filled. + // dense[0,1]=1 * mul[0,1]=20 = 20 + // dense[1,0]=2 * mul[1,0]=40 = 80 + // dense[1,2]=3 * mul[1,2]=60 = 180 + // dense[2,2]=4 * mul[2,2]=90 = 360 + expect := []float64{0, 20, 0, 80, 0, 180, 0, 0, 360} + for i, w := range expect { + v, _ := core.FloatAt(out, i/3, i%3) + if v != w { + t.Errorf("SpMul [%d]: got %v, want %v", i, v, w) + } + } +} + +func TestSparseMatMul(t *testing.T) { + // Sparse 2×3 times dense 3×2. + indices, err := core.FromInts([]int64{ + 0, 0, + 1, 2, + }, 2, 2) + if err != nil { + t.Fatal(err) + } + values, err := core.FromFloats([]float64{1, 2}, 2) + if err != nil { + t.Fatal(err) + } + s := &core.SparseCOO{Indices: indices, Values: values, Shape: []int{2, 3}} + dense, _ := core.FromFloats([]float64{ + 1, 2, + 3, 4, + 5, 6, + }, 3, 2) + out, err := core.SpMatMul(s, dense) + if err != nil { + t.Fatal(err) + } + // Row 0: [1, 0, 0] · [[1,2],[3,4],[5,6]] = [1, 2] + // Row 1: [0, 0, 2] · [[1,2],[3,4],[5,6]] = [10, 12] + for i, w := range []float64{1, 2, 10, 12} { + v, _ := core.FloatAt(out, i/2, i%2) + if v != w { + t.Errorf("SpMatMul [%d]: got %v, want %v", i, v, w) + } + } +} diff --git a/linalg/least_squares_test.go b/linalg/least_squares_test.go new file mode 100644 index 0000000..e847be7 --- /dev/null +++ b/linalg/least_squares_test.go @@ -0,0 +1,190 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/big" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The solve applies the Householder reflectors to the right-hand sides +// instead of forming Q, so its forward error is compared here against a +// minimiser solved exactly in rational arithmetic: the reference is not +// another floating-point route, it is the answer. + +// exactMinimiser returns the minimiser of ||Ax − b|| for an integer A and +// b, solved from the normal equations in exact rational arithmetic, along +// with the exact minimal residual norm. +func exactMinimiser(t *testing.T, a [][]float64, b []float64) ([]*big.Rat, float64) { + t.Helper() + m, n := len(a), len(a[0]) + g := make([][]*big.Rat, n) + for i := range n { + g[i] = make([]*big.Rat, n+1) + for j := range n + 1 { + g[i][j] = new(big.Rat) + } + } + exact := func(v float64) *big.Rat { + r := new(big.Rat).SetFloat64(v) + if r == nil { + t.Fatalf("value %g is not representable as a rational", v) + } + return r + } + for i := range n { + for j := range n { + s := new(big.Rat) + for r := range m { + s.Add(s, new(big.Rat).Mul(exact(a[r][i]), exact(a[r][j]))) + } + g[i][j] = s + } + s := new(big.Rat) + for r := range m { + s.Add(s, new(big.Rat).Mul(exact(a[r][i]), exact(b[r]))) + } + g[i][n] = s + } + for c := range n { + p := c + for r := c + 1; r < n; r++ { + if new(big.Rat).Abs(g[r][c]).Cmp(new(big.Rat).Abs(g[p][c])) > 0 { + p = r + } + } + g[c], g[p] = g[p], g[c] + for r := c + 1; r < n; r++ { + f := new(big.Rat).Quo(g[r][c], g[c][c]) + for k := c; k <= n; k++ { + g[r][k].Sub(g[r][k], new(big.Rat).Mul(f, g[c][k])) + } + } + } + x := make([]*big.Rat, n) + for i := n - 1; i >= 0; i-- { + s := new(big.Rat).Set(g[i][n]) + for j := i + 1; j < n; j++ { + s.Sub(s, new(big.Rat).Mul(g[i][j], x[j])) + } + x[i] = s.Quo(s, g[i][i]) + } + res := new(big.Rat) + for r := range m { + s := new(big.Rat) + for j := range n { + s.Add(s, new(big.Rat).Mul(exact(a[r][j]), x[j])) + } + s.Sub(s, exact(b[r])) + res.Add(res, new(big.Rat).Mul(s, s)) + } + rf, _ := new(big.Float).Sqrt(new(big.Float).SetRat(res)).Float64() + return x, rf +} + +// leastSquaresFixture builds an integer A with the given column scaling +// and a right-hand side that is A·x0 plus a small integer perturbation, so +// the exact minimiser is known and the residual is not zero. +func leastSquaresFixture(m, n int, scaleJ func(int) float64) ([][]float64, []float64) { + a := make([][]float64, m) + s := uint64(12345) + for i := range m { + a[i] = make([]float64, n) + for j := range n { + s = s*6364136223846793005 + 1442695040888963407 + a[i][j] = float64((s>>40)%9+1) * scaleJ(j) + } + } + x0 := make([]float64, n) + for j := range n { + x0[j] = float64((j*7)%5 - 2) + } + b := make([]float64, m) + for i := range m { + s = s*6364136223846793005 + 1442695040888963407 + v := 0.0 + for j := range n { + v += a[i][j] * x0[j] + } + b[i] = v + float64(int((s>>50)%3)-1) + } + return a, b +} + +func TestLeastSquaresAgainstRationalMinimiser(t *testing.T) { + cases := []struct { + name string + m, n int + scal func(int) float64 + }{ + {"512x32", 512, 32, func(int) float64 { return 1 }}, + {"1024x48", 1024, 48, func(int) float64 { return 1 }}, + {"512x32-scaled", 512, 32, func(j int) float64 { return math.Pow(0.5, float64(j)) }}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + a, b := leastSquaresFixture(c.m, c.n, c.scal) + xr, resExact := exactMinimiser(t, a, b) + am, err := core.FromFloats(flattenRows(a), c.m, c.n) + if err != nil { + t.Fatal(err) + } + bm, err := core.FromFloats(b, c.m, 1) + if err != nil { + t.Fatal(err) + } + sol, err := LeastSquares(am, bm) + if err != nil { + t.Fatalf("LeastSquares: %v", err) + } + x := sol.RawFloats() + worst, norm := 0.0, 0.0 + for i := range c.n { + ref, _ := xr[i].Float64() + norm = math.Max(norm, math.Abs(ref)) + worst = math.Max(worst, math.Abs(x[i]-ref)) + } + rel := worst / norm + // The housekeeping floor: a backward-stable solve of a problem + // this well conditioned lands within a few ulps per element, + // and the scaled case carries a condition number near 2^31, so + // the bound is the conditioning times the unit roundoff with a + // generous constant. + kond := 1.0 + if c.scal(1) != 1 { + kond = math.Pow(2, float64(c.n)) + } + bound := 64 * math.Max(1, kond) * 2.220446049250313e-16 + if rel > bound { + t.Errorf("relative solution error %g exceeds the bound %g", rel, bound) + } + // The residual must not be worse than the exact minimal one by + // more than the same order. + res := 0.0 + for i := range c.m { + v := b[i] + for j := range c.n { + v -= a[i][j] * x[j] + } + res += v * v + } + res = math.Sqrt(res) + if res > resExact*(1+1e-8) { + t.Errorf("residual %g exceeds the exact minimal residual %g", res, resExact) + } + t.Logf("relative solution error %.3e, residual %.6g (exact minimal %.6g)", rel, res, resExact) + }) + } +} + +func flattenRows(a [][]float64) []float64 { + out := make([]float64, 0, len(a)*len(a[0])) + for _, row := range a { + out = append(out, row...) + } + return out +} diff --git a/linalg/linalg.go b/linalg/linalg.go new file mode 100644 index 0000000..1bdc4ed --- /dev/null +++ b/linalg/linalg.go @@ -0,0 +1,343 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Linear algebra. Every entry point requires a square +// 2-D matrix. One LU decomposition with partial pivoting backs Det, +// Solve and Inv; the same kernel serves both real and complex element +// types through generics, so complex systems solve exactly where their +// real counterparts do. A singular matrix is an error for Solve and Inv, +// while the determinants return 0. Results are float64/complex128 +// approximations, as is standard for elimination with partial pivoting; +// near-singular matrices lose precision without failing. +// +// Type summary per the ladder: Solve and Inv promote to complex128 when +// any operand is complex; Det and Trace answer only real matrices and +// point complex callers at DetComplex and TraceComplex. + +// squareFloatMatrix converts a real array to a square float64 matrix, +// promoting ints and float32 on the way; it errors with the caller's +// name when the shape is wrong. The rows are views over one flat +// backing slice, so the conversion costs two allocations. The rows are +// handed to the solver, which may reorder the row headers in place, so +// callers must treat the result as consumed. +func squareFloatMatrix(a *core.Array, name string) ([][]float64, error) { + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + back := make([]float64, n*n) + if a.Dtype() == core.Float && !a.Strided() && len(a.RawFloats()) == n*n { + copy(back, a.RawFloats()) + } else { + for i := range n * n { + back[i] = a.FloatAt(i) + } + } + rows := make([][]float64, n) + for i := range n { + rows[i] = back[i*n : (i+1)*n] + } + return rows, nil +} + +// squareComplexMatrix converts a complex array to a square matrix of +// complex128 values with the same shape contract as its real twin. +func squareComplexMatrix(a *core.Array, name string) ([][]complex128, error) { + if a.Dtype() != core.Complex { + return nil, base.Errf("%s: needs a complex matrix", name) + } + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + back := make([]complex128, n*n) + if !a.Strided() && len(a.RawComplexes()) == n*n { + copy(back, a.RawComplexes()) + } else { + for i := range n * n { + back[i] = a.ComplexAt(i) + } + } + rows := make([][]complex128, n) + for i := range n { + rows[i] = back[i*n : (i+1)*n] + } + return rows, nil +} + +// identityColumns builds the columns of the n×n identity, the standard +// right-hand side of the inverse. The columns slice into one flat +// backing array: identical values, two allocations instead of n+1. +func identityColumns[T scalar](n int) [][]T { + back := make([]T, n*n) + cols := make([][]T, n) + var one T = 1 + for j := range n { + col := back[j*n : (j+1)*n] + col[j] = one + cols[j] = col + } + return cols +} + +// requireReal refuses the complex dtype at an entry point that reads +// its inputs as real numbers. Without the gate the first FloatAt would +// fall through to the complex array's empty int payload and index out +// of range, which is a panic rather than the error the caller expects. +func requireReal(name string, arrays ...*core.Array) error { + for _, a := range arrays { + if a.Dtype() == core.Complex { + return base.Errf("%s: complex arrays are not supported", name) + } + } + return nil +} + +// Det returns the determinant of a square real matrix (ints and +// float32 promote). A singular matrix yields 0, not an error; complex +// matrices are answered by DetComplex. +func Det(a *core.Array) (float64, error) { + if a.Dtype() == core.Complex { + return 0, base.Errf("Det: complex matrices are answered by DetComplex") + } + m, err := squareFloatMatrix(a, "Det") + if err != nil { + return 0, err + } + _, parity := base.Factor(m) + det := 1.0 * float64(parity) + for i := range m { + det *= m[i][i] + } + return det, nil +} + +// DetComplex returns the determinant of a square complex matrix as a +// complex128 value. A singular matrix yields 0, not an error. +func DetComplex(a *core.Array) (complex128, error) { + m, err := squareComplexMatrix(a, "DetComplex") + if err != nil { + return 0, err + } + _, parity := base.Factor(m) + det := complex(float64(parity), 0) + for i := range m { + det *= m[i][i] + } + return det, nil +} + +// solveRHS converts b into float columns of length n. b may be a vector +// of length n or a matrix with n rows. +func solveRHS(b *core.Array, n int) ([][]float64, error) { + switch { + case b.NDim() == 1 && b.Len() == n: + col := make([]float64, n) + if b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == n { + copy(col, b.RawFloats()) + } else { + for i := range n { + col[i] = b.FloatAt(i) + } + } + return [][]float64{col}, nil + case b.NDim() == 2 && b.Shape()[0] == n: + w := b.Shape()[1] + cols := make([][]float64, w) + dense := b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == n*w + for j := range w { + col := make([]float64, n) + for i := range n { + if dense { + col[i] = b.RawFloats()[i*w+j] + } else { + col[i] = b.FloatAt(i*w + j) + } + } + cols[j] = col + } + return cols, nil + } + return nil, base.Errf("Solve: b must be a vector of length %d or a matrix with %d rows, got shape %s", + n, n, base.ShapeText(b.Shape())) +} + +// complexRHS converts b into complex columns of length n under the same +// shape contract as solveRHS; a real right-hand side promotes wholesale. +func complexRHS(b *core.Array, n int) ([][]complex128, error) { + switch { + case b.NDim() == 1 && b.Len() == n: + col := make([]complex128, n) + if b.Dtype() == core.Complex && !b.Strided() && len(b.RawComplexes()) == n { + copy(col, b.RawComplexes()) + } else { + for i := range n { + col[i] = b.ComplexAt(i) + } + } + return [][]complex128{col}, nil + case b.NDim() == 2 && b.Shape()[0] == n: + w := b.Shape()[1] + cols := make([][]complex128, w) + dense := b.Dtype() == core.Complex && !b.Strided() && len(b.RawComplexes()) == n*w + for j := range w { + col := make([]complex128, n) + for i := range n { + if dense { + col[i] = b.RawComplexes()[i*w+j] + } else { + col[i] = b.ComplexAt(i*w + j) + } + } + cols[j] = col + } + return cols, nil + } + return nil, base.Errf("Solve: b must be a vector of length %d or a matrix with %d rows, got shape %s", + n, n, base.ShapeText(b.Shape())) +} + +// columnsToFloatArray assembles solved columns back into a row-major +// array shaped like the original right-hand side: like[0] rows in both +// the vector and the matrix case. +func columnsToFloatArray(cols [][]float64, like []int) *core.Array { + rows := like[0] + out := core.New(core.Float, like...) + for i := range rows { + base := i * len(cols) + for c, col := range cols { + out.RawFloats()[base+c] = col[i] + } + } + return out +} + +// columnsToComplexArray is the complex twin of columnsToFloatArray. +func columnsToComplexArray(cols [][]complex128, like []int) *core.Array { + out := core.New(core.Complex, like...) + rows := like[0] + for r := range rows { + for c, col := range cols { + out.RawComplexes()[r*len(cols)+c] = col[r] + } + } + return out +} + +// Solve returns x such that a·x = b, where b is a vector of length n or +// a matrix with n rows. Real inputs give a float result; any complex +// operand promotes the whole system to complex128 per the ladder. +func Solve(a, b *core.Array) (*core.Array, error) { + if a.Dtype() == core.Complex || b.Dtype() == core.Complex { + // Any complex operand promotes the whole system: a real A + // widens alongside a complex right-hand side. + if a.Dtype() != core.Complex { + conv, cerr := core.Astype(a, core.Complex) + if cerr != nil { + return nil, cerr + } + a = conv + } + ac, err := squareComplexMatrix(a, "Solve") + if err != nil { + return nil, err + } + rhs, err := complexRHS(b, len(ac)) + if err != nil { + return nil, err + } + x, err := base.SolveSystem("Solve", ac, rhs) + if err != nil { + return nil, err + } + return columnsToComplexArray(x, b.Shape()), nil + } + m, err := squareFloatMatrix(a, "Solve") + if err != nil { + return nil, err + } + rhs, err := solveRHS(b, len(m)) + if err != nil { + return nil, err + } + x, err := base.SolveSystem("Solve", m, rhs) + if err != nil { + return nil, err + } + return columnsToFloatArray(x, b.Shape()), nil +} + +// Inv returns the inverse of a square matrix. Real inputs give a float +// array; a complex matrix answers in complex128. +func Inv(a *core.Array) (*core.Array, error) { + if a.Dtype() == core.Complex { + ac, err := squareComplexMatrix(a, "Inv") + if err != nil { + return nil, err + } + x, err := base.SolveSystem("Inv", ac, identityColumns[complex128](len(ac))) + if err != nil { + return nil, err + } + return columnsToComplexArray(x, []int{len(ac), len(ac)}), nil + } + m, err := squareFloatMatrix(a, "Inv") + if err != nil { + return nil, err + } + x, err := base.SolveSystem("Inv", m, identityColumns[float64](len(m))) + if err != nil { + return nil, err + } + return columnsToFloatArray(x, []int{len(m), len(m)}), nil +} + +// FitPolynomial fits the coefficients (lowest power first) of a +// degree-th polynomial through (x, y) samples by QR least squares over +// the Vandermonde matrix, which stays stable where the normal +// equations would square the condition number. +func FitPolynomial(x, y *core.Array, degree int) (*core.Array, error) { + n := x.Len() + if err := requireReal("FitPolynomial", x, y); err != nil { + return nil, err + } + if degree < 0 { + return nil, base.Errf("FitPolynomial: degree %d is negative", degree) + } + // degree+1 would wrap at the int maximum and slip past the sample + // count check below as a negative. + if degree == math.MaxInt { + return nil, base.Errf("FitPolynomial: degree %d is too large", degree) + } + if n < degree+1 || y.Len() != n { + return nil, base.Errf("FitPolynomial: need n ≥ degree+1 matching samples") + } + aMat, err := zeros(core.Float, []int{n, degree + 1}) + if err != nil { + return nil, base.Errf("FitPolynomial: %w", err) + } + for r := range n { + pow := 1.0 + for d := range degree + 1 { + aMat.SetFloatAt(r*(degree+1)+d, pow) + pow *= x.FloatAt(r) + } + } + yFlat, err := zeros(core.Float, []int{n}) + if err != nil { + return nil, base.Errf("FitPolynomial: %w", err) + } + for i := range n { + yFlat.SetFloatAt(i, y.FloatAt(i)) + } + return LeastSquares(aMat, yFlat) +} diff --git a/linalg/linalg_complex_test.go b/linalg/linalg_complex_test.go new file mode 100644 index 0000000..09f97e7 --- /dev/null +++ b/linalg/linalg_complex_test.go @@ -0,0 +1,261 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +// complexMatrix builds a complex128 array from row-major values. +func complexMatrix(t *testing.T, shape []int, vals ...complex128) *core.Array { + t.Helper() + total := 1 + for _, d := range shape { + total *= d + } + if len(vals) != total { + t.Fatalf("value count %d does not fill %v", len(vals), shape) + } + a, err := core.ComplexFromArray(vals, shape...) + if err != nil { + t.Fatalf("ComplexFromArray(%v, %v): %v", vals, shape, err) + } + return a +} + +func TestDetComplexRotation(t *testing.T) { + // A rotation in the complex plane has unit determinant; scaling one + // row by (2+i) multiplies the determinant by the same factor. + angle := complex(0.6, -0.8) // unit magnitude + m := complexMatrix(t, []int{2, 2}, + angle, 0, + 0, 1, + ) + det, err := DetComplex(m) + if err != nil { + t.Fatal(err) + } + // The complex diagonal keeps its phase, so the unit-modulus claim + // holds on the magnitude only. + if math.Abs(cmplx.Abs(det)-1) > 1e-12 { + t.Fatalf("unitary |det| = %v", cmplx.Abs(det)) + } + + scaled := complexMatrix(t, []int{2, 2}, + 2+1i, 0, + 3-4i, angle, + ) + det2, err := DetComplex(scaled) + if err != nil { + t.Fatal(err) + } + want := (2 + 1i) * angle + if cmplx.Abs(det2-want) > 1e-12 { + t.Fatalf("scaled det = %v, want %v", det2, want) + } + + // A singular matrix yields zero without erroring. + singular := complexMatrix(t, []int{2, 2}, 1, 2i, 2, 4i) + detS, err := DetComplex(singular) + if err != nil { + t.Fatal(err) + } + if detS != 0 { + t.Fatalf("singular det = %v, want 0", detS) + } + + // The real Det keeps pointing complex callers at the complex twin. + realShaped, _ := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2) + if _, err := Det(realShaped); err != nil { + t.Fatalf("real det errored: %v", err) + } + if _, err := Det(complexMatrix(t, []int{2, 2}, 1, 0, 0, 1)); err == nil { + t.Fatal("real Det accepted a complex matrix") + } +} + +func TestSolveInvComplexRoundTrip(t *testing.T) { + a := complexMatrix(t, []int{3, 3}, + 2+1i, 0, 1-1i, + 0, 1-2i, 0.5, + 1i, 1, 3+2i, + ) + + inv, err := Inv(a) + if err != nil { + t.Fatal(err) + } + if inv.Dtype() != core.Complex { + t.Fatalf("inverse dtype %v", inv.Dtype()) + } + product, err := core.MatMul2D(a, inv) + if err != nil { + t.Fatal(err) + } + for i := range 9 { + row, col := i/3, i%3 + want := complex(0, 0) + if row == col { + want = 1 + } + if got := product.RawComplexes()[i]; cmplx.Abs(got-want) > 1e-10 { + t.Fatalf("A·A⁻¹[%d][%d] = %v, want %v", row, col, got, want) + } + } + + // Solve and verify by substitution for a vector right-hand side. + b, _ := core.FromFloats([]float64{1, 2, 3}, 3) // real b promotes to complex + x, err := Solve(a, b) + if err != nil { + t.Fatal(err) + } + check, err := core.MatMul2D(a, x) + if err != nil { + t.Fatal(err) + } + for i := range 3 { + if got := check.RawComplexes()[i]; cmplx.Abs(got-complex(float64(i+1), 0)) > 1e-10 { + t.Fatalf("(a·x)[%d] = %v, want %v", i, got, float64(i+1)) + } + } + + // Solving with an explicitly complex vector works symmetrically. + bc := complexMatrix(t, []int{3}, 2+2i, 0, -1i) + bcol, err := core.Reshape(bc, 3) + if err != nil { + t.Fatal(err) + } + xc, err := Solve(a, bcol) + if err != nil { + t.Fatal(err) + } + if xc.Dtype() != core.Complex || xc.Len() != 3 { + t.Fatalf("complex solve shape/dtype: %v %v", xc.Shape(), xc.Dtype()) + } + + // Singular systems error instead of returning garbage. + singular := complexMatrix(t, []int{2, 2}, 1, 1i, 2, 2i) + bad, _ := core.FromFloats([]float64{1, 1}, 2) + if _, err := Solve(singular, bad); err == nil { + t.Fatal("singular solve succeeded") + } + if _, err := Inv(singular); err == nil { + t.Fatal("singular inverse succeeded") + } +} + +func TestKronComplexBlocks(t *testing.T) { + a := complexMatrix(t, []int{2, 2}, 1+1i, 0, 0, 2-1i) + b := complexMatrix(t, []int{2, 2}, 1, 2i, 3, 0) + + out, err := core.Kron(a, b) + if err != nil { + t.Fatal(err) + } + if out.Dtype() != core.Complex { + t.Fatalf("kron dtype %v", out.Dtype()) + } + if got := out.Shape(); got[0] != 4 || got[1] != 4 { + t.Fatalf("kron shape %v", got) + } + // Top-left block scales b by (1+1i); it sits on flat positions + // k*4+l because blocks interleave in the outer product layout. + type slot struct { + idx int + want complex128 + } + for _, s := range []slot{ + {0, 1 + 1i}, + {1, 2i * (1 + 1i)}, + {4, 3 * (1 + 1i)}, + {5, 0}, + } { + if got := out.RawComplexes()[s.idx]; got != s.want { + t.Fatalf("block TL[%d] = %v, want %v", s.idx, got, s.want) + } + } + // Bottom-right block scales b by (2−1i); identity via mixed pair. + mixed, _ := core.FromFloats([]float64{1}, 1, 1) // real 1×1 + eye := complexMatrix(t, []int{1, 1}, 2-1i) + cross, err := core.Kron(mixed, eye) + if err != nil { + t.Fatal(err) + } + if cross.RawComplexes()[0] != 2-1i { + t.Fatalf("mixed kron element = %v", cross.RawComplexes()[0]) + } +} + +func TestTraceComplexSum(t *testing.T) { + m := complexMatrix(t, []int{2, 2}, 1+2i, 99, 99, 3-4i) + s, err := core.TraceComplex(m) + if err != nil { + t.Fatal(err) + } + if s != 4-2i { + t.Fatalf("trace = %v, want 4−2i", s) + } + nonSquare := complexMatrix(t, []int{1, 2}, 1, 2) + if _, err := core.TraceComplex(nonSquare); err == nil { + t.Fatal("non-square trace accepted") + } + // Real path keeps rejecting complexes by name. + realM, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := core.TraceComplex(realM); err == nil { + t.Fatal("TraceComplex accepted a real matrix") + } +} + +func TestPowIComplexExactness(t *testing.T) { + z := mustComplexes(t, []complex128{1 + 1i, 2 - 1i}, 2) + + cubed, err := core.PowI(z, 3) + if err != nil { + t.Fatal(err) + } + wantFirst := (1 + 1i) * (1 + 1i) * (1 + 1i) // −2+2i + if cubed.RawComplexes()[0] != wantFirst { + t.Fatalf("(1+i)^3 = %v, want %v", cubed.RawComplexes()[0], wantFirst) + } + inverse, err := core.PowI(z, -1) + if err != nil { + t.Fatal(err) + } + one, err := core.Mul(inverse, z) + if err != nil { + t.Fatal(err) + } + for i := range 2 { + if math.Abs(real(one.RawComplexes()[i])-1) > 1e-12 || math.Abs(imag(one.RawComplexes()[i])) > 1e-12 { + t.Fatalf("z·z⁻¹[%d] = %v", i, one.RawComplexes()[i]) + } + } +} + +// TestComplexSolveInvDetErrors moved with Solve, Inv and Det from the +// root package: Solve and Inv run on complex systems through the same +// LU kernel, while the real-only Det points at its complex twin. +func TestComplexSolveInvDetErrors(t *testing.T) { + sys := complexMatrix(t, []int{2, 2}, + 2+1i, 0, + 0, 3-1i, + ) + got, err := Inv(sys) + if err != nil { + t.Fatalf("Inv complex: %v", err) + } + if got.Dtype() != core.Complex { + t.Fatalf("Inv complex dtype %v", got.Dtype()) + } + if _, err := Solve(sys, sys); err != nil { + t.Fatalf("Solve complex: %v", err) + } + if _, err := Det(sys); err == nil || !strings.Contains(err.Error(), "DetComplex") { + t.Fatalf("Det complex: %v", err) + } +} diff --git a/linalg/linalg_test.go b/linalg/linalg_test.go new file mode 100644 index 0000000..ba21119 --- /dev/null +++ b/linalg/linalg_test.go @@ -0,0 +1,187 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +func approx(t *testing.T, name string, got, want float64) { + t.Helper() + if math.Abs(got-want) > 1e-9 { + t.Fatalf("%s: got %v, want %v", name, got, want) + } +} + +func TestDet(t *testing.T) { + m := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) + d, err := Det(m) + if err != nil { + t.Fatalf("Det: %v", err) + } + approx(t, "Det", d, -2) + + id, _ := core.Identity(core.Float, 3) + d, err = Det(id) + if err != nil { + t.Fatalf("Det identity: %v", err) + } + approx(t, "Det identity", d, 1) + + // A singular matrix yields 0, not an error. + singular := mustFromFloats(t, []float64{1, 2, 2, 4}, 2, 2) + d, err = Det(singular) + if err != nil { + t.Fatalf("Det singular: %v", err) + } + approx(t, "Det singular", d, 0) + + // int matrices convert. + im := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + d, err = Det(im) + if err != nil { + t.Fatalf("Det int: %v", err) + } + approx(t, "Det int", d, -2) + + nonSquare := mustFromFloats(t, []float64{1, 2, 3}, 1, 3) + if _, err := Det(nonSquare); err == nil || !strings.Contains(err.Error(), "square 2-D matrix") { + t.Fatalf("Det shape: %v", err) + } +} + +func TestSolve(t *testing.T) { + // 3x + y = 9, x + 2y = 8 so x = 2, y = 3. + a := mustFromFloats(t, []float64{3, 1, 1, 2}, 2, 2) + b := mustFromFloats(t, []float64{9, 8}, 2) + + x, err := Solve(a, b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + if x.Dtype() != core.Float || x.Shape()[0] != 2 { + t.Fatalf("Solve shape: %s", x) + } + v0, _ := core.FloatAt(x, 0) + v1, _ := core.FloatAt(x, 1) + approx(t, "Solve x", v0, 2) + approx(t, "Solve y", v1, 3) + + // Matrix right-hand side solves column by column: rows [9,1] and [8,2] + // give the columns [9,8] and [1,2]. + bm := mustFromFloats(t, []float64{9, 1, 8, 2}, 2, 2) + xm, err := Solve(a, bm) + if err != nil { + t.Fatalf("Solve matrix: %v", err) + } + // Second column: 3x+y=1, x+2y=2 so x=0, y=1. + c00, _ := core.FloatAt(xm, 0, 0) + c01, _ := core.FloatAt(xm, 0, 1) + c10, _ := core.FloatAt(xm, 1, 0) + c11, _ := core.FloatAt(xm, 1, 1) + approx(t, "Solve matrix 00", c00, 2) + approx(t, "Solve matrix 01", c01, 0) + approx(t, "Solve matrix 10", c10, 3) + approx(t, "Solve matrix 11", c11, 1) + + // int operands convert on both sides. + ia := mustFromInts(t, []int64{3, 1, 1, 2}, 2, 2) + ib := mustFromInts(t, []int64{9, 8}, 2) + ix, err := Solve(ia, ib) + if err != nil || ix.Dtype() != core.Float { + t.Fatalf("Solve int: %s %v", ix, err) + } + iv, _ := core.FloatAt(ix, 0) + approx(t, "Solve int x", iv, 2) + + singular := mustFromFloats(t, []float64{1, 2, 2, 4}, 2, 2) + if _, err := Solve(singular, b); err == nil || !strings.Contains(err.Error(), "singular") { + t.Fatalf("Solve singular: %v", err) + } + + wrong := mustFromFloats(t, []float64{1, 2, 3}, 3) + if _, err := Solve(a, wrong); err == nil || !strings.Contains(err.Error(), "b must be") { + t.Fatalf("Solve b shape: %v", err) + } + if _, err := Solve(wrong, a); err == nil || !strings.Contains(err.Error(), "square 2-D matrix") { + t.Fatalf("Solve shape: %v", err) + } +} + +func TestInv(t *testing.T) { + m := mustFromFloats(t, []float64{4, 7, 2, 6}, 2, 2) + inv, err := Inv(m) + if err != nil { + t.Fatalf("Inv: %v", err) + } + // 1/10 · [[6, -7], [-2, 4]] + v00, _ := core.FloatAt(inv, 0, 0) + v01, _ := core.FloatAt(inv, 0, 1) + v10, _ := core.FloatAt(inv, 1, 0) + v11, _ := core.FloatAt(inv, 1, 1) + approx(t, "Inv 00", v00, 0.6) + approx(t, "Inv 01", v01, -0.7) + approx(t, "Inv 10", v10, -0.2) + approx(t, "Inv 11", v11, 0.4) + + // A·A⁻¹ is the identity. + prod, err := core.MatMul2D(m, inv) + if err != nil { + t.Fatalf("Inv check: %v", err) + } + for i := range 2 { + for j := range 2 { + want := 0.0 + if i == j { + want = 1 + } + got, _ := core.FloatAt(prod, i, j) + approx(t, "Inv product", got, want) + } + } + + // A permutation matrix exercises the pivot path. + p := mustFromFloats(t, []float64{0, 1, 1, 0}, 2, 2) + pinv, err := Inv(p) + if err != nil { + t.Fatalf("Inv permutation: %v", err) + } + g00, _ := core.FloatAt(pinv, 0, 0) + g01, _ := core.FloatAt(pinv, 0, 1) + approx(t, "Inv permutation 00", g00, 0) + approx(t, "Inv permutation 01", g01, 1) + + singular := mustFromFloats(t, []float64{1, 2, 2, 4}, 2, 2) + if _, err := Inv(singular); err == nil || !strings.Contains(err.Error(), "singular") { + t.Fatalf("Inv singular: %v", err) + } + if _, err := Inv(mustFromFloats(t, []float64{1, 2, 3}, 1, 3)); err == nil { + t.Fatalf("Inv shape must error") + } +} + +// TestSolveRealMatrixComplexRHS pins the promotion contract: a real +// system with a complex right-hand side promotes the whole solve to +// complex128 instead of erroring. +func TestSolveRealMatrixComplexRHS(t *testing.T) { + a := mustFromFloats(t, []float64{2, 0, 0, 4}, 2, 2) + b, _ := core.FromComplexes([]complex128{2 + 4i, 8}, 2) + x, err := Solve(a, b) + if err != nil { + t.Fatalf("Solve real x complex: %v", err) + } + if x.Dtype() != core.Complex { + t.Fatalf("Solve promote dtype: %s", x.Dtype()) + } + // A is diagonal: x = b / diag = [1+2i, 2]. + want := []complex128{1 + 2i, 2} + for i := range want { + if v, _ := core.ComplexAt(x, i); v != want[i] { + t.Fatalf("Solve [%d]: got %v, want %v", i, v, want[i]) + } + } +} diff --git a/linalg/mat_scratch_test.go b/linalg/mat_scratch_test.go new file mode 100644 index 0000000..10e54d0 --- /dev/null +++ b/linalg/mat_scratch_test.go @@ -0,0 +1,40 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// TestMatMul2DFloat32ScratchCleaned pins the regression where the +// pooled accumulation buffer was handed back dirty: leftover values +// from an earlier borrower used to leak into every float32 product sum, +// making results depend on allocation history instead of inputs alone. +func TestMatMul2DFloat32ScratchCleaned(t *testing.T) { + const size = 35 + warm := engine.GetFloat64Buf(size) + for i := range warm { + warm[i] = 1.5 // leave obvious residue in the pooled window + } + engine.PutFloat64Buf(warm) + + a := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 2, 2) + b := mustFromFloat32s(t, []float32{0.5, 0, 2, 1}, 2, 2) + want := []float32{4.5, 2, 9.5, 4} + + for run := range 2 { + got, err := core.MatMul2D(a, b) + if err != nil { + t.Fatal(err) + } + for i := range want { + if g := got.RawFloat32s()[i]; g != want[i] { + t.Fatalf("run %d: c[%d] = %v, want %v", run, i, g, want[i]) + } + } + } +} diff --git a/linalg/mat_test.go b/linalg/mat_test.go new file mode 100644 index 0000000..272a0a6 --- /dev/null +++ b/linalg/mat_test.go @@ -0,0 +1,248 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +func TestIdentity(t *testing.T) { + id, err := core.Identity(core.Int, 2) + if err != nil { + t.Fatalf("Identity: %v", err) + } + want := mustFromInts(t, []int64{1, 0, 0, 1}, 2, 2) + if !core.Equal(want, id) { + t.Fatalf("Identity int: %s", id) + } + fid, err := core.Identity(core.Float, 3) + if err != nil { + t.Fatalf("Identity float: %v", err) + } + if fid.Dtype() != core.Float { + t.Fatalf("Identity dtype: %s", fid.Dtype()) + } + if v, _ := core.FloatAt(fid, 2, 2); v != 1 { + t.Fatalf("Identity corner: %v", v) + } + if _, err := core.Identity(core.Int, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") { + t.Fatalf("Identity negative: %v", err) + } +} + +func TestMatMul2D(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + b := mustFromInts(t, []int64{7, 8, 9, 10, 11, 12}, 3, 2) + + c, err := core.MatMul2D(a, b) + if err != nil { + t.Fatalf("MatMul: %v", err) + } + if c.Dtype() != core.Int { + t.Fatalf("MatMul dtype: %s", c.Dtype()) + } + want := mustFromInts(t, []int64{58, 64, 139, 154}, 2, 2) + if !core.Equal(want, c) { + t.Fatalf("MatMul: %s", c) + } + + // Promotion and float values. + f := mustFromFloats(t, []float64{0.5, 1.5, 2.5}, 1, 3) + fc, err := core.MatMul2D(f, b) + if err != nil { + t.Fatalf("MatMul float: %v", err) + } + if fc.Dtype() != core.Float { + t.Fatalf("MatMul promote: %s", fc.Dtype()) + } + // 0.5*7 + 1.5*9 + 2.5*11 = 44.5, 0.5*8 + 1.5*10 + 2.5*12 = 49 + if v, _ := core.FloatAt(fc, 0, 0); v != 44.5 { + t.Fatalf("MatMul float value: %v", v) + } + if v, _ := core.FloatAt(fc, 0, 1); v != 49 { + t.Fatalf("MatMul float value: %v", v) + } +} + +// TestMatMul2DFloat32MixedInt pins the mixed float32 x int product: an +// core.Int right operand used to reach the float32 kernel and slice b's nil +// floats32 payload, panicking instead of promoting. +func TestMatMul2DFloat32MixedInt(t *testing.T) { + a := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 2, 2) + b := mustFromInts(t, []int64{5, 6, 7, 8}, 2, 2) + + c, err := core.MatMul2D(a, b) + if err != nil { + t.Fatalf("MatMul float32 x int: %v", err) + } + if c.Dtype() != core.Float32 { + t.Fatalf("MatMul float32 x int dtype: %s", c.Dtype()) + } + // [1*5+2*7, 1*6+2*8, 3*5+4*7, 3*6+4*8] = [19, 22, 43, 50] + want := mustFromFloat32s(t, []float32{19, 22, 43, 50}, 2, 2) + if !core.Equal(want, c) { + t.Fatalf("MatMul float32 x int: %s", c) + } + + // The mirrored int x float32 product keeps working. + d, err := core.MatMul2D(b, a) + if err != nil { + t.Fatalf("MatMul int x float32: %v", err) + } + wantD := mustFromFloat32s(t, []float32{23, 34, 31, 46}, 2, 2) + if !core.Equal(wantD, d) { + t.Fatalf("MatMul int x float32: %s", d) + } +} + +func TestMatMulVectorShapes(t *testing.T) { + m := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + v := mustFromInts(t, []int64{5, 6}, 2) + + mv, err := core.MatMul2D(m, v) + if err != nil { + t.Fatalf("matrix × vector: %v", err) + } + // [1*5+2*6, 3*5+4*6] = [17, 39] + want := mustFromInts(t, []int64{17, 39}, 2) + if !core.Equal(want, mv) { + t.Fatalf("matrix × vector: %s", mv) + } + + vm, err := core.MatMul2D(v, m) + if err != nil { + t.Fatalf("vector × matrix: %v", err) + } + // [5*1+6*3, 5*2+6*4] = [23, 34] + wantVM := mustFromInts(t, []int64{23, 34}, 2) + if !core.Equal(wantVM, vm) { + t.Fatalf("vector × matrix: %s", vm) + } +} + +func TestMatMulErrors(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + wrong2D := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 3, 2) + if _, err := core.MatMul2D(a, wrong2D); err == nil || !strings.Contains(err.Error(), "inner dimensions must agree") { + t.Fatalf("inner mismatch: %v", err) + } + wrongV := mustFromInts(t, []int64{1, 2, 3}, 3) + if _, err := core.MatMul2D(a, wrongV); err == nil || !strings.Contains(err.Error(), "inner dimensions must agree") { + t.Fatalf("vector mismatch: %v", err) + } + if _, err := core.MatMul2D(wrongV, a); err == nil || !strings.Contains(err.Error(), "inner dimensions must agree") { + t.Fatalf("vector × matrix mismatch: %v", err) + } + v := mustFromInts(t, []int64{1}, 1) + if _, err := core.MatMul2D(v, v); err == nil || !strings.Contains(err.Error(), "unsupported shapes") { + t.Fatalf("1-D × 1-D is Dot: %v", err) + } + c3 := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2) + if _, err := core.MatMul2D(c3, c3); err == nil || !strings.Contains(err.Error(), "unsupported shapes") { + t.Fatalf("3-D matmul: %v", err) + } +} + +func TestTranspose(t *testing.T) { + m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + tt := core.Transpose(m) + want := mustFromInts(t, []int64{1, 4, 2, 5, 3, 6}, 3, 2) + if !core.Equal(want, tt) { + t.Fatalf("Transpose: %s", tt) + } + + // 1-D transposition is a copy. + v := mustFromInts(t, []int64{1, 2}, 2) + if !core.Equal(v, core.Transpose(v)) { + t.Fatalf("Transpose 1-D: %s", core.Transpose(v)) + } + + // 3-D reverses all dimensions. + c := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2) + ct := core.Transpose(c) + if shape := ct.Shape(); shape[0] != 2 || shape[1] != 2 || shape[2] != 2 { + t.Fatalf("Transpose 3-D shape: %v", shape) + } + if v, _ := core.IntAt(ct, 0, 0, 1); v != 5 { + t.Fatalf("Transpose 3-D value: %d", v) + } +} + +func TestReshape(t *testing.T) { + a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + r, err := core.Reshape(a, 3, 2) + if err != nil { + t.Fatalf("Reshape: %v", err) + } + want := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 3, 2) + if !core.Equal(want, r) { + t.Fatalf("Reshape: %s", r) + } + flat, err := core.Reshape(a, 6) + if err != nil || flat.Len() != 6 { + t.Fatalf("Reshape flat: %s %v", flat, err) + } + if _, err := core.Reshape(a, 4); err == nil || !strings.Contains(err.Error(), "do not fill the shape") { + t.Fatalf("Reshape count: %v", err) + } + if _, err := core.Reshape(a); err == nil || !strings.Contains(err.Error(), "at least one dimension") { + t.Fatalf("Reshape empty: %v", err) + } +} + +func TestRowCol(t *testing.T) { + m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) + + row, err := core.Row(m, 1) + if err != nil { + t.Fatalf("Row: %v", err) + } + wantRow := mustFromInts(t, []int64{4, 5, 6}, 3) + if !core.Equal(wantRow, row) { + t.Fatalf("Row: %s", row) + } + + col, err := core.Col(m, 1) + if err != nil { + t.Fatalf("Col: %v", err) + } + wantCol := mustFromInts(t, []int64{2, 5}, 2) + if !core.Equal(wantCol, col) { + t.Fatalf("Col: %s", col) + } + + if _, err := core.Row(m, 2); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("Row range: %v", err) + } + if _, err := core.Col(m, 3); err == nil || !strings.Contains(err.Error(), "out of range") { + t.Fatalf("Col range: %v", err) + } + v := mustFromInts(t, []int64{1}, 1) + if _, err := core.Row(v, 0); err == nil || !strings.Contains(err.Error(), "needs a 2-D array") { + t.Fatalf("Row 1-D: %v", err) + } + if _, err := core.Col(v, 0); err == nil || !strings.Contains(err.Error(), "needs a 2-D array") { + t.Fatalf("Col 1-D: %v", err) + } +} + +// TestKronIntExact pins the int Kronecker product: products above 2^53 +// used to round-trip through float64 and lose their low bits. +func TestKronIntExact(t *testing.T) { + big := int64(1) << 53 + a := mustFromInts(t, []int64{big + 1}, 1, 1) + b := mustFromInts(t, []int64{2}, 1, 1) + got, err := core.Kron(a, b) + if err != nil { + t.Fatal(err) + } + if got.Dtype() != core.Int { + t.Fatalf("Kron int dtype: %s", got.Dtype()) + } + if v, _ := core.IntAt(got, 0, 0); v != (big+1)*2 { + t.Fatalf("Kron int exact: got %d, want %d", v, (big+1)*2) + } +} diff --git a/linalg/matrixfunc.go b/linalg/matrixfunc.go new file mode 100644 index 0000000..b17234a --- /dev/null +++ b/linalg/matrixfunc.go @@ -0,0 +1,403 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Matrix functions through the eigen-decomposition. MatrixSqrt and +// MatrixLog evaluate √A and ln(A) for a symmetric positive +// semi-definite real matrix by diagonalising: A = VΛVᵀ, so +// f(A) = V·f(Λ)·Vᵀ. Every other input, a Hermitian complex matrix +// included, runs through the complex Schur decomposition described at +// the individual functions. The positive semi-definite requirement +// matters: the eigenvalues of such a matrix are non-negative, which +// makes the square root and the logarithm real. A matrix with a +// significantly negative eigenvalue is refused rather than silently +// producing complex or NaN entries. + +// spectrumRelTol is the relative tolerance MatrixSqrt and the Schur +// routes of both matrix functions use to judge their spectrum: an +// eigenvalue counts as non-positive or singular when it sits within +// this fraction of the spectrum's magnitude. Purely relative on +// purpose: an absolute floor would refuse legitimate small-scale +// inputs such as diag(1e-12, 2e-12). The symmetric route of MatrixLog +// judges negativity against the rounding floor n·eps·λmax instead, +// which admits far smaller positive eigenvalues; see MatrixLog. +const spectrumRelTol = 1e-10 + +// MatrixSqrt returns the principal square root √A of a symmetric +// positive semi-definite real matrix, by diagonalisation: A = VΛVᵀ, +// so √A = V·√Λ·Vᵀ. Eigenvalues below −tol·‖A‖ are refused (the +// matrix is not positive semi-definite); small negative values from +// rounding are clamped to zero. +// For nonsymmetric matrices the answer runs through the complex Schur +// decomposition: A = Q·T·Qᴴ reduces the function to the triangular +// recurrence on T, the Schur-Parlett route. A real matrix answers in +// real entries whenever the principal function is real, which for the +// square root means no negative real eigenvalues and for the +// logarithm none on the non-positive real axis; otherwise the call is +// an error, honestly, rather than a complex answer in real clothing. +func MatrixSqrt(a *core.Array) (*core.Array, error) { + n := a.Shape()[0] + if a.NDim() != 2 || n != a.Shape()[1] { + return nil, base.Errf("MatrixSqrt: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + if !isSymmetricMatrix(a) { + return matrixSqrtSchur(a) + } + vals, vecs, err := Eigen(a) + if err != nil { + return nil, base.Errf("MatrixSqrt: %w", err) + } + // Eigen returns dense float arrays, so the reconstruction reads the + // payloads directly rather than paying floatAt's dispatch per entry + // of this O(n³) triple product. + scale := 0.0 + for _, lam := range vals.RawFloats() { + if v := math.Abs(lam); v > scale { + scale = v + } + } + tol := spectrumRelTol * scale + sqrtVals := make([]float64, n) + for i, lam := range vals.RawFloats() { + if lam < -tol { + return nil, base.Errf("MatrixSqrt: eigenvalue %v is negative; the matrix is not positive semi-definite", lam) + } + if lam < 0 { + lam = 0 + } + sqrtVals[i] = math.Sqrt(lam) + } + // √A = V·diag(√λ)·Vᵀ. + out := core.New(core.Float, []int{n, n}...) + vMat := vecs.RawFloats() + for i := range n { + for j := range n { + s := 0.0 + for k := range n { + s += vMat[i*n+k] * sqrtVals[k] * vMat[j*n+k] + } + out.RawFloats()[i*n+j] = s + } + } + return out, nil +} + +// MatrixLog returns the principal matrix logarithm ln(A) of a real +// symmetric positive-definite matrix, by the same diagonalisation +// route as MatrixSqrt with ln applied to the eigenvalues. A negative +// eigenvalue beyond the rounding floor n·eps·λmax is refused as +// non-positive, and so is an exact zero: the principal logarithm of a +// matrix with eigenvalues on the non-positive real axis leaves the +// reals. A positive eigenvalue is logged whatever its size against +// the spectrum, since ln 1e-300 is finite and exact to its own +// digits. +func MatrixLog(a *core.Array) (*core.Array, error) { + n := a.Shape()[0] + if a.NDim() != 2 || n != a.Shape()[1] { + return nil, base.Errf("MatrixLog: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape())) + } + if !isSymmetricMatrix(a) { + return matrixLogSchur(a) + } + vals, vecs, err := Eigen(a) + if err != nil { + return nil, base.Errf("MatrixLog: %w", err) + } + // The negativity floor is the eigensolver's own rounding scale, + // n·eps·λmax, the floor the QR rank guard uses: a negative value + // beyond it is a genuine sign of a non-positive spectrum, one + // inside it is rounding of a zero, and a positive value is logged + // honestly however far below the spectrum's top it sits. + floor := float64(n) * base.EpsF * eigenMaxAbs(vals, n) + logVals := make([]float64, n) + for i, lam := range vals.RawFloats() { + if lam < -floor { + return nil, base.Errf("MatrixLog: eigenvalue %v is negative; the matrix is not positive definite", lam) + } + if lam <= 0 { + return nil, base.Errf("MatrixLog: eigenvalue %v is non-positive; the principal logarithm needs a positive-definite matrix", lam) + } + logVals[i] = math.Log(lam) + } + out := core.New(core.Float, []int{n, n}...) + vMat := vecs.RawFloats() + for i := range n { + for j := range n { + s := 0.0 + for k := range n { + s += vMat[i*n+k] * logVals[k] * vMat[j*n+k] + } + out.RawFloats()[i*n+j] = s + } + } + return out, nil +} + +// eigenMaxAbs returns the largest absolute eigenvalue. +func eigenMaxAbs(vals *core.Array, n int) float64 { + m := 0.0 + for _, v := range vals.RawFloats() { + if v := math.Abs(v); v > m { + m = v + } + } + return m +} + +// isSymmetricMatrix reports whether a square real matrix is symmetric +// to rounding, which routes the matrix functions to the cheaper +// symmetric eigendecomposition. The tolerance is purely relative to +// the matrix scale, deliberately without an absolute floor: a floor +// would call a small-scale matrix symmetric on asymmetry it carries in +// full relative measure, and the Eigen route then applies its own +// relative check and refuses the matrix the routing just approved. +// Relative-only keeps the two decisions in agreement, because any +// asymmetry this check passes is within the fraction of the scale the +// route's own guard tolerates. +func isSymmetricMatrix(a *core.Array) bool { + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] || a.Dtype() == core.Complex { + return false + } + n := a.Shape()[0] + scale := 0.0 + for i := range n * n { + scale = math.Max(scale, math.Abs(a.FloatAt(i))) + } + tol := 1e-12 * scale + for i := range n { + for j := i + 1; j < n; j++ { + if math.Abs(a.FloatAt(i*n+j)-a.FloatAt(j*n+i)) > tol { + return false + } + } + } + return true +} + +// matrixSqrtSchur evaluates the principal square root through the +// complex Schur decomposition and the triangular recurrence. A real +// matrix with a negative real eigenvalue has a complex principal root +// and is refused; a real answer is otherwise recovered to rounding. +func matrixSqrtSchur(a *core.Array) (*core.Array, error) { + const name = "MatrixSqrt" + n := a.Shape()[0] + t, q, err := SchurComplex(a) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + tm := make([]complex128, n*n) + copy(tm, t.RawComplexes()) + if a.Dtype() != core.Complex { + tol := spectrumRelTol * cmplxAbsMax(tm, n) + for i := range n { + lam := tm[i*n+i] + if math.Abs(imag(lam)) <= tol && real(lam) < -tol { + return nil, base.Errf("%s: eigenvalue %v is negative real; the principal square root is complex", name, lam) + } + } + } + if err := triSqrt(tm, n); err != nil { + return nil, base.Errf("%s: %w", name, err) + } + return schurBackTransform(q, tm, n, a.Dtype() != core.Complex, name) +} + +// matrixLogSchur evaluates the principal logarithm through the complex +// Schur decomposition by inverse scaling and squaring: repeated +// triangular square roots walk the matrix near the identity, where the +// Mercator series finishes the job, and the squarings unwind. A real +// matrix with an eigenvalue on the non-positive real axis has no real +// principal logarithm and is refused. +func matrixLogSchur(a *core.Array) (*core.Array, error) { + const name = "MatrixLog" + n := a.Shape()[0] + t, q, err := SchurComplex(a) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + tm := make([]complex128, n*n) + copy(tm, t.RawComplexes()) + logScale := 0.0 + // Spectrum screen. Real input refuses eigenvalues on the + // non-positive real axis, where the principal logarithm leaves the + // reals; complex input refuses eigenvalues at or below the + // rounding floor of the spectrum, where the square-root walk below + // would diverge into garbage instead of converging. + tol := spectrumRelTol * cmplxAbsMax(tm, n) + for i := range n { + lam := tm[i*n+i] + if a.Dtype() != core.Complex { + if math.Abs(imag(lam)) <= tol && real(lam) <= tol { + return nil, base.Errf("%s: eigenvalue %v is on the non-positive real axis; the principal logarithm is complex", name, lam) + } + } else if cmplx.Abs(lam) <= tol { + return nil, base.Errf("%s: eigenvalue %v is at the rounding floor of the spectrum; the matrix is singular and has no principal logarithm", name, lam) + } + } + // Scale the spectrum towards one so the square-root walk converges + // in a handful of steps; the scalar shift unwinds afterwards. + c := 0.0 + for i := range n { + c = math.Max(c, cmplx.Abs(tm[i*n+i])) + } + if c > 0 && (c < 0.5 || c > 2) { + logScale = math.Log(c) + for i := range n { + for j := range n { + tm[i*n+j] /= complex(c, 0) + } + } + } + // Square-root walk towards the identity. + squarings := 0 + for range 60 { + dev := triDeviation(tm, n) + if dev <= 0.25 { + break + } + if err := triSqrt(tm, n); err != nil { + return nil, base.Errf("%s: %w", name, err) + } + squarings++ + } + // A spectrum the walk cannot bring near the identity (typically a + // near-singular one the screen let through) would send the Mercator + // series below into open-ended divergence; refuse instead. + if dev := triDeviation(tm, n); dev > 0.25 { + return nil, base.Errf("%s: the square-root walk did not converge in 60 passes (deviation %.3g); the matrix is outside the domain of the principal logarithm", name, dev) + } + // Mercator series for log(I + E) on the near-identity triangle. + e := make([]complex128, n*n) + for i := range n { + for j := range n { + e[i*n+j] = tm[i*n+j] + } + e[i*n+i] -= 1 + } + lg := make([]complex128, n*n) + // The series multiplies term by e each step, so the two buffers + // alternate instead of the product landing in a fresh matrix: every + // value read is the previous step's product exactly as it was when + // the product was copied back, because only the upper triangle is + // ever written and both buffers start with a zero lower triangle. + term, next := make([]complex128, n*n), make([]complex128, n*n) + copy(term, e) + for m := 1; m <= 40; m++ { + factor := complex(1/float64(m), 0) + if m%2 == 0 { + factor = complex(-1/float64(m), 0) + } + for i := range n * n { + lg[i] += factor * term[i] + } + if m == 40 { + break + } + for i := range n { + for j := i; j < n; j++ { + s := complex(0, 0) + for k := i; k <= j; k++ { + s += term[i*n+k] * e[k*n+j] + } + next[i*n+j] = s + } + } + tnorm := 0.0 + for i := range n * n { + tnorm = math.Max(tnorm, cmplx.Abs(next[i])) + } + term, next = next, term + if tnorm <= 1e-17 { + break + } + } + // Unwind the walk: log(T) = 2^k·L + log(c)·I. + for i := range n * n { + lg[i] *= complex(math.Pow(2, float64(squarings)), 0) + } + if logScale != 0 { + for i := range n { + lg[i*n+i] += complex(logScale, 0) + } + } + return schurBackTransform(q, lg, n, a.Dtype() != core.Complex, name) +} + +// schurBackTransform maps the triangular result back through the Schur +// unitary, q·S·qᴴ, and returns real entries for real input when the +// imaginary rounding cancels. +func schurBackTransform(q *core.Array, s []complex128, n int, wantReal bool, name string) (*core.Array, error) { + qm := q.RawComplexes() + w := make([]complex128, n*n) + for i := range n { + for j := range n { + sum := complex(0, 0) + for k := i; k < n; k++ { + sum += s[i*n+k] * cmplx.Conj(qm[j*n+k]) + } + w[i*n+j] = sum + } + } + out := make([]complex128, n*n) + for i := range n { + for j := range n { + sum := complex(0, 0) + for k := range n { + sum += qm[i*n+k] * w[k*n+j] + } + out[i*n+j] = sum + } + } + if !wantReal { + res, _ := core.FromComplexes(out, n, n) + return res, nil + } + worstIm, worstRe := 0.0, 0.0 + for i := range n * n { + worstIm = math.Max(worstIm, math.Abs(imag(out[i]))) + worstRe = math.Max(worstRe, math.Abs(real(out[i]))) + } + if worstIm > 1e-8*math.Max(1, worstRe) { + return nil, base.Errf("%s: the result is complex (imaginary magnitude %.3g); the real answer does not exist", name, worstIm) + } + res := core.New(core.Float, []int{n, n}...) + for i := range n * n { + res.RawFloats()[i] = real(out[i]) + } + return res, nil +} + +// cmplxAbsMax returns the largest diagonal magnitude of a complex +// triangular matrix. +func cmplxAbsMax(t []complex128, n int) float64 { + m := 0.0 + for i := range n { + m = math.Max(m, cmplx.Abs(t[i*n+i])) + } + return m +} + +// triDeviation measures how far an upper triangular matrix still sits +// from the identity: the largest magnitude among the diagonal offsets +// and the strict upper part, the convergence gauge of the square-root +// walk in matrixLogSchur. +func triDeviation(t []complex128, n int) float64 { + dev := 0.0 + for i := range n { + dev = math.Max(dev, cmplx.Abs(t[i*n+i]-1)) + } + for i := range n { + for j := i + 1; j < n; j++ { + dev = math.Max(dev, cmplx.Abs(t[i*n+j])) + } + } + return dev +} diff --git a/linalg/matrixfunc_test.go b/linalg/matrixfunc_test.go new file mode 100644 index 0000000..533e6a4 --- /dev/null +++ b/linalg/matrixfunc_test.go @@ -0,0 +1,125 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +func TestMatrixSqrtLog(t *testing.T) { + // A = [[2,1],[1,2]] has eigenvalues 3 and 1, both positive. + a := mustFloats(t, []float64{2, 1, 1, 2}, 2, 2) + sq, err := MatrixSqrt(a) + if err != nil { + t.Fatalf("MatrixSqrt: %v", err) + } + // √A·√A = A. + prod, _ := core.MatMul2D(sq, sq) + for i := range 4 { + if math.Abs(prod.FloatAt(i)-a.FloatAt(i)) > 1e-10 { + t.Fatalf("√A·√A[%d] = %v, want %v", i, prod.FloatAt(i), a.FloatAt(i)) + } + } + lg, err := MatrixLog(a) + if err != nil { + t.Fatalf("MatrixLog: %v", err) + } + // e^{ln A} = A: exponentiate the log element-wise and rebuild. + // Simpler: check ln A against the eigen-decomposition directly. + evals, evecs, _ := Eigen(a) + lnA := make([]float64, 4) + for i := range 2 { + for j := range 2 { + s := 0.0 + for k := range 2 { + s += evecs.FloatAt(i*2+k) * math.Log(evals.FloatAt(k)) * evecs.FloatAt(j*2+k) + } + lnA[i*2+j] = s + } + } + for i := range 4 { + if math.Abs(lg.FloatAt(i)-lnA[i]) > 1e-10 { + t.Fatalf("ln A[%d] = %v, want %v", i, lg.FloatAt(i), lnA[i]) + } + } + // Negative eigenvalue refused. + neg := mustFloats(t, []float64{1, 2, 2, 1}, 2, 2) + if _, err := MatrixSqrt(neg); err == nil { + t.Fatal("expected an error for a non-PSD matrix") + } + if _, err := MatrixLog(neg); err == nil { + t.Fatal("expected an error for a matrix with a non-positive eigenvalue") + } +} + +// TestMatrixFunctionsTinyScale pins the purely relative spectrum +// tolerance: a tiny positive-definite matrix must pass the +// positive-definiteness screen and an indefinite one of the same scale +// must be refused, where the old max(1, scale) floors did the +// opposite. +func TestMatrixFunctionsTinyScale(t *testing.T) { + // MatrixLog of diag(1e-12, 2e-12): ln of the eigenvalues. + tiny := mustFloats(t, []float64{1e-12, 0, 0, 2e-12}, 2, 2) + lg, err := MatrixLog(tiny) + if err != nil { + t.Fatalf("MatrixLog(diag(1e-12, 2e-12)): %v", err) + } + want := []float64{math.Log(1e-12), math.Log(2e-12)} + for i := range 2 { + if math.Abs(lg.FloatAt(i*2+i)-want[i]) > 1e-9 { + t.Fatalf("ln A[%d][%d] = %.16g, want %.16g", i, i, lg.FloatAt(i*2+i), want[i]) + } + } + for _, i := range []int{1, 2} { + if math.Abs(lg.FloatAt(i)) > 1e-12 { + t.Fatalf("ln A off-diagonal [%d] = %v, want 0", i, lg.FloatAt(i)) + } + } + // MatrixSqrt of the same matrix: √ of the eigenvalues. + sq, err := MatrixSqrt(tiny) + if err != nil { + t.Fatalf("MatrixSqrt(diag(1e-12, 2e-12)): %v", err) + } + for i, w := range []float64{1e-6, math.Sqrt(2) * 1e-6} { + if math.Abs(sq.FloatAt(i*2+i)-w) > 1e-15 { + t.Fatalf("√A[%d][%d] = %.16g, want %.16g", i, i, sq.FloatAt(i*2+i), w) + } + } + // An indefinite matrix at the same tiny scale is refused, not + // silently clamped. + indefinite := mustFloats(t, []float64{-1e-12, 0, 0, 2e-12}, 2, 2) + if _, err := MatrixSqrt(indefinite); err == nil { + t.Fatal("MatrixSqrt of a tiny indefinite matrix: want an error") + } + if _, err := MatrixLog(indefinite); err == nil { + t.Fatal("MatrixLog of a tiny indefinite matrix: want an error") + } +} + +// TestMatrixLogComplex pins the complex route of the principal +// logarithm: a nonsingular complex matrix logs eigenvalue-wise, while +// a singular one is refused instead of walking into diverging-series +// garbage. +func TestMatrixLogComplex(t *testing.T) { + diag := core.New(core.Complex, 2, 2) + diag.RawComplexes()[0] = 3 + diag.RawComplexes()[3] = 4 + lg, err := MatrixLog(diag) + if err != nil { + t.Fatalf("MatrixLog(complex diag(3, 4)): %v", err) + } + want := []float64{math.Log(3), math.Log(4)} + for i := range 2 { + if v := lg.ComplexAt(i*2 + i); math.Abs(real(v)-want[i]) > 1e-10 || math.Abs(imag(v)) > 1e-10 { + t.Fatalf("ln A[%d][%d] = %v, want ≈ %v", i, i, v, want[i]) + } + } + singular := core.New(core.Complex, 2, 2) + singular.RawComplexes()[0] = 1 + if _, err := MatrixLog(singular); err == nil { + t.Fatal("MatrixLog of a singular complex matrix: want an error") + } +} diff --git a/linalg/nonfinite_refusal_pins_test.go b/linalg/nonfinite_refusal_pins_test.go new file mode 100644 index 0000000..f39ea82 --- /dev/null +++ b/linalg/nonfinite_refusal_pins_test.go @@ -0,0 +1,197 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression pins for non-finite and counting guards: the matrix +// exponential refuses a non-finite entry, the minimum-degree ordering +// matches its brute-force reference, the sparse LU non-zero counts +// follow the unit triangle, and the ILU intake refuses overflow. + +// TestMatrixExpRejectsNonFinite pins the loud refusal for a +// non-finite entry, which the theta ladder used to read through an +// implementation-defined conversion into an all-NaN answer. +func TestMatrixExpRejectsNonFinite(t *testing.T) { + a := mustF(t, []float64{math.Inf(1), 0, 0, 1}, 2, 2) + if _, err := MatrixExp(a); err == nil { + t.Fatal("expected an error for an Inf entry") + } + b := mustF(t, []float64{math.NaN(), 0, 0, 1}, 2, 2) + if _, err := MatrixExp(b); err == nil { + t.Fatal("expected an error for a NaN entry") + } + c, _ := core.FromComplexes([]complex128{complex(math.Inf(1), 0), 0, 0, 1}, 2, 2) + if _, err := MatrixExp(c); err == nil { + t.Fatal("expected an error for an Inf complex entry") + } +} + +// TestMinimumDegreeMixedPattern pins the ordering against a +// brute-force minimum-degree reference on a mixed-degree pattern, +// where the degree-sorted adjacency lists and the index-sorted set +// union used to disagree and corrupt the elimination. The sample is +// the measured failing case: the pre-fix code eliminated vertex 10 +// before 1 and swapped the tail of the order. +func TestMinimumDegreeMixedPattern(t *testing.T) { + rows := []int{0, 0, 0, 0, 1, 1, 1, 1, 2, 4, 5, 5, 6, 6, 7} + cols := []int{1, 2, 5, 8, 3, 5, 6, 10, 6, 6, 9, 10, 7, 10, 10} + const n = 11 + coo := edgesCOO(t, rows, cols, n) + csc, err := CSCFromCOO(coo) + if err != nil { + t.Fatalf("CSCFromCOO: %v", err) + } + got, err := minimumDegree(csc) + if err != nil { + t.Fatalf("minimumDegree: %v", err) + } + // Brute-force reference: repeatedly eliminate the uneliminated + // vertex with the fewest uneliminated neighbours (ties to the + // smaller index), unioning neighbourhoods exactly. + adj := map[int]map[int]bool{} + addEdge := func(i, j int) { + if adj[i] == nil { + adj[i] = map[int]bool{} + } + if adj[j] == nil { + adj[j] = map[int]bool{} + } + adj[i][j], adj[j][i] = true, true + } + for e := range rows { + addEdge(rows[e], cols[e]) + } + eliminated := map[int]bool{} + var want []int + for range n { + best, bestDeg := -1, math.MaxInt + for v := range n { + if eliminated[v] { + continue + } + d := 0 + for u := range adj[v] { + if !eliminated[u] { + d++ + } + } + if d < bestDeg { + best, bestDeg = v, d + } + } + want = append(want, best) + eliminated[best] = true + nb := map[int]bool{} + for u := range adj[best] { + if !eliminated[u] { + nb[u] = true + } + } + for u := range nb { + for w := range nb { + if u != w { + adj[u][w] = true + } + } + delete(adj[u], best) + } + } + for i := range n { + if got[i] != want[i] { + t.Fatalf("order[%d] = %d, want %d (full %v vs %v)", i, got[i], want[i], got, want) + } + } +} + +// TestSparseLUNNZCountsUTriangle pins that NNZ includes U's +// strict triangle, which the column walk used to miss. +func TestSparseLUNNZCountsUTriangle(t *testing.T) { + // A tridiagonal matrix: L holds the subdiagonal, U the diagonal and + // the superdiagonal, so the factor stores exactly 3n - 2 entries. + const n = 8 + idx := make([]int64, 0, 6*n) + vals := make([]float64, 0, 3*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 2) + if i+1 < n { + add(i, i+1, -1) + add(i+1, i, -1) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + f, err := NewSparseLU(coo) + if err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + if want := 3*n - 2; f.NNZ() != want { + t.Fatalf("NNZ = %d, want %d (L strict + U strict + diagonal)", f.NNZ(), want) + } +} + +// TestSparseILUOverflowRefused pins the overflow refusal on +// finite input, mirroring the LU and Cholesky guards. +func TestSparseILUOverflowRefused(t *testing.T) { + idx := make([]int64, 0, 6) + vals := make([]float64, 0, 3) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + add(0, 0, 1e-200) + add(0, 1, 1e100) + add(1, 0, 1e100) + add(1, 1, 1e200) + indices, err := core.FromInts(idx, 4, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{4}), []int{2, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseILU(coo); err == nil { + t.Fatal("expected an overflow error from the elimination") + } +} + +// edgesCOO builds a symmetric-pattern COO from edge lists. +func edgesCOO(t *testing.T, rows, cols []int, n int) *core.SparseCOO { + t.Helper() + idx := make([]int64, 0, 2*len(rows)) + vals := make([]float64, 0, 2*len(rows)) + add := func(r, c int) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, 1) + } + for e := range rows { + add(rows[e], cols[e]) + add(cols[e], rows[e]) + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} diff --git a/linalg/perf_bench_test.go b/linalg/perf_bench_test.go new file mode 100644 index 0000000..d000a43 --- /dev/null +++ b/linalg/perf_bench_test.go @@ -0,0 +1,132 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" + + core "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Dense solver benchmarks guard the factorisation kernels. Inputs are +// diagonally dominant so every path stays on the well-conditioned side +// and the timings measure the kernel, not pivoting churn. + +func benchMatrix(b *testing.B, n int, spd bool) *core.Array { + b.Helper() + v := make([]float64, n*n) + for i := range n { + for j := range n { + v[i*n+j] = float64((i*7+j*13)%11) - 5 + } + v[i*n+i] += float64(n) + if spd { + // Symmetrise, then push the diagonal past the row sum so + // positive definiteness is guaranteed by strict diagonal + // dominance. + for j := i + 1; j < n; j++ { + avg := (v[i*n+j] + v[j*n+i]) / 2 + v[i*n+j], v[j*n+i] = avg, avg + } + row := float64(0) + for j := range n { + if j != i { + row += math.Abs(v[i*n+j]) + } + } + v[i*n+i] += row + } + } + a, err := core.FromFloats(v, n, n) + if err != nil { + b.Fatal(err) + } + return a +} + +func benchRHS(b *testing.B, n int) *core.Array { + b.Helper() + v := make([]float64, n) + for i := range v { + v[i] = float64(i%9) - 4 + } + x, err := core.FromFloats(v, n) + if err != nil { + b.Fatal(err) + } + return x +} + +func BenchmarkSolve64(b *testing.B) { + a := benchMatrix(b, 64, false) + x := benchRHS(b, 64) + b.ReportAllocs() + for b.Loop() { + if _, err := Solve(a, x); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSolve256(b *testing.B) { + a := benchMatrix(b, 256, false) + x := benchRHS(b, 256) + b.ReportAllocs() + for b.Loop() { + if _, err := Solve(a, x); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkInv256(b *testing.B) { + a := benchMatrix(b, 256, false) + b.ReportAllocs() + for b.Loop() { + if _, err := Inv(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCholesky256(b *testing.B) { + a := benchMatrix(b, 256, true) + b.ReportAllocs() + for b.Loop() { + if _, err := Cholesky(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkQR256(b *testing.B) { + a := benchMatrix(b, 256, false) + b.ReportAllocs() + for b.Loop() { + if _, _, err := QR(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSVD128(b *testing.B) { + a := benchMatrix(b, 128, true) + b.ReportAllocs() + for b.Loop() { + if _, _, _, err := SVD(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkDet256(b *testing.B) { + a := benchMatrix(b, 256, false) + b.ReportAllocs() + for b.Loop() { + if _, err := Det(a); err != nil { + b.Fatal(err) + } + } +} diff --git a/linalg/pipeline.go b/linalg/pipeline.go new file mode 100644 index 0000000..f83f72e --- /dev/null +++ b/linalg/pipeline.go @@ -0,0 +1,393 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Pipeline is a thin eager-evaluation wrapper around the package-level +// functions. Each step allocates a new array; the pipeline is +// just sugar over the same functions. It exists so callers can read +// sequential transformations as a single expression without losing the +// package-level style: +// +// out, err := Pipe(a). +// AddF(1). +// Sqrt(). +// Result() +// +// The pipeline does not provide automatic error propagation: a step +// that returns an error is recorded, and subsequent steps become +// no-ops. Call Result() to retrieve both the final array and the first +// error observed (or nil). +// +// There is no lazy graph, no autograd, no backprop: just a chain of +// eager calls with the error short-circuit on top. + +// Pipeline holds the current array under transformation and the first +// error encountered during the chain. +type Pipeline struct { + current *core.Array + err error +} + +// Pipe starts a pipeline from a. +func Pipe(a *core.Array) *Pipeline { + return &Pipeline{current: a} +} + +// Result returns the current array and the first error observed during +// the chain. If an error occurred earlier, current is the array that +// triggered it. +func (p *Pipeline) Result() (*core.Array, error) { + return p.current, p.err +} + +// shortCircuit returns true when the pipeline already has an error: +// the caller becomes a no-op so subsequent steps stay composable. +func (p *Pipeline) shortCircuit() bool { + return p.err != nil +} + +// assign updates the pipeline's current array under the wrapping rule +// (errors short-circuit, otherwise the new array is stored). +func (p *Pipeline) assign(out *core.Array, err error) { + if p.err != nil { + return + } + if err != nil { + p.err = err + return + } + p.current = out +} + +// Binary array-array operations. +func (p *Pipeline) Add(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Add(p.current, b) + p.assign(out, err) + return p +} + +func (p *Pipeline) Sub(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Sub(p.current, b) + p.assign(out, err) + return p +} + +func (p *Pipeline) Mul(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Mul(p.current, b) + p.assign(out, err) + return p +} + +func (p *Pipeline) Div(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Div(p.current, b) + p.assign(out, err) + return p +} + +// Scalar arithmetic (the integer flavours only: float scalar variants +// would shadow Go's float64 promotion, so callers reach for AddF etc. +// explicitly). +func (p *Pipeline) AddI(v int64) *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.AddI(p.current, v), nil) + return p +} + +func (p *Pipeline) SubI(v int64) *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.SubI(p.current, v), nil) + return p +} + +func (p *Pipeline) MulI(v int64) *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.MulI(p.current, v), nil) + return p +} + +func (p *Pipeline) AddF(v float64) *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.AddF(p.current, v), nil) + return p +} + +func (p *Pipeline) SubF(v float64) *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.SubF(p.current, v), nil) + return p +} + +func (p *Pipeline) MulF(v float64) *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.MulF(p.current, v), nil) + return p +} + +// Reshape and transpose. +func (p *Pipeline) Reshape(shape ...int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Reshape(p.current, shape...) + p.assign(out, err) + return p +} + +func (p *Pipeline) Transpose() *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.Transpose(p.current), nil) + return p +} + +func (p *Pipeline) TransposeAxes(dims ...int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.TransposeAxes(p.current, dims...) + p.assign(out, err) + return p +} + +func (p *Pipeline) Flatten(startDim, endDim int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Flatten(p.current, startDim, endDim) + p.assign(out, err) + return p +} + +func (p *Pipeline) Squeeze(dim int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Squeeze(p.current, dim) + p.assign(out, err) + return p +} + +func (p *Pipeline) Unsqueeze(dim int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Unsqueeze(p.current, dim) + p.assign(out, err) + return p +} + +// Reductions and extrema. +func (p *Pipeline) Maximum(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Maximum(p.current, b) + p.assign(out, err) + return p +} + +func (p *Pipeline) Minimum(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Minimum(p.current, b) + p.assign(out, err) + return p +} + +func (p *Pipeline) SumAxis(dim int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.SumAxis(p.current, dim) + p.assign(out, err) + return p +} + +func (p *Pipeline) MeanAxis(dim int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.MeanAxis(p.current, dim) + p.assign(out, err) + return p +} + +func (p *Pipeline) ClipI(lo, hi int64) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.ClipI(p.current, lo, hi) + p.assign(out, err) + return p +} + +func (p *Pipeline) ClipF(lo, hi float64) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.ClipF(p.current, lo, hi) + p.assign(out, err) + return p +} + +// Activations. +func (p *Pipeline) Tanh() *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Tanh(p.current) + p.assign(out, err) + return p +} + +func (p *Pipeline) Sigmoid() *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Sigmoid(p.current) + p.assign(out, err) + return p +} + +// Matrix and linear algebra. +func (p *Pipeline) MatMul2D(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.MatMul2D(p.current, b) + p.assign(out, err) + return p +} + +func (p *Pipeline) Inv() *Pipeline { + if p.shortCircuit() { + return p + } + out, err := Inv(p.current) + p.assign(out, err) + return p +} + +func (p *Pipeline) Solve(b *core.Array) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := Solve(p.current, b) + p.assign(out, err) + return p +} + +// Math functions. +func (p *Pipeline) Abs() *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.Abs(p.current), nil) + return p +} + +func (p *Pipeline) Neg() *Pipeline { + if p.shortCircuit() { + return p + } + p.assign(core.MulI(p.current, -1), nil) + return p +} + +func (p *Pipeline) Sqrt() *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Sqrt(p.current) + p.assign(out, err) + return p +} + +func (p *Pipeline) Exp() *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Exp(p.current) + p.assign(out, err) + return p +} + +func (p *Pipeline) Log() *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Log(p.current) + p.assign(out, err) + return p +} + +func (p *Pipeline) Floor() *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.Floor(p.current) + p.assign(out, err) + return p +} + +// Reductions and extrema. +func (p *Pipeline) ArgMaxAxis(dim int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.ArgMaxAxis(p.current, dim) + p.assign(out, err) + return p +} + +func (p *Pipeline) ArgMinAxis(dim int) *Pipeline { + if p.shortCircuit() { + return p + } + out, err := core.ArgMinAxis(p.current, dim) + p.assign(out, err) + return p +} + +// TopK keeps the top-k values along dim; the matching indices are +// discarded; call the package-level TopK directly when the indices +// matter. +func (p *Pipeline) TopK(k int, dim int) *Pipeline { + if p.shortCircuit() { + return p + } + values, _, err := core.TopK(p.current, k, dim) + p.assign(values, err) + return p +} diff --git a/linalg/pipeline_test.go b/linalg/pipeline_test.go new file mode 100644 index 0000000..a2f08f7 --- /dev/null +++ b/linalg/pipeline_test.go @@ -0,0 +1,112 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// mustF builds a float array or fails the test. +func mustF(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 +} + +// TestPipelineHappyChain pins the values of a chained transformation +// against the same calls made directly. +func TestPipelineHappyChain(t *testing.T) { + a := mustF(t, []float64{0.25, 4, 9, 16}, 2, 2) + b := mustF(t, []float64{1, 1, 1, 1}, 2, 2) + got, err := Pipe(a).AddF(1).Sqrt().Mul(b).Transpose().Reshape(4).Result() + if err != nil { + t.Fatalf("pipeline: %v", err) + } + want := core.Transpose(mustF(t, []float64{math.Sqrt(1.25), math.Sqrt(5), math.Sqrt(10), math.Sqrt(17)}, 2, 2)) + if got.Len() != 4 { + t.Fatalf("result length %d, want 4", got.Len()) + } + for i := range 4 { + if math.Abs(got.FloatAt(i)-want.FloatAt(i)) > 1e-12 { + t.Fatalf("[%d] = %g, want %g", i, got.FloatAt(i), want.FloatAt(i)) + } + } +} + +// TestPipelineShortCircuit pins the error contract: a failing step +// records the first error, Result returns the array that triggered it, +// and every later step is a no-op however wrong its arguments are. +func TestPipelineShortCircuit(t *testing.T) { + a := mustF(t, []float64{1, 2, 3, 4}, 2, 2) + // Reshape to an incompatible count fails. + out, err := Pipe(a).Reshape(3, 5).Sqrt().AddF(1).MatMul2D(nil).Result() + if err == nil { + t.Fatal("expected the reshape error to surface") + } + if out != a { + t.Fatal("Result must return the array that triggered the error") + } + // A singular Inv fails and freezes the chain. + singular := mustF(t, []float64{1, 2, 2, 4}, 2, 2) + out2, err2 := Pipe(singular).Inv().MulF(3).Result() + if err2 == nil { + t.Fatal("expected the singular Inv error to surface") + } + if out2 != singular { + t.Fatal("Result must return the singular matrix itself") + } + // Reducing the only dimension of a 1-D array fails and freezes the + // chain. + vec := mustF(t, []float64{1, 2}, 2) + out3, err3 := Pipe(vec).SumAxis(0).Exp().Result() + if err3 == nil { + t.Fatal("expected the SumAxis error to surface") + } + if out3 != vec { + t.Fatal("Result must return the array the SumAxis failed on") + } +} + +// TestPipelineScalarsAndExtrema pins the scalar and reduction steps +// against the same calls made directly. +func TestPipelineScalarsAndExtrema(t *testing.T) { + a := mustF(t, []float64{1, -2, 3, -4}, 2, 2) + got, err := Pipe(a).MulF(2).ClipF(-5, 5).SumAxis(0).Result() + if err != nil { + t.Fatalf("pipeline: %v", err) + } + mul := core.MulF(a, 2) + clipped, err := core.ClipF(mul, -5, 5) + if err != nil { + t.Fatalf("ClipF: %v", err) + } + want, err := core.SumAxis(clipped, 0) + if err != nil { + t.Fatalf("SumAxis: %v", err) + } + if got.Len() != want.Len() { + t.Fatalf("length %d, want %d", got.Len(), want.Len()) + } + for i := range want.Len() { + if got.FloatAt(i) != want.FloatAt(i) { + t.Fatalf("[%d] = %g, want %g", i, got.FloatAt(i), want.FloatAt(i)) + } + } + b := mustF(t, []float64{5, -6, 7, -8}, 4) + got2, err2 := Pipe(b).Abs().Neg().Result() + if err2 != nil { + t.Fatalf("pipeline: %v", err2) + } + for i := range 4 { + if got2.FloatAt(i) != float64(-i-5) { + t.Fatalf("[%d] = %g", i, got2.FloatAt(i)) + } + } +} diff --git a/linalg/polyroots.go b/linalg/polyroots.go new file mode 100644 index 0000000..2213d94 --- /dev/null +++ b/linalg/polyroots.go @@ -0,0 +1,82 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Polynomial roots through the companion matrix. The roots of a +// polynomial are exactly the eigenvalues of its companion matrix, so +// the general nonsymmetric eigensolver answers the question directly: +// no Aberth iteration, no bracketing, one direct construction and a +// free ride on EigenGeneral's shifted QR. (The companion matrix is +// not balanced; coefficients spread over many magnitudes condition +// the roots through the eigenvalue problem as it stands.) + +// PolynomialRoots returns the roots of the polynomial whose +// coefficients are given in ascending power order, lowest power first, +// the same convention EvaluatePolynomial uses. The answer is a complex +// vector sorted descending by magnitude, as EigenGeneral orders its +// values. Trailing zero coefficients raise nothing: they are stripped +// before the companion matrix is built, so the degree is the true one. +// A nonzero constant has no roots and answers an empty vector; the +// zero polynomial has every point as a root and is an error, as are +// coefficients that are not a vector and an empty coefficient list. +func PolynomialRoots(coeffs *core.Array) (*core.Array, error) { + const name = "PolynomialRoots" + if coeffs.NDim() != 1 { + return nil, base.Errf("%s: coefficients must be a vector, got shape %s", + name, base.ShapeText(coeffs.Shape())) + } + n := coeffs.Len() + if n == 0 { + return nil, base.Errf("%s: the coefficient vector must not be empty", name) + } + c := make([]complex128, n) + for i := range n { + if coeffs.Dtype() == core.Complex { + c[i] = coeffs.ComplexAt(i) + } else { + c[i] = complex(coeffs.FloatAt(i), 0) + } + } + // Strip trailing zeros to reach the true degree. + for n > 0 && c[n-1] == 0 { + n-- + } + if n == 0 { + return nil, base.Errf("%s: the zero polynomial has every point as a root", name) + } + if n == 1 { + return core.FromComplexes(nil, 0) + } + // The Frobenius companion of the monic polynomial: ones on the + // subdiagonal, the negated scaled coefficients down the last + // column. Its characteristic polynomial is p(x)/c_{n-1}, so its + // eigenvalues are the roots. + degree := n - 1 + companion := make([]complex128, degree*degree) + for row := 1; row < degree; row++ { + companion[row*degree+row-1] = 1 + } + for k := range degree { + companion[k*degree+degree-1] = -c[k] / c[degree] + } + values, _, err := EigenGeneral(fromComplexesMust(companion, degree, degree)) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + return values, nil +} + +// fromComplexesMust wraps a construction that cannot fail: the value +// count always matches the two-dimensional shape. It is unexported on +// purpose: a library that panics on a caller's input is a defect, and +// this caller cannot fail. +func fromComplexesMust(vals []complex128, rows, cols int) *core.Array { + a, _ := core.FromComplexes(vals, rows, cols) + return a +} diff --git a/linalg/polyroots_test.go b/linalg/polyroots_test.go new file mode 100644 index 0000000..0f37220 --- /dev/null +++ b/linalg/polyroots_test.go @@ -0,0 +1,135 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "testing" + +// rootSetMatch reports whether the computed roots equal the wanted set +// within tolerance: every wanted root has a computed partner and every +// computed root is claimed. Magnitude-sorted input order carries no +// meaning for a conjugate pair beyond rounding. +func rootSetMatch(got, want []complex128, tol float64) bool { + if len(got) != len(want) { + return false + } + used := make([]bool, len(got)) + for _, w := range want { + found := false + for i, g := range got { + if !used[i] && abs2c(g-w) <= tol*tol { + used[i] = true + found = true + break + } + } + if !found { + return false + } + } + return true +} + +func abs2c(z complex128) float64 { + return real(z)*real(z) + imag(z)*imag(z) +} + +// TestPolynomialRootsFactored pins roots read off factored forms: a +// real pair, a cubic with a conjugate pair, and a degree-one case. +func TestPolynomialRootsFactored(t *testing.T) { + cases := []struct { + coeffs []float64 + want []complex128 + }{ + {[]float64{6, -5, 1}, []complex128{2, 3}}, + {[]float64{-13, 17, -5, 1}, []complex128{1, 2 + 3i, 2 - 3i}}, + {[]float64{-4, 2}, []complex128{2}}, + } + for k, tc := range cases { + roots, err := PolynomialRoots(mustFloats(t, tc.coeffs, len(tc.coeffs))) + if err != nil { + t.Fatalf("case %d: PolynomialRoots: %v", k, err) + } + got := make([]complex128, roots.Len()) + for i := range got { + got[i] = roots.ComplexAt(i) + } + if !rootSetMatch(got, tc.want, 1e-8) { + t.Fatalf("case %d: roots %v, want %v", k, got, tc.want) + } + } +} + +// TestPolynomialRootsRepeated checks a double root: the companion is +// defective, so the squared-off accuracy is all a similarity-based +// solver can give and the tolerance says so. +func TestPolynomialRootsRepeated(t *testing.T) { + roots, err := PolynomialRoots(mustFloats(t, []float64{4, -4, 1}, 3)) + if err != nil { + t.Fatalf("PolynomialRoots: %v", err) + } + for i := range roots.Len() { + if d := roots.ComplexAt(i) - 2; abs2c(d) > 1e-8 { + t.Fatalf("root %d = %v, want 2 ± 1e-8", i, roots.ComplexAt(i)) + } + } +} + +// TestPolynomialRootsTrailingZeros strips trailing zero coefficients: +// the degree is the true one and the roots are unchanged. +func TestPolynomialRootsTrailingZeros(t *testing.T) { + roots, err := PolynomialRoots(mustFloats(t, []float64{6, -5, 1, 0, 0}, 5)) + if err != nil { + t.Fatalf("PolynomialRoots: %v", err) + } + if roots.Len() != 2 { + t.Fatalf("got %d roots, want 2", roots.Len()) + } + if !rootSetMatch([]complex128{roots.ComplexAt(0), roots.ComplexAt(1)}, + []complex128{2, 3}, 1e-8) { + t.Fatalf("roots %v, want {2, 3}", roots.Shape()) + } +} + +// TestPolynomialRootsComplexCoefficients exercises the complex path: +// x − (1+2i) has the obvious root. +func TestPolynomialRootsComplexCoefficients(t *testing.T) { + coeffs, err := core.FromComplexes([]complex128{-(1 + 2i), 1}, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + roots, err := PolynomialRoots(coeffs) + if err != nil { + t.Fatalf("PolynomialRoots: %v", err) + } + if roots.Len() != 1 || abs2c(roots.ComplexAt(0)-(1+2i)) > 1e-12 { + t.Fatalf("root = %v, want 1+2i", roots.ComplexAt(0)) + } +} + +// TestPolynomialRootsErrors pins the degenerate contracts: a constant +// answers an empty vector, the zero polynomial is an error, and so are +// wrong shapes. +func TestPolynomialRootsErrors(t *testing.T) { + roots, err := PolynomialRoots(mustFloats(t, []float64{3}, 1)) + if err != nil { + t.Fatalf("constant polynomial: %v", err) + } + if roots.Len() != 0 { + t.Fatalf("constant polynomial returned %d roots, want 0", roots.Len()) + } + if _, err := PolynomialRoots(mustFloats(t, []float64{0, 0}, 2)); err == nil { + t.Fatal("expected an error for the zero polynomial") + } + if _, err := PolynomialRoots(mustFloats(t, nil)); err == nil { + t.Fatal("expected an error for empty coefficients") + } + matrix, _ := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2) + if _, err := PolynomialRoots(matrix); err == nil { + t.Fatal("expected an error for a rank-2 coefficient array") + } +} diff --git a/linalg/rrqr.go b/linalg/rrqr.go new file mode 100644 index 0000000..d73d1d4 --- /dev/null +++ b/linalg/rrqr.go @@ -0,0 +1,422 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "slices" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Rank-revealing QR with column pivoting (Businger and Golub). The +// plain Householder QR of `QR` answers a factorisation but hides the +// rank: an exactly or nearly dependent column leaves a rounding-level +// R diagonal and `LeastSquares` has to refuse the system outright. +// Pivoting sweeps the column of largest remaining norm to the front at +// every step, which forces the dependence into the trailing block: the +// |R| diagonal comes out non-increasing and its decay is the rank +// decision, the same tolerance convention `LeastSquares` applies to +// its unpivoted diagonal. + +// RRQR returns the rank-revealing QR factorisation with column +// pivoting of an m×n matrix a with m ≥ n: A·P = Q·R, with Q m×m +// orthogonal, R m×n upper triangular, and perm the column permutation +// P, meaning the j-th column of A·P is column perm[j] of A. The rank +// is the count of |R[i,i]| above n·eps·max|R|, the tolerance +// `LeastSquares` applies; with column pivoting the diagonal decays +// monotonically, so the count is where the decay crosses the floor. +// Complex matrices, non-2-D inputs and underdetermined shapes are +// refused, matching the dense surface. +func RRQR(a *core.Array) (q, r *core.Array, perm []int, rank int, err error) { + const name = "RRQR" + if a.Dtype() == core.Complex { + return nil, nil, nil, 0, base.Errf("%s: complex matrices are not supported", name) + } + if a.NDim() != 2 { + return nil, nil, nil, 0, base.Errf("%s: needs a 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m == 0 || n == 0 { + return nil, nil, nil, 0, base.Errf("%s: zero-sized matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + if m < n { + return nil, nil, nil, 0, base.Errf("%s: needs m ≥ n, got shape %s", name, base.ShapeText(a.Shape())) + } + rMat, qMat, perm := rrqrFactor(denseFloats(a, m, n), m, n, true) + return floatsToArray(qMat, []int{m, m}), floatsToArray(rMat, []int{m, n}), perm, rrqrRankOf(rMat, n), nil +} + +// RRQRRank returns the numerical rank a pivoted factorisation of a +// reports under the package's tolerance convention: the count of +// |R[i,i]| above n·eps·max|R|. It is the rank `RRQR` returns, computed +// without the orthogonal factor for callers that only want the count. +func RRQRRank(a *core.Array) (int, error) { + const name = "RRQRRank" + if a.Dtype() == core.Complex { + return 0, base.Errf("%s: complex matrices are not supported", name) + } + if a.NDim() != 2 { + return 0, base.Errf("%s: needs a 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m == 0 || n == 0 { + return 0, base.Errf("%s: zero-sized matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + if m < n { + return 0, base.Errf("%s: needs m ≥ n, got shape %s", name, base.ShapeText(a.Shape())) + } + rMat, _, _ := rrqrFactor(denseFloats(a, m, n), m, n, false) + return rrqrRankOf(rMat, n), nil +} + +// rrqrRankOf counts the diagonal entries of the flat m×n R above +// n·eps·max|R|, the rank convention `LeastSquares` states for its +// unpivoted diagonal. +func rrqrRankOf(rMat []float64, n int) int { + rMax := 0.0 + for i := range n { + if v := math.Abs(rMat[i*n+i]); v > rMax { + rMax = v + } + } + rankTol := float64(n) * base.EpsF * rMax + rank := 0 + for i := range n { + if math.Abs(rMat[i*n+i]) > rankTol { + rank++ + } + } + return rank +} + +// rrqrColumnBlock is the number of R columns one blocked reflector pass +// covers, the blocking qrRaw applies to the same sweep and for the same +// reason: R is row-major, so a lone column's update walks a stride-n +// pattern that touches one cache line per element, where a block of +// rrqrColumnBlock columns spans one line of float64 and walks each row +// once. +const rrqrColumnBlock = 8 + +// rrqrMinWork is the element-touch floor below which one worker's share +// of a reflector application or of a column-norm scan falls back to the +// calling goroutine: below it the dispatch costs more than the work. +const rrqrMinWork = 2048 + +// rrqrFactor runs the pivoted Householder sweep on a flat m×n matrix. +// The column of largest remaining norm moves to the front at every +// step; norms are recomputed each step rather than downdated, so no +// cancellation drift and no squared-overflow window: the cost is one +// pass over the working block per step, the same order as the +// reflector sweep itself. The Householder reflectors and their +// application follow `qrRaw`: scale before squaring, one reflector per +// step, strictly in sweep order, so the serial result is the +// factorisation the parallel plain QR also answers. +func rrqrFactor(aMat []float64, m, n int, wantQ bool) (rMat, qMat []float64, perm []int) { + rMat = make([]float64, m*n) + copy(rMat, aMat) + perm = make([]int, n) + for i := range perm { + perm[i] = i + } + if wantQ { + qMat = eye(m) + } + col := make([]float64, m) + norms := make([]float64, n) + for k := range n { + best := rrqrPivot(rMat, m, n, k, perm, norms) + if best == 0 { + // Every remaining column is exactly zero: the trailing + // block of R stays zero and the sweep is done. + break + } + lv := col[:m-k] + for i := range lv { + lv[i] = rMat[(k+i)*n+k] + } + hh := householderVectorInto(lv, lv) + if hh.beta == 0 { + continue + } + rrqrApply(rMat, m, n, k, hh, qMat, nil) + } + return rMat, qMat, perm +} + +// rrqrPivot selects the column of largest remaining norm among columns +// k..n−1, swaps it into column k and records the swap in perm. The +// norms themselves are computed one worker per contiguous column slab, +// each column by the serial two-pass accumulation stridedNorm2 always +// ran, and the largest is then picked by a serial scan whose strict +// comparison keeps the first of equal maxima, so the tie rule is +// unchanged. A zero return means every remaining column is exactly zero. +// +// A slab is a run of consecutive columns, which is what makes the scan +// affordable at all: column j and column j+1 share seven of every eight +// cache lines their strided walks touch, so a slab re-reads them from +// the cache a lone column at a time cannot. The crew only forms when one +// worker's slab holds a full dispatch floor of entries. +func rrqrPivot(rMat []float64, m, n, k int, perm []int, norms []float64) float64 { + cols, rows := n-k, m-k + engine.ParallelMin(cols, max((rrqrMinWork+rows-1)/rows, 1), func(start, end int) { + for j := start; j < end; j++ { + norms[j] = stridedNorm2(rMat, k*n+k+j, n, rows) + } + }) + best, bestJ := 0.0, k + for j := range cols { + if nrm := norms[j]; nrm > best { + best, bestJ = nrm, k+j + } + } + if best == 0 { + return 0 + } + if bestJ != k { + rrqrSwapCols(rMat, m, n, bestJ, k) + perm[bestJ], perm[k] = perm[k], perm[bestJ] + } + return best +} + +// rrqrApply applies the reflector to the R columns k..n−1 in column +// blocks and, when qMat is non-nil, to the m rows of the orthonormal +// accumulator, or when bv is non-nil to the tail of the right-hand side. +// The targets are disjoint buffers, so they share one crew; the R blocks +// and the accumulator rows take contiguous slices, and every dot keeps +// its own i-ascending chain with the subtraction repeating the same +// operand order, so no crew size and no block boundary moves a bit. +func rrqrApply(rMat []float64, m, n, k int, hh householder, qMat, bv []float64) { + ln := m - k + blocks := (n - k + rrqrColumnBlock - 1) / rrqrColumnBlock + items := blocks + if qMat != nil { + items += m + } + if bv != nil { + items++ + } + weight := (n - k) + m + w := min(weight*ln/rrqrMinWork+1, engine.WorkersFor(items)) + if w < 2 { + rrqrApplyWorker(rMat, m, n, k, hh, qMat, bv, blocks, items, 0, 1) + return + } + var wg sync.WaitGroup + for gi := range w { + wg.Go(func() { + rrqrApplyWorker(rMat, m, n, k, hh, qMat, bv, blocks, items, gi, w) + }) + } + wg.Wait() +} + +// rrqrApplyWorker runs crew member gi of nw over its share of the +// reflector's targets: a contiguous slab of the R column blocks, a +// contiguous slab of the accumulator rows, and the right-hand side. +func rrqrApplyWorker(rMat []float64, m, n, k int, hh householder, qMat, bv []float64, blocks, items, gi, nw int) { + ln := m - k + // Target layout in the item index space: the R blocks first, then the + // accumulator rows when there is an accumulator, then the right-hand + // side when there is one. + qEnd := blocks + if qMat != nil { + qEnd += m + } + var d [rrqrColumnBlock]float64 + for it := gi * items / nw; it < (gi+1)*items/nw; it++ { + switch { + case it < blocks: + j0 := k + it*rrqrColumnBlock + j1 := min(j0+rrqrColumnBlock, n) + wd := j1 - j0 + _ = d[wd-1] // bounds-check proof: the block never exceeds rrqrColumnBlock + clear(d[:wd]) + for jj := range wd { + jcol := j0 + jj + for i := range ln { + d[jj] += hh.v[i] * rMat[(k+i)*n+jcol] + } + } + for j := range wd { + d[j] = hh.beta * d[j] + } + for i := range ln { + vi := hh.v[i] + base := (k+i)*n + j0 + for j := range wd { + rMat[base+j] -= vi * d[j] + } + } + case it < qEnd: + row := it - blocks + base := row*m + k + dot := 0.0 + for i := range ln { + dot += hh.v[i] * qMat[base+i] + } + wq := hh.beta * dot + for i := range ln { + qMat[base+i] -= hh.v[i] * wq + } + default: + dot := 0.0 + for i := range ln { + dot += hh.v[i] * bv[k+i] + } + wb := hh.beta * dot + for i := range ln { + bv[k+i] -= hh.v[i] * wb + } + } + } +} + +// rrqrSwapCols exchanges columns j1 and j2 of a flat m×n matrix. +func rrqrSwapCols(rMat []float64, m, n, j1, j2 int) { + for i := range m { + rMat[i*n+j1], rMat[i*n+j2] = rMat[i*n+j2], rMat[i*n+j1] + } +} + +// stridedNorm2 returns the L2 norm of count entries walking stride +// step from start, summed relative to the largest magnitude so a +// column near 1e154 does not overflow on the way to its norm. +func stridedNorm2(vals []float64, start, step, count int) float64 { + maxAbs := 0.0 + for i := range count { + if a := math.Abs(vals[start+i*step]); a > maxAbs { + maxAbs = a + } + } + if maxAbs == 0 { + return 0 + } + s := 0.0 + for i := range count { + t := vals[start+i*step] / maxAbs + s += t * t + } + return maxAbs * math.Sqrt(s) +} + +// rrqrValidPerm reports whether p is a permutation of 0..n−1, the +// property the returned permutation's contract rests on. +func rrqrValidPerm(p []int) bool { + seen := make([]bool, len(p)) + for _, v := range p { + if v < 0 || v >= len(p) || seen[v] { + return false + } + seen[v] = true + } + return !slices.Contains(seen, false) +} + +// SolveRRQR solves min ‖A·x − b‖₂ through the pivoted factorisation. +// On full rank the answer is the ordinary back-substitution through R +// against c = Qᵀb. On rank-deficient input it is the minimum-norm +// least-squares solution: the rank r is read off the R diagonal decay, +// and the leading r×n trapezoidal block R₁ answers its underdetermined +// system in minimum norm through x = R₁ᵀ(R₁R₁ᵀ)⁻¹c₁, which squares the +// condition of the leading block exactly the way the package's +// BᵀB-based SVD squares the spectrum: honest for a rank decision taken +// from a decay that has already been observed. +func SolveRRQR(a, b *core.Array) (*core.Array, error) { + const name = "SolveRRQR" + if a.Dtype() == core.Complex { + return nil, base.Errf("%s: complex matrices are not supported", name) + } + if b.Dtype() == core.Complex { + return nil, base.Errf("%s: complex right-hand sides are not supported", name) + } + if a.NDim() != 2 { + return nil, base.Errf("%s: needs a 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if m == 0 || n == 0 { + return nil, base.Errf("%s: zero-sized matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + if m < n { + return nil, base.Errf("%s: needs m ≥ n, got shape %s", name, base.ShapeText(a.Shape())) + } + if b.NDim() != 1 || b.Len() != m { + return nil, base.Errf("%s: right-hand side must be a vector of length %d, got shape %s", + name, m, base.ShapeText(b.Shape())) + } + // The sweep below carries b through the reflectors as they are + // built, so c = Qᵀb falls out of the factorisation without the + // orthogonal factor ever being materialised. + rMat := make([]float64, m*n) + copy(rMat, denseFloats(a, m, n)) + bv := vectorF64(b, m) + perm := make([]int, n) + for i := range perm { + perm[i] = i + } + col := make([]float64, m) + norms := make([]float64, n) + for k := range n { + if rrqrPivot(rMat, m, n, k, perm, norms) == 0 { + break + } + lv := col[:m-k] + for i := range lv { + lv[i] = rMat[(k+i)*n+k] + } + hh := householderVectorInto(lv, lv) + if hh.beta == 0 { + continue + } + rrqrApply(rMat, m, n, k, hh, nil, bv) + } + x := make([]float64, n) + switch rank := rrqrRankOf(rMat, n); { + case rank == n: + // Full rank: back-substitution through R. + for i := n - 1; i >= 0; i-- { + s := bv[i] + for j := i + 1; j < n; j++ { + s -= rMat[i*n+j] * x[j] + } + x[i] = s / rMat[i*n+i] + } + default: + // Minimum norm through the leading trapezoidal block R₁: + // R₁R₁ᵀw = c₁, x = R₁ᵀw. R₁R₁ᵀ is symmetric positive definite + // at the rank the decay showed, and the dense `Solve` carries + // the small r×r system. + if rank > 0 { + mm := make([]float64, rank*rank) + for i := range rank { + for j := range rank { + s := 0.0 + for l := range n { + s += rMat[i*n+l] * rMat[j*n+l] + } + mm[i*rank+j] = s + } + } + w, err := Solve(floatsToArray(mm, []int{rank, rank}), floatsToArray(append([]float64(nil), bv[:rank]...), []int{rank})) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + for i := range rank { + for j := range n { + x[j] += rMat[i*n+j] * w.FloatAt(i) + } + } + } + } + // x lives in permuted coordinates; undo the permutation. + out := make([]float64, n) + for j := range n { + out[perm[j]] = x[j] + } + return floatsToArray(out, []int{n}), nil +} diff --git a/linalg/rrqr_test.go b/linalg/rrqr_test.go new file mode 100644 index 0000000..fd1cc5b --- /dev/null +++ b/linalg/rrqr_test.go @@ -0,0 +1,304 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// randomDense builds a deterministic m×n matrix with entries in [−1, 1) +// from the seeded generator, so no test randomness leaks. +func randomDense(t *testing.T, m, n int, seed int64) *core.Array { + t.Helper() + g := core.NewGenerator(seed) + vals := make([]float64, m*n) + for i := range vals { + vals[i] = float64(g.Next()%2000)/1000 - 1 + } + return floatsToArray(vals, []int{m, n}) +} + +func TestRRQRFactorisation(t *testing.T) { + const m, n = 14, 8 + a := randomDense(t, m, n, 11) + q, r, perm, rank, err := RRQR(a) + if err != nil { + t.Fatalf("RRQR: %v", err) + } + if rank != n { + t.Fatalf("rank %d for a full-rank %d×%d matrix", rank, m, n) + } + if !rrqrValidPerm(perm) { + t.Fatalf("the permutation %v is not a permutation", perm) + } + // A·P = Q·R to machine precision. + worst, aMax := 0.0, 0.0 + for i := range m { + for j := range n { + ap := a.FloatAt(i*n + perm[j]) + qr := 0.0 + for l := range m { + qr += q.FloatAt(i*m+l) * r.FloatAt(l*n+j) + } + if d := math.Abs(ap - qr); d > worst { + worst = d + } + if v := math.Abs(a.FloatAt(i*n + j)); v > aMax { + aMax = v + } + } + } + if worst > 1e-12*math.Max(1, aMax)*float64(m) { + t.Fatalf("A·P = Q·R off by %.3g", worst) + } + // Q orthogonal. + for i := range m { + for j := range m { + s := 0.0 + for l := range m { + s += q.FloatAt(l*m+i) * q.FloatAt(l*m+j) + } + want := 0.0 + if i == j { + want = 1 + } + if math.Abs(s-want) > 1e-13 { + t.Fatalf("QᵀQ[%d,%d] = %.16g, want %.16g", i, j, s, want) + } + } + } + // R upper triangular. + for i := range m { + for j := range n { + if j < i && math.Abs(r.FloatAt(i*n+j)) > 1e-12 { + t.Fatalf("R[%d,%d] = %.3g below the diagonal", i, j, r.FloatAt(i*n+j)) + } + } + } + // The |R| diagonal decays monotonically, the pivoting guarantee. + for i := 1; i < n; i++ { + hi := math.Abs(r.FloatAt(i*n + i)) + lo := math.Abs(r.FloatAt((i-1)*n + i - 1)) + if hi > lo*(1+1e-12) { + t.Fatalf("|R[%d,%d]| = %.6g exceeds |R[%d,%d]| = %.6g", i, i, hi, i-1, i-1, lo) + } + } + // RRQRRank is the same count without the orthogonal factor. + rank2, err := RRQRRank(a) + if err != nil { + t.Fatalf("RRQRRank: %v", err) + } + if rank2 != n { + t.Fatalf("RRQRRank = %d, want %d", rank2, n) + } +} + +func TestRRQRRankDetection(t *testing.T) { + // A hand-built 6×4: three independent columns and a fourth that is + // exactly 2·c0 + c1 in float arithmetic. + const m, n = 6, 4 + g := core.NewGenerator(5) + raw := make([]float64, m*n) + for i := range m { + raw[i*n+0] = float64(g.Next()%100)/50 - 1 + raw[i*n+1] = float64(g.Next()%100)/50 - 1 + raw[i*n+2] = float64(g.Next()%100)/50 - 1 + raw[i*n+3] = 2*raw[i*n+0] + raw[i*n+1] + } + a := floatsToArray(raw, []int{m, n}) + _, _, perm, rank, err := RRQR(a) + if err != nil { + t.Fatalf("RRQR: %v", err) + } + if rank != 3 { + t.Fatalf("rank %d for a matrix whose fourth column is exactly dependent", rank) + } + if !rrqrValidPerm(perm) { + t.Fatalf("the permutation %v is not a permutation", perm) + } + // The SVD count must agree, independently. + svdRank, err := MatrixRank(a, 0) + if err != nil { + t.Fatalf("MatrixRank: %v", err) + } + if svdRank != 3 { + t.Fatalf("the SVD counts rank %d where the pivoted diagonal shows 3", svdRank) + } + // An exact zero column is the cleanest possible direction. + zero := floatsToArray([]float64{ + 1, 0, 2, 0, + 2, 1, 1, 0, + 1, 1, 3, 0, + 2, 0, 1, 0, + 1, 2, 2, 0, + 3, 1, 4, 0, + }, []int{m, n}) + rankZero, err := RRQRRank(zero) + if err != nil { + t.Fatalf("RRQRRank: %v", err) + } + if rankZero != 3 { + t.Fatalf("rank %d for a matrix with an exact zero column", rankZero) + } + // The zero matrix has no rank at all. + empty := floatsToArray(make([]float64, 3*2), []int{3, 2}) + rankEmpty, err := RRQRRank(empty) + if err != nil { + t.Fatalf("RRQRRank: %v", err) + } + if rankEmpty != 0 { + t.Fatalf("rank %d for the zero matrix", rankEmpty) + } +} + +func TestSolveRRQRMinimumNorm(t *testing.T) { + // Rank-deficient and consistent: the minimum-norm answer, checked + // against the SVD's Pinverse and, on a hand-checkable 3×3, against + // the exact (0, 1, 1). + const m, n = 6, 4 + g := core.NewGenerator(9) + raw := make([]float64, m*n) + for i := range m { + raw[i*n+0] = float64(g.Next()%100)/50 - 1 + raw[i*n+1] = float64(g.Next()%100)/50 - 1 + raw[i*n+2] = float64(g.Next()%100)/50 - 1 + raw[i*n+3] = 2*raw[i*n+0] + raw[i*n+1] + } + a := floatsToArray(raw, []int{m, n}) + xTrue := make([]float64, n) + for i := range n { + xTrue[i] = math.Cos(0.3*float64(i)) + float64(i%3) + } + b := core.New(core.Float, m) + for i := range m { + s := 0.0 + for j := range n { + s += raw[i*n+j] * xTrue[j] + } + b.RawFloats()[i] = s + } + x, err := SolveRRQR(a, b) + if err != nil { + t.Fatalf("SolveRRQR: %v", err) + } + pinv, err := Pinverse(a, 0) + if err != nil { + t.Fatalf("Pinverse: %v", err) + } + worst := 0.0 + for i := range n { + s := 0.0 + for j := range m { + s += pinv.FloatAt(i*m+j) * b.FloatAt(j) + } + if d := math.Abs(x.FloatAt(i) - s); d > worst { + worst = d + } + } + if worst > 1e-8 { + t.Fatalf("the pivoted answer misses the SVD minimum norm by %.3g", worst) + } + // The hand-checkable 3×3 system from the sparse solver's own pin. + tiny := floatsToArray([]float64{1, 0, 1, 0, 1, 1, 1, 1, 2}, []int{3, 3}) + tb := mustFloats(t, []float64{1, 2, 3}, 3) + xt, err := SolveRRQR(tiny, tb) + if err != nil { + t.Fatalf("SolveRRQR: %v", err) + } + want := []float64{0, 1, 1} + for i := range want { + if math.Abs(xt.FloatAt(i)-want[i]) > 1e-10 { + t.Fatalf("x[%d] = %.12g, want the minimum-norm %.12g", i, xt.FloatAt(i), want[i]) + } + } + // Full rank and inconsistent: the ordinary least-squares answer. + full := randomDense(t, 5, 3, 13) + fb := mustFloats(t, []float64{1, -1, 2, 0.5, 3}, 5) + xf, err := SolveRRQR(full, fb) + if err != nil { + t.Fatalf("SolveRRQR: %v", err) + } + ref, err := LeastSquares(full, fb) + if err != nil { + t.Fatalf("LeastSquares: %v", err) + } + for i := range 3 { + if math.Abs(xf.FloatAt(i)-ref.FloatAt(i)) > 1e-9 { + t.Fatalf("full-rank solve x[%d] = %.12g, want %.12g", i, xf.FloatAt(i), ref.FloatAt(i)) + } + } +} + +func TestSolveRRQRErrors(t *testing.T) { + wide := floatsToArray([]float64{1, 2, 3, 4, 5, 6}, []int{2, 3}) + b2 := mustFloats(t, []float64{1, 2}, 2) + if _, err := SolveRRQR(wide, b2); err == nil { + t.Fatal("an underdetermined system was accepted") + } + complexA, err := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + if _, err := SolveRRQR(complexA, b2); err == nil { + t.Fatal("a complex matrix was accepted") + } + good := randomDense(t, 4, 3, 17) + complexB, err := core.FromComplexes([]complex128{1, 2, 3}, 3) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + if _, err := SolveRRQR(good, complexB); err == nil { + t.Fatal("a complex right-hand side was accepted") + } + if _, err := SolveRRQR(good, mustFloats(t, []float64{1, 2}, 2)); err == nil { + t.Fatal("a short right-hand side was accepted") + } + if _, err := SolveRRQR(good, core.New(core.Float, 2, 2)); err == nil { + t.Fatal("a rank-2 right-hand side was accepted") + } +} + +func TestRRQRErrors(t *testing.T) { + oneD := mustFloats(t, []float64{1, 2, 3}, 3) + if _, _, _, _, err := RRQR(oneD); err == nil { + t.Fatal("a rank-1 input was accepted") + } + complexA, err := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + if _, _, _, _, err := RRQR(complexA); err == nil { + t.Fatal("a complex matrix was accepted") + } + wide := floatsToArray([]float64{1, 2, 3, 4, 5, 6}, []int{2, 3}) + if _, _, _, _, err := RRQR(wide); err == nil { + t.Fatal("an underdetermined matrix was accepted") + } + if _, err := RRQRRank(oneD); err == nil { + t.Fatal("RRQRRank accepted a rank-1 input") + } + if _, err := RRQRRank(complexA); err == nil { + t.Fatal("RRQRRank accepted a complex matrix") + } + if _, err := RRQRRank(wide); err == nil { + t.Fatal("RRQRRank accepted an underdetermined matrix") + } + empty := floatsToArray([]float64{1}, []int{1, 0}) + if _, _, _, _, err := RRQR(empty); err == nil { + t.Fatal("an empty matrix was accepted") + } + // The permutation helper's contract. + if !rrqrValidPerm([]int{2, 0, 1}) { + t.Fatal("a rotation was rejected as a permutation") + } + if rrqrValidPerm([]int{0, 0, 1}) { + t.Fatal("a repeated index passed as a permutation") + } + if rrqrValidPerm([]int{0, 2}) { + t.Fatal("an out-of-range index passed as a permutation") + } +} diff --git a/linalg/schur.go b/linalg/schur.go new file mode 100644 index 0000000..76209ea --- /dev/null +++ b/linalg/schur.go @@ -0,0 +1,263 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "math/cmplx" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The complex Schur decomposition and the principal matrix functions +// built on it. A = Q·T·Qᴴ with Q unitary and T upper triangular turns +// f(A) into Q·f(T)·Qᴴ, and f of a triangular matrix is a recurrence: +// the diagonal maps entrywise and the off-diagonals follow from the +// function's own defining equation. That is the Schur-Parlett route to +// the principal square root and logarithm of nonsymmetric matrices, +// where the symmetric eigendecomposition the easy route needs does not +// exist. + +// SchurComplex returns the complex Schur decomposition of the square +// matrix a: t upper triangular and q unitary with a = q·t·qᴴ. For a +// real matrix whose eigenvalues come in conjugate pairs t is complex; +// the real Schur form is a different construction. The diagonal of t +// carries the eigenvalues, matching EigenGeneral's set. An input +// other than a square 2-D matrix or a non-converging iteration is an +// error. +func SchurComplex(a *core.Array) (t, q *core.Array, err error) { + const name = "SchurComplex" + if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { + return nil, nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + n := a.Shape()[0] + if n == 0 { + return nil, nil, base.Errf("%s: zero-sized matrix", name) + } + h := make([]complex128, n*n) + if a.Dtype() == core.Complex && !a.Strided() && len(a.RawComplexes()) == n*n { + copy(h, a.RawComplexes()) + } else { + for i := range n * n { + h[i] = a.ComplexAt(i) + } + } + // Outside the safe window the Hessenberg reduction's squared + // magnitudes and the Givens denominators leave the normal range: a + // tiny matrix has norm == 0 for every reflector, so the reduction is + // skipped and the iteration then runs on a matrix that is not + // Hessenberg, and a huge one has den == 0 destroy the subdiagonals. + // The matrix is moved into the window for the reduction and the + // iteration, and the triangular factor, which carries the scale, is + // moved back before it is returned; the unitary factor is + // scale-free. + ws := windowScale(maxMagComplex(h)) + if ws != 1 { + scaleComplexes(h, ws) + } + vecs := schurHessenberg(h, n) + if err := schurQR(h, vecs, n); err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + if ws != 1 { + unscaleComplexes(h, ws) + } + tArr := core.New(core.Complex, []int{n, n}...) + copy(tArr.RawComplexes(), h) + qArr := core.New(core.Complex, []int{n, n}...) + copy(qArr.RawComplexes(), vecs) + return tArr, qArr, nil +} + +// schurHessenberg reduces a in place to upper Hessenberg form by +// complex Householder similarity transforms, accumulating the unitary +// into q (seeded with the identity). +func schurHessenberg(a []complex128, n int) []complex128 { + q := make([]complex128, n*n) + for i := range n { + q[i*n+i] = 1 + } + // Reflector scratch reused across the sweep; each step uses the + // first n−k−1 entries. + vecBuf := make([]complex128, n) + for k := 0; k < n-2; k++ { + norm := 0.0 + for i := k + 1; i < n; i++ { + norm += base.Real2(a[i*n+k]) + } + norm = math.Sqrt(norm) + if norm == 0 { + continue + } + head := a[(k+1)*n+k] + beta := -norm + if real(head) < 0 { + beta = norm + } + vec := vecBuf[:n-k-1] + for i := k + 1; i < n; i++ { + vec[i-k-1] = a[i*n+k] + if i == k+1 { + vec[0] -= complex(beta, 0) + } + } + if base.Real2(vec[0]) == 0 { + continue + } + vhx := complex(0, 0) + for i := k + 1; i < n; i++ { + vhx += cmplx.Conj(vec[i-k-1]) * a[i*n+k] + } + tau := complex(1, 0) / vhx + // Left: Ḣ·x = βe₁ with the plain tau, so the similarity is + // H <- Ḣ·H·Ḣᴴ and the accumulated unitary picks up Ḣᴴ. + for j := k; j < n; j++ { + s := complex(0, 0) + for i := k + 1; i < n; i++ { + s += cmplx.Conj(vec[i-k-1]) * a[i*n+j] + } + s *= tau + for i := k + 1; i < n; i++ { + a[i*n+j] -= s * vec[i-k-1] + } + } + // Right: all rows over columns k+1..n-1, by Ḣᴴ. + for i := range n { + s := complex(0, 0) + for j := k + 1; j < n; j++ { + s += a[i*n+j] * vec[j-k-1] + } + s *= cmplx.Conj(tau) + for j := k + 1; j < n; j++ { + a[i*n+j] -= s * cmplx.Conj(vec[j-k-1]) + } + } + for r := range n { + s := complex(0, 0) + for j := k + 1; j < n; j++ { + s += q[r*n+j] * vec[j-k-1] + } + s *= cmplx.Conj(tau) + for j := k + 1; j < n; j++ { + q[r*n+j] -= s * cmplx.Conj(vec[j-k-1]) + } + } + } + return q +} + +// schurQR drives the shifted QR iteration on the Hessenberg a to upper +// triangular form, folding each rotation into the unitary q. Explicit +// Givens QR of the active block with the Wilkinson shift keeps the +// bookkeeping simple; convergence deflates subdiagonals to zero. +func schurQR(a []complex128, q []complex128, n int) error { + // Rotation scratch reused across sweeps; each sweep uses its first + // 2·(r−l+1) entries. + rots := make([]complex128, 2*n) + for iter := 0; iter < 120*n+200; iter++ { + // Deflate negligible subdiagonals. + for i := 1; i < n; i++ { + if cmplx.Abs(a[i*n+i-1]) <= base.EpsF*(cmplx.Abs(a[(i-1)*n+i-1])+cmplx.Abs(a[i*n+i])) { + a[i*n+i-1] = 0 + } + } + // Active trailing block [l..r]. + r := n - 1 + for r > 0 && a[r*n+r-1] == 0 { + r-- + } + if r == 0 { + return nil + } + l := r + for l > 0 && a[l*n+l-1] != 0 { + l-- + } + // Wilkinson shift from the trailing two-by-two. + t11, t12 := a[(r-1)*n+r-1], a[(r-1)*n+r] + t21, t22 := a[r*n+r-1], a[r*n+r] + disc := cmplx.Sqrt((t11-t22)*(t11-t22) + 4*t21*t12) + root1 := (t11 + t22 + disc) / 2 + root2 := (t11 + t22 - disc) / 2 + mu := root1 + if cmplx.Abs(root2-t22) < cmplx.Abs(root1-t22) { + mu = root2 + } + // Explicit shifted QR of the block by Givens rotations: the + // shift leaves the diagonal first and returns at the end, so + // the factorisation runs on H − μI exactly. + for i := l; i <= r; i++ { + a[i*n+i] -= mu + } + rr := rots[:2*(r-l+1)] + for i := l; i < r; i++ { + aa, ab := a[i*n+i], a[(i+1)*n+i] + // The Givens denominator is the real hypot of the two + // entries; a complex square root here would leave the + // rotation non-unitary and the iteration divergent. + den := math.Hypot(cmplx.Abs(aa), cmplx.Abs(ab)) + if den == 0 { + rr[2*(i-l)] = 1 + rr[2*(i-l)+1] = 0 + continue + } + c := aa / complex(den, 0) + s := ab / complex(den, 0) + rr[2*(i-l)] = c + rr[2*(i-l)+1] = s + // Rows i, i+1 of the block columns l..n-1. + for j := i; j < n; j++ { + h1, h2 := a[i*n+j], a[(i+1)*n+j] + a[i*n+j] = cmplx.Conj(c)*h1 + cmplx.Conj(s)*h2 + a[(i+1)*n+j] = -s*h1 + c*h2 + } + a[(i+1)*n+i] = 0 + } + // Multiply back by the rotations from the right: R·Gᴴ. + for i := l; i < r; i++ { + c, s := rr[2*(i-l)], rr[2*(i-l)+1] + for j := 0; j <= min(i+1, n-1); j++ { + h1, h2 := a[j*n+i], a[j*n+i+1] + a[j*n+i] = h1*c + h2*s + a[j*n+i+1] = -h1*cmplx.Conj(s) + h2*cmplx.Conj(c) + } + // Accumulate q <- q·Gᴴ. + for row := range n { + q1, q2 := q[row*n+i], q[row*n+i+1] + q[row*n+i] = q1*c + q2*s + q[row*n+i+1] = -q1*cmplx.Conj(s) + q2*cmplx.Conj(c) + } + } + for i := l; i <= r; i++ { + a[i*n+i] += mu + } + } + return base.Errf("the Schur iteration did not converge") +} + +// triSqrt overwrites the upper triangular t with its principal square +// root by the recurrence s·s = t, diagonal first. A vanishing +// s_ii + s_jj denominator means the principal root does not exist for +// this spectrum and the sweep reports it. +func triSqrt(t []complex128, n int) error { + s := make([]complex128, n*n) + for i := range n { + s[i*n+i] = cmplx.Sqrt(t[i*n+i]) + } + for j := 1; j < n; j++ { + for i := j - 1; i >= 0; i-- { + sum := t[i*n+j] + for k := i + 1; k < j; k++ { + sum -= s[i*n+k] * s[k*n+j] + } + den := s[i*n+i] + s[j*n+j] + if den == 0 { + return base.Errf("the principal square root does not exist: repeated zero eigenvalues on the diagonal") + } + s[i*n+j] = sum / den + } + } + copy(t, s) + return nil +} diff --git a/linalg/schur_test.go b/linalg/schur_test.go new file mode 100644 index 0000000..c87ad57 --- /dev/null +++ b/linalg/schur_test.go @@ -0,0 +1,197 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// nonsymSample builds a deterministic nonsymmetric square matrix. +func nonsymSample(n int) *core.Array { + vals := make([]float64, n*n) + for i := range n * n { + vals[i] = math.Sin(float64(3*i + 2)) + } + a, _ := core.FromFloats(vals, n, n) + return a +} + +// TestSchurComplexContracts pins the decomposition: a = q·t·qᴴ to +// rounding, q unitary, t upper triangular, and the diagonal of t the +// same set of eigenvalues EigenGeneral reports. +func TestSchurComplexContracts(t *testing.T) { + n := 5 + a := nonsymSample(n) + tm, q, err := SchurComplex(a) + if err != nil { + t.Fatalf("SchurComplex: %v", err) + } + scale := 0.0 + for i := range n * n { + scale = math.Max(scale, absComplex(tm.ComplexAt(i))) + } + // Reconstruction a ≈ q·t·qᴴ. + w := make([]complex128, n*n) + for i := range n { + for j := range n { + s := complex(0, 0) + for k := range n { + s += q.ComplexAt(i*n+k) * tm.ComplexAt(k*n+j) + } + w[i*n+j] = s + } + } + recon := 0.0 + for i := range n { + for j := range n { + s := complex(0, 0) + for k := range n { + s += w[i*n+k] * cmplxConj(q.ComplexAt(j*n+k)) + } + recon = math.Max(recon, absComplex(s-a.ComplexAt(i*n+j))) + } + } + if recon > 1e-10*math.Max(1, scale) { + t.Fatalf("reconstruction error %.3g", recon) + } + // t strictly upper triangular below the diagonal. + for i := 1; i < n; i++ { + for j := 0; j < i; j++ { + if absComplex(tm.ComplexAt(i*n+j)) > 1e-10*math.Max(1, scale) { + t.Fatalf("t[%d][%d] = %v not deflated", i, j, tm.ComplexAt(i*n+j)) + } + } + } + // Eigenvalue set match against EigenGeneral. + want, _, err := EigenGeneral(a) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + got := make([]complex128, n) + for i := range n { + got[i] = tm.ComplexAt(i*n + i) + } + wants := make([]complex128, n) + for i := range n { + wants[i] = want.ComplexAt(i) + } + if !rootSetMatch(got, wants, 1e-6) { + t.Fatalf("Schur eigenvalues %v do not match EigenGeneral's %v", got, wants) + } +} + +func cmplxConj(z complex128) complex128 { return complex(real(z), -imag(z)) } + +// TestMatrixSqrtNonsymmetric pins exact principal roots: the square of +// an upper triangular matrix recovers it, and the rotation-like pair +// whose square has purely imaginary eigenvalues comes back too. +func TestMatrixSqrtNonsymmetric(t *testing.T) { + // [[2,1],[0,3]]² = [[4,5],[0,9]] and the principal root is the + // original: eigenvalues 2, 3 sit on the positive real axis. + a := mustFloats(t, []float64{4, 5, 0, 9}, 2, 2) + s, err := MatrixSqrt(a) + if err != nil { + t.Fatalf("MatrixSqrt: %v", err) + } + if math.Abs(s.FloatAt(0)-2) > 1e-8 || math.Abs(s.FloatAt(1)-1) > 1e-8 || + math.Abs(s.FloatAt(2)) > 1e-8 || math.Abs(s.FloatAt(3)-3) > 1e-8 { + t.Fatalf("sqrt = [[%.10g, %.10g], [%.10g, %.10g]], want [[2, 1], [0, 3]]", + s.FloatAt(0), s.FloatAt(1), s.FloatAt(2), s.FloatAt(3)) + } + // [[1,1],[−1,1]]² = [[0,2],[−2,0]]: the principal root of the + // rotation-doubler is the original, whose eigenvalues 1±i are the + // principal square roots of ±2i. + b := mustFloats(t, []float64{0, 2, -2, 0}, 2, 2) + s2, err := MatrixSqrt(b) + if err != nil { + t.Fatalf("MatrixSqrt rotation pair: %v", err) + } + if math.Abs(s2.FloatAt(0)-1) > 1e-8 || math.Abs(s2.FloatAt(1)-1) > 1e-8 || + math.Abs(s2.FloatAt(2)+1) > 1e-8 || math.Abs(s2.FloatAt(3)-1) > 1e-8 { + t.Fatalf("sqrt = [[%.10g, %.10g], [%.10g, %.10g]], want [[1, 1], [−1, 1]]", + s2.FloatAt(0), s2.FloatAt(1), s2.FloatAt(2), s2.FloatAt(3)) + } +} + +// TestMatrixLogNonsymmetric walks the round trip through the matrix +// exponential: log(exp(B)) must return B for a nonsymmetric B, and the +// logarithm of a triangular matrix must match the analytic entries. +func TestMatrixLogNonsymmetric(t *testing.T) { + b := mustFloats(t, []float64{1, 0.5, 0, 2}, 2, 2) + a, err := MatrixExp(b) + if err != nil { + t.Fatalf("MatrixExp: %v", err) + } + lg, err := MatrixLog(a) + if err != nil { + t.Fatalf("MatrixLog: %v", err) + } + for i := range 4 { + if math.Abs(lg.FloatAt(i)-b.FloatAt(i)) > 1e-6 { + t.Fatalf("log(exp(B))[%d] = %.12g, want %.12g", i, lg.FloatAt(i), b.FloatAt(i)) + } + } + // Rotation generator: exp(B) is a plane rotation with eigenvalues + // e^{±iθ}, squarely in the nonsymmetric complex-eigenvalue case. + const theta = 0.7 + rot := mustFloats(t, []float64{ + math.Cos(theta), -math.Sin(theta), + math.Sin(theta), math.Cos(theta), + }, 2, 2) + lg2, err := MatrixLog(rot) + if err != nil { + t.Fatalf("MatrixLog rotation: %v", err) + } + back, err := MatrixExp(lg2) + if err != nil { + t.Fatalf("MatrixExp round trip: %v", err) + } + for i := range 4 { + if math.Abs(back.FloatAt(i)-rot.FloatAt(i)) > 1e-6 { + t.Fatalf("exp(log(R))[%d] = %.12g, want %.12g", i, back.FloatAt(i), rot.FloatAt(i)) + } + } +} + +// TestMatrixFunctionRefusals pins the honest refusals: a negative real +// eigenvalue blocks the principal square root and anything on the +// non-positive real axis blocks the principal logarithm. +func TestMatrixFunctionRefusals(t *testing.T) { + negSqrt := mustFloats(t, []float64{-4, 1, 0, 1}, 2, 2) + if _, err := MatrixSqrt(negSqrt); err == nil { + t.Fatal("expected an error for a square root with a negative real eigenvalue") + } + negLog := mustFloats(t, []float64{-2, 1, 0, 4}, 2, 2) + if _, err := MatrixLog(negLog); err == nil { + t.Fatal("expected an error for a logarithm with a negative real eigenvalue") + } + singLog := mustFloats(t, []float64{0, 1, 0, 4}, 2, 2) + if _, err := MatrixLog(singLog); err == nil { + t.Fatal("expected an error for a logarithm at zero") + } +} + +// TestMatrixSqrtComplexInput checks the complex path: the principal +// root of diag(1+i, 4) squares back to the input. +func TestMatrixSqrtComplexInput(t *testing.T) { + a, _ := core.FromComplexes([]complex128{1 + 1i, 0, 0, 4}, 2, 2) + s, err := MatrixSqrt(a) + if err != nil { + t.Fatalf("MatrixSqrt complex: %v", err) + } + if s.Dtype() != core.Complex { + t.Fatal("the complex answer must stay complex") + } + sq, err := core.MatMul2D(s, s) + if err != nil { + t.Fatalf("MatMul2D: %v", err) + } + for i := range 4 { + if absComplex(sq.ComplexAt(i)-a.ComplexAt(i)) > 1e-10 { + t.Fatalf("S·S[%d] = %v, want %v", i, sq.ComplexAt(i), a.ComplexAt(i)) + } + } +} diff --git a/linalg/solver_dtype_pins_test.go b/linalg/solver_dtype_pins_test.go new file mode 100644 index 0000000..6025555 --- /dev/null +++ b/linalg/solver_dtype_pins_test.go @@ -0,0 +1,169 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// Regression pins for the solver dtype contracts: GMRES refuses a +// complex right-hand side by name and solves in float32, the sparse +// eigenvalue and exponential paths survive tiny scales, a NaN operator +// is a loud failure, and a negative polynomial degree is refused. + +// TestGMRESRejectsComplexRHS pins the dtype contract: a complex b must +// be a loud error, never the silent zero solution a nil RawFloats +// payload used to produce. +func TestGMRESRejectsComplexRHS(t *testing.T) { + b, err := core.FromComplexes([]complex128{1 + 1i, 2}, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + op := func(v *core.Array) (*core.Array, error) { return v, nil } + if _, err := GMRES(op, b, 0, 0, 0); err == nil { + t.Fatal("GMRES accepted a complex right-hand side") + } +} + +// TestGMRESFloat32RHS verifies the float32 path solves instead of +// returning the zero vector (the silent RawFloats-nil failure). +func TestGMRESFloat32RHS(t *testing.T) { + b, err := core.FromFloat32s([]float32{1, 2}, 2) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + d, err := core.FromFloats([]float64{2, 0, 0, 3}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + op := func(v *core.Array) (*core.Array, error) { return core.MatMul2D(d, v) } + x, err := GMRES(op, b, 0, 0, 1e-12) + if err != nil { + t.Fatalf("GMRES: %v", err) + } + want := []float64{0.5, 2.0 / 3.0} + for i := range 2 { + if math.Abs(x.FloatAt(i)-want[i]) > 1e-10 { + t.Fatalf("x[%d] = %g, want %g", i, x.FloatAt(i), want[i]) + } + } +} + +// TestSpEigenTinyScale pins the Lanczos deflation threshold against a +// matrix whose whole spectrum sits far below the old absolute floor: +// the two eigenvalues must stay distinct instead of collapsing into +// random Rayleigh quotients. +func TestSpEigenTinyScale(t *testing.T) { + const scale = 1e-13 + dense, err := core.FromFloats([]float64{2 * scale, scale, scale, 2 * scale}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.SparseFrom(dense) + if err != nil { + t.Fatalf("SparseFrom: %v", err) + } + vals, _, err := SpEigen(sp, 2, core.NewGenerator(99)) + if err != nil { + t.Fatalf("SpEigen: %v", err) + } + hi, lo := vals.FloatAt(0), vals.FloatAt(1) + if hi < lo { + hi, lo = lo, hi + } + if math.Abs(hi-3*scale) > 1e-14*scale/1e-13 || math.Abs(lo-1*scale) > 1e-14*scale/1e-13 { + t.Fatalf("eigenvalues (%g, %g), want (%g, %g)", hi, lo, 3*scale, 1*scale) + } +} + +// TestSpExpApplyTinyScale pins the Krylov exhaustion threshold: on a +// matrix of norm far below the old absolute floor, exp(A)·v must keep +// the first-order correction instead of collapsing to e^{α₁}·v. +func TestSpExpApplyTinyScale(t *testing.T) { + const s = 1e-13 + dense, err := core.FromFloats([]float64{2 * s, s, s, 2 * s}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.SparseFrom(dense) + if err != nil { + t.Fatalf("SparseFrom: %v", err) + } + v, err := core.FromFloats([]float64{1, 0}, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + out, err := SpExpApply(sp, v, 30) + if err != nil { + t.Fatalf("SpExpApply: %v", err) + } + // exp(s·[[2,1],[1,2]])·(1,0) = e^{2s}·(cosh s, sinh s) + // = (1+2s+O(s²), s+O(s²)); the old absolute floor collapsed this + // to (e^{2s}, 0), losing the whole first-order off-diagonal term. + // The factored form avoids the cancellation of (e^{3s}−e^s)/2. + want0 := math.Exp(2*s) * math.Cosh(s) + want1 := math.Exp(2*s) * math.Sinh(s) + if math.Abs(out.FloatAt(0)-want0) > 1e-15 || math.Abs(out.FloatAt(1)-want1) > 1e-24 { + t.Fatalf("exp(A)v = (%g, %g), want (%g, %g)", out.FloatAt(0), out.FloatAt(1), want0, want1) + } +} + +// TestEigenComplexTinyHermitian pins the Jacobi convergence threshold: +// a Hermitian matrix with norm below the old max(1, norm) floor must +// still be diagonalised, not returned as-is with identity vectors. +func TestEigenComplexTinyHermitian(t *testing.T) { + const e = 1e-14 + a, err := core.FromComplexes([]complex128{0, e, e, 0}, 2, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + vals, _, err := EigenComplex(a) + if err != nil { + t.Fatalf("EigenComplex: %v", err) + } + if math.Abs(math.Abs(vals.FloatAt(0))-e) > 1e-20 || math.Abs(math.Abs(vals.FloatAt(1))-e) > 1e-20 { + t.Fatalf("eigenvalues (%g, %g), want magnitudes %g", vals.FloatAt(0), vals.FloatAt(1), e) + } +} + +// TestBiCGSTABNaNOperator pins the NaN guard: an operator that +// produces a non-finite residual must fail loudly, never return the +// half-step iterate as a converged solution. +func TestBiCGSTABNaNOperator(t *testing.T) { + b, err := core.FromFloats([]float64{1, 1}, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + dense, err := core.FromFloats([]float64{1, 0, 0, math.NaN()}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.SparseFrom(dense) + if err != nil { + t.Fatalf("SparseFrom: %v", err) + } + if _, err := SpSolveBiCGSTAB(sp, b, 1e-10, 50); err == nil { + t.Fatal("BiCGSTAB reported convergence on a NaN operator") + } +} + +// TestFitPolynomialNegativeDegree pins the input validation: a +// negative degree must be an error, not a nil-dereference panic. +func TestFitPolynomialNegativeDegree(t *testing.T) { + x, err := core.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + y, err := core.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + for _, d := range []int{-1, -2, -5} { + if _, err := FitPolynomial(x, y, d); err == nil { + t.Fatalf("FitPolynomial accepted degree %d", d) + } + } +} diff --git a/linalg/solver_guard_pins_test.go b/linalg/solver_guard_pins_test.go new file mode 100644 index 0000000..76f6b5f --- /dev/null +++ b/linalg/solver_guard_pins_test.go @@ -0,0 +1,135 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression pins: NaN states that read as converged solves, and +// contract gaps the sibling solvers had already closed. + +func cooFrom(t *testing.T, idx []int64, vals []float64, shape []int) *core.SparseCOO { + t.Helper() + i, err := core.FromInts(idx, len(idx)/len(shape), len(shape)) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + v, err := core.FromFloats(vals, len(vals)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.NewSparseCOO(i, v, shape) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return sp +} + +// TestSpSolveOverflowCurvature: with every input finite, the Jacobi +// preconditioner overflowed the curvature to +Inf, Inf/Inf gave a NaN +// alpha, and the all-NaN residual read as a converged solve through +// the NaN-skipping norm. +func TestSpSolveOverflowCurvature(t *testing.T) { + sp := cooFrom(t, []int64{0, 0, 1, 1}, []float64{1e-155, 1e-155}, []int{2, 2}) + b, err := core.FromFloats([]float64{1e155, 1e155}, 2) + if err != nil { + t.Fatal(err) + } + x, err := SpSolve(sp, b, 0, 0) + if err == nil { + t.Fatalf("an overflowed solve returned %v with no error", x.FloatAt(0)) + } + if x != nil { + t.Fatalf("SpSolve returned a result beside the error") + } +} + +// TestSpSolveBiCGSTABOverflowOmega: a tiny diagonal entry overflowed +// the preconditioned stabiliser to Inf and its dots to NaN/Inf, omega +// became NaN, and the finiteness guard sat behind the convergence +// return it was written to protect. +func TestSpSolveBiCGSTABOverflowOmega(t *testing.T) { + // A = [[1e-155, 1], [1, 1]]: the Jacobi preconditioner divides the + // first residual by 1e-155 and the stabiliser leg overflows. + sp := cooFrom(t, []int64{0, 0, 0, 1, 1, 0, 1, 1}, + []float64{1e-155, 1, 1, 1}, []int{2, 2}) + b, err := core.FromFloats([]float64{1, 0}, 2) + if err != nil { + t.Fatal(err) + } + if _, err := SpSolveBiCGSTAB(sp, b, 0, 0); err == nil { + t.Fatal("an overflowed stabiliser returned a solution with no error") + } +} + +// TestGMRESRejectsRank2: every sibling entry point refuses a non-vector +// right-hand side; GMRES flattened it silently. +func TestGMRESRejectsRank2(t *testing.T) { + b, err := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatal(err) + } + ident := func(v *core.Array) (*core.Array, error) { return core.Copy(v), nil } + if _, err := GMRES(ident, b, 0, 0, 0); err == nil || !strings.Contains(err.Error(), "rank-1") { + t.Fatalf("GMRES on a rank-2 b: err = %v", err) + } +} + +// TestGMRESRejectsNonFiniteOp: an operator answering NaN must be an +// error, not a zero-norm residual that reads as convergence. +func TestGMRESRejectsNonFiniteOp(t *testing.T) { + b, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + nan := math.NaN() + broken := func(v *core.Array) (*core.Array, error) { return core.FromFloats([]float64{nan}, 1) } + if _, err := GMRES(broken, b, 0, 0, 0); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("GMRES with a NaN op: err = %v", err) + } +} + +// TestGMRESTinyOperator: the breakdown floor was relative to ‖b‖, so a +// legitimate system with ‖A‖ ≪ ‖b‖ collapsed every cycle at its first +// column and ran out of cycles without converging. +func TestGMRESTinyOperator(t *testing.T) { + b, err := core.FromFloats([]float64{1, 1}, 2) + if err != nil { + t.Fatal(err) + } + scale := 1e-20 + tiny := func(v *core.Array) (*core.Array, error) { + x := v.FloatAt(0) * scale + y := v.FloatAt(1) * scale + return core.FromFloats([]float64{x, y}, 2) + } + x, err := GMRES(tiny, b, 0, 0, 1e-8) + if err != nil { + t.Fatalf("GMRES on a tiny-norm operator: %v", err) + } + if math.Abs(x.FloatAt(0)-1/scale) > 1e-3/scale { + t.Fatalf("x[0] = %g, want %g", x.FloatAt(0), 1/scale) + } +} + +// TestNewCubicSplineRejectsRank2: a matrix input was silently +// reinterpreted as a flattened vector. +func TestNewCubicSplineRejectsRank2(t *testing.T) { + xs, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + ys, err := core.FromFloats([]float64{1, 4, 9, 16, 25, 36}, 2, 3) + if err != nil { + t.Fatal(err) + } + if _, err := NewCubicSpline(xs, ys); err == nil || !strings.Contains(err.Error(), "1-D") { + t.Fatalf("NewCubicSpline on rank-2 inputs: err = %v", err) + } +} diff --git a/linalg/sparse_guard_pins_test.go b/linalg/sparse_guard_pins_test.go new file mode 100644 index 0000000..5a08a47 --- /dev/null +++ b/linalg/sparse_guard_pins_test.go @@ -0,0 +1,519 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "fmt" + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression pins: a depth-first seed that ignored the live row +// permutation in the sparse LU, overflow states that slipped past the +// pivot and downdate guards, non-finite entries that sailed through +// the ILU intake and the symmetry screens, a minimum-degree +// absorption that handed a vertex its own index back, a zero +// hypotenuse in the complex SVD bulge chase, and a MatrixLog screen +// that refused positive definite matrices with small positive +// eigenvalues and a false "non-positive" verdict. + +// TestSparseLUPermutedSeedFactorisation pins the seed fix: once an +// earlier column has swapped rows, the depth-first search over column +// k must enter through the rows' current positions in the factor, not +// through the labels they were stored under. The 4×4 matrix below +// swaps original row 2 to position 0 and row 3 to position 1 before +// column 1 is eliminated, so column 1's stored rows 2 and 3 now sit +// at positions 0 and 3. +func TestSparseLUPermutedSeedFactorisation(t *testing.T) { + trips := [][3]float64{ + {0, 0, 1}, {0, 2, 1}, {0, 3, 1}, + {1, 0, 2}, {1, 2, 1}, {1, 3, 1}, + {2, 0, 5}, {2, 1, 7}, {2, 2, 8}, {2, 3, 1}, + {3, 1, 4}, {3, 2, 1}, {3, 3, 9}, + } + idx := make([]int64, 0, 2*len(trips)) + vals := make([]float64, 0, len(trips)) + for _, e := range trips { + idx = append(idx, int64(e[0]), int64(e[1])) + vals = append(vals, e[2]) + } + coo, err := core.NewSparseCOO(mustInts(t, idx, len(trips), 2), + floatsToArray(vals, []int{len(trips)}), []int{4, 4}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + f, err := NewSparseLU(coo) + if err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + if f.piv[0] != 2 || f.piv[1] != 3 { + t.Fatalf("permutation %v, want it to start [2 3 ...]", f.piv) + } + // The swapped row 2 (now position 0) carries A[2,1] = 7 into + // U row 0: the seed walk must have found its column. + found := false + for p, c := range f.rowCols[0] { + if c == 1 { + found = true + if math.Abs(f.rowVals[0][p]-7) > 1e-12 { + t.Fatalf("U[0,1] = %g, want 7", f.rowVals[0][p]) + } + } + } + if !found { + t.Fatal("U row 0 has no column-1 entry; the seed walk ignored the permutation") + } + // The factor must answer what the dense elimination answers. + x, err := f.Solve(mustFloats(t, []float64{8, 9, 47, 47})) + if err != nil { + t.Fatalf("Solve: %v", err) + } + want := []float64{1, 2, 3, 4} + for i := range 4 { + if math.Abs(x.FloatAt(i)-want[i]) > 1e-12 { + t.Fatalf("solve[%d] = %.15g, want %.15g", i, x.FloatAt(i), want[i]) + } + } +} + +// TestSparseLUOverflowPivotRefused pins the pivot guard: an +// elimination that squares the range (the update 1e308 − 1·1e308 +// reads −2e308, which rounds to −Inf) used to slide past the NaN and +// zero tests and store an infinite pivot without a word. +func TestSparseLUOverflowPivotRefused(t *testing.T) { + coo := cooFrom(t, []int64{0, 0, 0, 1, 1, 0, 1, 1}, + []float64{1e308, 1e308, 1e308, -1e308}, []int{2, 2}) + f, err := NewSparseLU(coo) + if err == nil { + t.Fatalf("NewSparseLU stored an infinite pivot (diag = [%g, %g]) with no error", f.diag[0], f.diag[1]) + } + if !strings.Contains(err.Error(), "overflow") { + t.Fatalf("error = %v, want an overflow refusal", err) + } +} + +// TestSparseCholeskyOverflowPivotRefused pins the same guard on the +// Cholesky side: L[1,0] = 1e100/1e-100 = 1e200 and the pivot update +// squares it to 1e400, which overflows. The matrix is indefinite, but +// the rounded arithmetic cannot honestly deliver a +// positive-definiteness verdict, and the report is the overflow. +func TestSparseCholeskyOverflowPivotRefused(t *testing.T) { + coo := cooFrom(t, []int64{0, 0, 0, 1, 1, 0, 1, 1}, + []float64{1e-200, 1e100, 1e100, 1e200}, []int{2, 2}) + _, err := NewSparseCholesky(coo, SparseOrderingNatural) + if err == nil { + t.Fatal("NewSparseCholesky stored an infinite pivot with no error") + } + if !strings.Contains(err.Error(), "overflow") { + t.Fatalf("error = %v, want an overflow refusal", err) + } +} + +// TestSparseILURefusesNonFinite pins the intake the complete sparse +// factorisations already have: a NaN or infinite entry is refused +// before the elimination turns it into a poisoned preconditioner. +func TestSparseILURefusesNonFinite(t *testing.T) { + idx := []int64{0, 0, 0, 1, 1, 0, 1, 1, 1, 2, 2, 1, 2, 2} + for name, bad := range map[string]float64{"NaN": math.NaN(), "Inf": math.Inf(1)} { + t.Run(name, func(t *testing.T) { + coo := cooFrom(t, idx, []float64{4, bad, -1, 4, -1, -1, 4}, []int{3, 3}) + if _, err := NewSparseILU(coo); err == nil { + t.Fatal("NewSparseILU accepted a non-finite entry") + } else if !strings.Contains(err.Error(), "not finite") { + t.Fatalf("error = %v, want a non-finite refusal", err) + } + }) + } + t.Run("finite", func(t *testing.T) { + coo := cooFrom(t, idx, []float64{4, 1, -1, 4, -1, -1, 4}, []int{3, 3}) + if _, err := NewSparseILU(coo); err != nil { + t.Fatalf("NewSparseILU refused a finite matrix: %v", err) + } + }) +} + +// TestMinimumDegreeAbsorptionDropsSelf pins the absorption rule: +// eliminating a vertex merges its adjacency into each surviving +// neighbour as adj(j) ∪ adj(p) \ {j}. The plain union hands j its own +// index back (it sits on p's list), which inflates the degree the +// selection scan reads and warps the order. +func TestMinimumDegreeAbsorptionDropsSelf(t *testing.T) { + // Path 0 - 1 - 2: adj = [[1], [0, 2], [1]]. + coo := cooFrom(t, []int64{0, 1, 1, 0, 1, 2, 2, 1, 0, 0, 1, 1, 2, 2}, + []float64{-1, -1, -1, -1, 4, 4, 4}, []int{3, 3}) + csc, err := CSCFromCOO(coo) + if err != nil { + t.Fatalf("CSCFromCOO: %v", err) + } + adj, err := symmetrisedAdjacency(csc) + if err != nil { + t.Fatalf("symmetrisedAdjacency: %v", err) + } + eliminated := []bool{true, false, false} + degree := []int{1, 2, 1} + frontier := °reeFrontier{} + absorbElement(adj, degree, eliminated, 0, frontier, &intArena{}) + for _, u := range adj[1] { + if u == 1 { + t.Fatalf("adj[1] = %v lists 1 itself after absorbing 0's element", adj[1]) + } + } + if degree[1] != 1 { + t.Fatalf("vertex 1 keeps degree %d after the absorption, want 1", degree[1]) + } + deg := 0 + for _, u := range adj[1] { + if !eliminated[u] { + deg++ + } + } + if deg != 1 { + t.Fatalf("vertex 1 reads degree %d after the absorption, want 1", deg) + } + if !slices.IsSorted(adj[1]) { + t.Fatalf("adj[1] = %v is not the sorted form the union expects", adj[1]) + } + // With honest degrees, vertices 1 and 2 tie at degree 1 and the + // tie breaks to the smaller index. + order, err := minimumDegree(csc) + if err != nil { + t.Fatalf("minimumDegree: %v", err) + } + if !slices.Equal(order, []int{0, 1, 2}) { + t.Fatalf("minimum degree order %v, want [0 1 2]", order) + } +} + +// TestSVDComplexRankOneRectangular pins the bulge chase's zero step: +// when a deflated diagonal meets a deflated bulge the rotation is the +// identity, where d[k]/0 raised NaNs that spilled through every +// factor. The rank-1 rectangular matrices below reach exactly that +// step. +func TestSVDComplexRankOneRectangular(t *testing.T) { + t.Run("2x3explicit", func(t *testing.T) { + raw := []complex128{1, 2, 3, 2, 4, 6} + a := mustComplex(t, raw, 2, 3) + u, sigma, vh, err := SVDComplex(a) + if err != nil { + t.Fatalf("SVDComplex: %v", err) + } + want := math.Sqrt(70.0) + if rel := math.Abs(sigma.FloatAt(0)-want) / want; rel > 1e-12 { + t.Fatalf("sigma[0] = %.15g, want %.15g (relative error %.3g)", sigma.FloatAt(0), want, rel) + } + if rel := sigma.FloatAt(1) / want; rel > 1e-12 { + t.Fatalf("sigma[1] = %.15g, want 0 (relative %.3g)", sigma.FloatAt(1), rel) + } + checkUnitaryRows(t, u, sigma, vh, raw, 2, 3) + }) + for _, shape := range [][2]int{{2, 3}, {2, 4}, {2, 5}, {3, 2}, {3, 4}, {4, 2}, {5, 2}} { + t.Run(fmt.Sprintf("%dx%d", shape[0], shape[1]), func(t *testing.T) { + m, n := shape[0], shape[1] + a := make([]complex128, m*n) + for i := range m { + for j := range n { + ui := complex(float64(i+1), float64(i)) + vj := complex(float64(-j), float64(j+1)) + a[i*n+j] = ui * vj + } + } + aa, err := core.FromComplexes(a, m, n) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + u, sigma, vh, err := SVDComplex(aa) + if err != nil { + t.Fatalf("SVDComplex: %v", err) + } + // A rank-1 matrix carries its whole Frobenius norm in the + // leading singular value and nothing in the rest. + fro := 0.0 + for _, z := range a { + fro += real(z)*real(z) + imag(z)*imag(z) + } + fro = math.Sqrt(fro) + if rel := math.Abs(sigma.FloatAt(0)-fro) / fro; rel > 1e-12 { + t.Fatalf("sigma[0] = %.15g, want %.15g (relative error %.3g)", sigma.FloatAt(0), fro, rel) + } + if rel := sigma.FloatAt(1) / fro; rel > 1e-12 { + t.Fatalf("sigma[1] = %.15g, want 0 (relative %.3g)", sigma.FloatAt(1), rel) + } + for k := 1; k < sigma.Len(); k++ { + if math.IsNaN(sigma.FloatAt(k)) || math.IsInf(sigma.FloatAt(k), 0) || sigma.FloatAt(k) < 0 { + t.Fatalf("sigma[%d] = %g is not a plain non-negative number", k, sigma.FloatAt(k)) + } + } + checkUnitaryRows(t, u, sigma, vh, a, m, n) + }) + } +} + +// checkUnitaryRows verifies the factors a rank-decomposed SVD must +// return: Uᴴ·U = I on the thin U, orthonormal rows of Vᴴ, and the +// reconstruction U·diag(σ)·Vᴴ back to a. +func checkUnitaryRows(t *testing.T, u, sigma, vh *core.Array, a []complex128, m, n int) { + t.Helper() + r := min(m, n) + uc := u.RawComplexes() + ucols := u.Shape()[1] + urows := u.Shape()[0] + for j := range r { + for i := range r { + s := complex(0, 0) + for k := range urows { + s += cmplxConjTest(uc[k*ucols+i]) * uc[k*ucols+j] + } + want := 0.0 + if i == j { + want = 1 + } + if d := cmplxAbsTest(s - complex(want, 0)); d > 1e-12 { + t.Fatalf("(Uᴴ·U)[%d,%d] = %v, want %v", i, j, s, want) + } + } + } + vc := vh.RawComplexes() + vcols := vh.Shape()[1] + for j := range r { + for i := range r { + s := complex(0, 0) + for k := range vcols { + s += vc[i*vcols+k] * cmplxConjTest(vc[j*vcols+k]) + } + want := 0.0 + if i == j { + want = 1 + } + if d := cmplxAbsTest(s - complex(want, 0)); d > 1e-12 { + t.Fatalf("(Vᴴ·(Vᴴ)ᴴ)[%d,%d] = %v, want %v", i, j, s, want) + } + } + } + worst := 0.0 + for i := range m { + for j := range n { + s := complex(0, 0) + for k := range r { + s += uc[i*ucols+k] * complex(sigma.FloatAt(k), 0) * vc[k*vcols+j] + } + if d := cmplxAbsTest(s - a[i*n+j]); d > worst { + worst = d + } + } + } + if worst > 1e-12*sigma.FloatAt(0) { + t.Fatalf("reconstruction error %g exceeds %g", worst, 1e-12*sigma.FloatAt(0)) + } +} + +func cmplxConjTest(z complex128) complex128 { return complex(real(z), -imag(z)) } + +func cmplxAbsTest(z complex128) float64 { return math.Hypot(real(z), imag(z)) } + +// TestMatrixLogSymmetricSmallSpectrum pins the symmetric route's +// negativity floor: a positive eigenvalue is logged whatever its size +// against the spectrum, where the old 1e-10 relative screen refused +// fine positive eigenvalues with a false "non-positive" verdict. +func TestMatrixLogSymmetricSmallSpectrum(t *testing.T) { + tiny := mustFloats(t, []float64{1, 0, 0, 1e-12}, 2, 2) + lg, err := MatrixLog(tiny) + if err != nil { + t.Fatalf("MatrixLog(diag(1, 1e-12)): %v", err) + } + for i, w := range []float64{0, math.Log(1e-12)} { + if math.Abs(lg.FloatAt(i*2+i)-w) > 1e-9 { + t.Fatalf("ln A[%d][%d] = %.16g, want %.16g", i, i, lg.FloatAt(i*2+i), w) + } + } + // The smaller eigenvalue is 1e-16 of the larger, far below any + // relative screen, and still logs honestly. + big := mustFloats(t, []float64{1e20, 0, 0, 1e4}, 2, 2) + lg, err = MatrixLog(big) + if err != nil { + t.Fatalf("MatrixLog(diag(1e20, 1e4)): %v", err) + } + for i, w := range []float64{math.Log(1e20), math.Log(1e4)} { + if rel := math.Abs(lg.FloatAt(i*2+i)-w) / math.Abs(w); rel > 1e-12 { + t.Fatalf("ln A[%d][%d] = %.16g, want %.16g", i, i, lg.FloatAt(i*2+i), w) + } + } + for _, i := range []int{1, 2} { + if math.Abs(lg.FloatAt(i)) > 1e-6 { + t.Fatalf("ln A off-diagonal [%d] = %g, want 0", i, lg.FloatAt(i)) + } + } + // A genuinely negative eigenvalue is still refused. + neg := mustFloats(t, []float64{0, 1, 1, 0}, 2, 2) + if _, err := MatrixLog(neg); err == nil { + t.Fatal("MatrixLog([[0,1],[1,0]]): want an error for the negative eigenvalue") + } + // MatrixSqrt keeps its own screen: positive spectra root, negative + // spectra refuse, nothing moved. + sq, err := MatrixSqrt(tiny) + if err != nil { + t.Fatalf("MatrixSqrt(diag(1, 1e-12)): %v", err) + } + for i, w := range []float64{1, 1e-6} { + if math.Abs(sq.FloatAt(i*2+i)-w) > 1e-12 { + t.Fatalf("√A[%d][%d] = %.16g, want %.16g", i, i, sq.FloatAt(i*2+i), w) + } + } + sq, err = MatrixSqrt(big) + if err != nil { + t.Fatalf("MatrixSqrt(diag(1e20, 1e4)): %v", err) + } + for i, w := range []float64{1e10, 100} { + if rel := math.Abs(sq.FloatAt(i*2+i)-w) / w; rel > 1e-12 { + t.Fatalf("√A[%d][%d] = %.16g, want %.16g", i, i, sq.FloatAt(i*2+i), w) + } + } + if _, err := MatrixSqrt(neg); err == nil { + t.Fatal("MatrixSqrt of an indefinite matrix: want an error") + } +} + +// TestCholeskyDowndateOverflowRefused pins the downdate's range test: +// with a diagonal at 1e160 the squared terms overflow, and r2 reads +// NaN (Inf − Inf) or +Inf depending on the vector, both of which used +// to slide past the cone test into a silent NaN or infinite factor. +func TestCholeskyDowndateOverflowRefused(t *testing.T) { + l := mustFloats(t, []float64{1e160}, 1, 1) + for name, xv := range map[string]float64{"NaN": 1e160, "Inf": 1} { + t.Run(name, func(t *testing.T) { + out, err := CholeskyDowndate(l, mustFloats(t, []float64{xv})) + if err == nil { + t.Fatalf("downdate returned a factor with diagonal %g and no error", out.FloatAt(0)) + } + if !strings.Contains(err.Error(), "overflow") { + t.Fatalf("error = %v, want an overflow refusal", err) + } + }) + } + // A rank-one update and a downdate inside the range are untouched. + small := mustFloats(t, []float64{4}, 1, 1) + up, err := CholeskyUpdate(small, mustFloats(t, []float64{1})) + if err != nil { + t.Fatalf("CholeskyUpdate: %v", err) + } + if math.Abs(up.FloatAt(0)-math.Sqrt(17)) > 1e-12 { + t.Fatalf("update diagonal %g, want %.15g", up.FloatAt(0), math.Sqrt(17)) + } + down, err := CholeskyDowndate(small, mustFloats(t, []float64{1})) + if err != nil { + t.Fatalf("CholeskyDowndate: %v", err) + } + if math.Abs(down.FloatAt(0)-math.Sqrt(15)) > 1e-12 { + t.Fatalf("downdate diagonal %g, want %.15g", down.FloatAt(0), math.Sqrt(15)) + } +} + +// TestSparseSymmetryChecksRefuseNonFinite pins the symmetry and +// Hermitian screens: a NaN compares unequal to everything, so the +// mirror test alone waved non-finite entries through as symmetric. +func TestSparseSymmetryChecksRefuseNonFinite(t *testing.T) { + idx := []int64{0, 0, 0, 1, 1, 0, 1, 1} + t.Run("symmetricNaN", func(t *testing.T) { + coo := cooFrom(t, idx, []float64{4, math.NaN(), math.NaN(), 4}, []int{2, 2}) + c, err := cooToCSR(coo, "test") + if err != nil { + t.Fatalf("cooToCSR: %v", err) + } + if err := c.checkSymmetric("SpSolve"); err == nil { + t.Fatal("checkSymmetric accepted a NaN entry") + } else if !strings.Contains(err.Error(), "not finite") { + t.Fatalf("error = %v, want a non-finite refusal", err) + } + }) + t.Run("symmetricInf", func(t *testing.T) { + coo := cooFrom(t, idx, []float64{4, math.Inf(1), math.Inf(1), 4}, []int{2, 2}) + c, err := cooToCSR(coo, "test") + if err != nil { + t.Fatalf("cooToCSR: %v", err) + } + if err := c.checkSymmetric("SpSolve"); err == nil { + t.Fatal("checkSymmetric accepted an infinite entry") + } else if !strings.Contains(err.Error(), "not finite") { + t.Fatalf("error = %v, want a non-finite refusal", err) + } + }) + t.Run("hermitianNaN", func(t *testing.T) { + coo := cooComplexFrom(t, idx, []complex128{4, cmplxNaN(), cmplxNaN(), 4}, []int{2, 2}) + c, err := cooToComplexCSR(coo, "test") + if err != nil { + t.Fatalf("cooToComplexCSR: %v", err) + } + if err := c.checkHermitian("SpSolveComplexCG"); err == nil { + t.Fatal("checkHermitian accepted a NaN entry") + } else if !strings.Contains(err.Error(), "not finite") { + t.Fatalf("error = %v, want a non-finite refusal", err) + } + }) + t.Run("hermitianInf", func(t *testing.T) { + inf := complex(math.Inf(1), 0) + coo := cooComplexFrom(t, idx, []complex128{4, inf, inf, 4}, []int{2, 2}) + c, err := cooToComplexCSR(coo, "test") + if err != nil { + t.Fatalf("cooToComplexCSR: %v", err) + } + if err := c.checkHermitian("SpSolveComplexCG"); err == nil { + t.Fatal("checkHermitian accepted an infinite entry") + } else if !strings.Contains(err.Error(), "not finite") { + t.Fatalf("error = %v, want a non-finite refusal", err) + } + }) + t.Run("legalInputs", func(t *testing.T) { + coo := cooFrom(t, idx, []float64{4, 1, 1, 4}, []int{2, 2}) + c, err := cooToCSR(coo, "test") + if err != nil { + t.Fatalf("cooToCSR: %v", err) + } + if err := c.checkSymmetric("SpSolve"); err != nil { + t.Fatalf("checkSymmetric refused a finite symmetric matrix: %v", err) + } + hcoo := cooComplexFrom(t, idx, []complex128{4, 1 + 2i, 1 - 2i, 4}, []int{2, 2}) + hc, err := cooToComplexCSR(hcoo, "test") + if err != nil { + t.Fatalf("cooToComplexCSR: %v", err) + } + if err := hc.checkHermitian("SpSolveComplexCG"); err != nil { + t.Fatalf("checkHermitian refused a finite Hermitian matrix: %v", err) + } + }) + t.Run("publicEntries", func(t *testing.T) { + sp := cooFrom(t, idx, []float64{4, math.NaN(), math.NaN(), 4}, []int{2, 2}) + if _, err := SpSolve(sp, mustFloats(t, []float64{1, 1}), 0, 0); err == nil { + t.Fatal("SpSolve accepted a NaN entry") + } + hsp := cooComplexFrom(t, idx, []complex128{4, cmplxNaN(), cmplxNaN(), 4}, []int{2, 2}) + b := mustComplex(t, []complex128{1, 1}, 2) + if _, err := SpSolveComplexCG(hsp, b, 0, 0); err == nil { + t.Fatal("SpSolveComplexCG accepted a NaN entry") + } + }) +} + +// cooComplexFrom builds a complex-valued sparse COO, failing the test +// on a bad shape. +func cooComplexFrom(t *testing.T, idx []int64, vals []complex128, shape []int) *core.SparseCOO { + t.Helper() + i, err := core.FromInts(idx, len(idx)/len(shape), len(shape)) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + v, err := core.FromComplexes(vals, len(vals)) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + sp, err := core.NewSparseCOO(i, v, shape) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return sp +} + +func cmplxNaN() complex128 { return complex(math.NaN(), math.NaN()) } diff --git a/linalg/sparsebicg.go b/linalg/sparsebicg.go new file mode 100644 index 0000000..f2af694 --- /dev/null +++ b/linalg/sparsebicg.go @@ -0,0 +1,142 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Nonsymmetric sparse solve. SpSolve's conjugate gradient leans on +// symmetry for its short recurrences; the convection terms that make +// systems nonsymmetric break that structure, and the fix is van der +// Vorst's BiCGSTAB: two nested three-term recurrences whose second +// half smooths the erratic convergence the plain biconjugate gradient +// shows. Every step costs two sparse products and two preconditioner +// applications, the same order as one CG step pair. +// +// The Jacobi preconditioner needs a nonzero diagonal, which for the +// systems this solver targets (shifted Laplacians with convection, +// discretised transport) is present. Breakdowns (the recurrences +// collapsing) are reported as errors with the step they happened at, +// never swallowed. + +// SpSolveBiCGSTAB returns the vector x solving A·x = b for a general +// real square sparse A, by preconditioned BiCGSTAB. Unlike SpSolve it +// asks for no symmetry; like SpSolve the stopping rule is the +// relative residual ‖b − A·x‖₂ ≤ tol·‖b‖₂ (tol ≤ 0 means 1e-10), +// maxIter ≤ 0 means n steps, and an unconverged solve is an error +// naming the residual achieved. The default preconditioner is Jacobi +// scaling, which divides by the diagonal and refuses a zero or +// missing entry; passing an ILU(0) factorisation from NewSparseILU +// replaces it and usually cuts the step count on hard systems. +func SpSolveBiCGSTAB(a *core.SparseCOO, b *core.Array, tol float64, maxIter int, precond ...*SparseILU) (*core.Array, error) { + const name = "SpSolveBiCGSTAB" + n, err := checkSparseSquare(name, a, b) + if err != nil { + return nil, err + } + c, err := cooToCSR(a, name) + if err != nil { + return nil, err + } + ilu, diag, err := pickPreconditioner(name, c, precond) + if err != nil { + return nil, err + } + if tol <= 0 { + tol = spSolveTol + } + if maxIter <= 0 { + maxIter = n + } + + r := vectorF64(b, n) + bNorm := norm2F64(r) + if bNorm == 0 { + return core.Zeros(core.Float, n) + } + rHat := append([]float64(nil), r...) + x := make([]float64, n) + v := make([]float64, n) + p := make([]float64, n) + s := make([]float64, n) + t := make([]float64, n) + z := make([]float64, n) + y := make([]float64, n) + mt := make([]float64, n) + + rho := 1.0 + alpha := 1.0 + omega := 1.0 + for iter := range maxIter { + rhoNext := dotF64(rHat, r) + if !finiteF64(rhoNext) || rhoNext == 0 { + return nil, base.Errf("SpSolveBiCGSTAB: breakdown at step %d (rho vanished)", iter+1) + } + if iter > 0 { + beta := rhoNext / rho * (alpha / omega) + for i := range n { + p[i] = r[i] + beta*(p[i]-omega*v[i]) + } + } else { + copy(p, r) + } + iluPrecondition(z, p, ilu, diag) + c.matVec(z, v) + rhatV := dotF64(rHat, v) + if rhatV == 0 { + return nil, base.Errf("SpSolveBiCGSTAB: breakdown at step %d (direction orthogonal to shadow)", iter+1) + } + alpha = rhoNext / rhatV + for i := range n { + s[i] = r[i] - alpha*v[i] + } + if !vecFinite(s) { + return nil, base.Errf("SpSolveBiCGSTAB: breakdown at step %d (non-finite residual)", iter+1) + } + // The half step may already be exact. The check is affirmative so + // a NaN norm (which fails every comparison) can never read as + // converged. + if norm2F64(s) > tol*bNorm { + iluPrecondition(y, s, ilu, diag) + c.matVec(y, t) + iluPrecondition(mt, t, ilu, diag) + mtS := dotF64(mt, s) + mtT := dotF64(mt, t) + if mtT == 0 { + return nil, base.Errf("SpSolveBiCGSTAB: breakdown at step %d (stabiliser vanished)", iter+1) + } + omega = mtS / mtT + // The finiteness half of the guard belongs before the + // update: a NaN omega poisons x and r below, and the + // all-NaN residual then reads as an exact solve through + // the NaN-skipping norm. + if !finiteF64(omega) { + return nil, base.Errf("SpSolveBiCGSTAB: breakdown at step %d (stabiliser not finite)", iter+1) + } + for i := range n { + x[i] += alpha*z[i] + omega*y[i] + r[i] = s[i] - omega*t[i] + } + if !vecFinite(r) { + return nil, base.Errf("SpSolveBiCGSTAB: non-finite residual at step %d", iter+1) + } + if norm2F64(r) <= tol*bNorm { + return floatsToArray(x, []int{n}), nil + } + if omega == 0 { + return nil, base.Errf("SpSolveBiCGSTAB: stagnation at step %d (omega zero)", iter+1) + } + } else { + for i := range n { + x[i] += alpha * z[i] + } + return floatsToArray(x, []int{n}), nil + } + rho = rhoNext + } + return nil, base.Errf("SpSolveBiCGSTAB: no convergence in %d steps, residual %.3g (tolerance %.3g)", + maxIter, norm2F64(r), tol*bNorm) +} diff --git a/linalg/sparsebicg_test.go b/linalg/sparsebicg_test.go new file mode 100644 index 0000000..45af9d3 --- /dev/null +++ b/linalg/sparsebicg_test.go @@ -0,0 +1,195 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestSpSolveBiCGSTABNonsymmetric checks the solver against the dense +// Solve on a 4×4 system with genuinely asymmetric entries. +func TestSpSolveBiCGSTABNonsymmetric(t *testing.T) { + dense := [][]float64{ + {10, 1, 0, 2}, + {-1, 12, 3, 0}, + {0, -2, 15, 1}, + {2, 0, -1, 8}, + } + entries := make([][3]float64, 0, 16) + for i := range 4 { + for j := range 4 { + if dense[i][j] != 0 { + entries = append(entries, [3]float64{float64(i), float64(j), dense[i][j]}) + } + } + } + idx := make([]int64, 0, len(entries)) + vals := make([]float64, 0, len(entries)) + for _, e := range entries { + idx = append(idx, int64(e[0]), int64(e[1])) + vals = append(vals, e[2]) + } + indices, err := core.FromInts(idx, len(entries), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(entries)}), []int{4, 4}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + b := mustFloats(t, []float64{1, 2, 3, 4}) + x, err := SpSolveBiCGSTAB(coo, b, 1e-12, 0) + if err != nil { + t.Fatalf("SpSolveBiCGSTAB: %v", err) + } + ref, err := Solve(mustFloats(t, []float64{ + 10, 1, 0, 2, + -1, 12, 3, 0, + 0, -2, 15, 1, + 2, 0, -1, 8, + }, 4, 4), b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + for i := range 4 { + if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-9 { + t.Fatalf("x[%d] = %.14g, want %.14g", i, x.FloatAt(i), ref.FloatAt(i)) + } + } +} + +// TestSpSolveBiCGSTABConvective builds a 50×50 shifted Laplacian with +// an asymmetric convection term, the shape transport discretisations +// produce, and checks the residual the solver promises. +func TestSpSolveBiCGSTABConvective(t *testing.T) { + const n = 50 + idx := make([]int64, 0, 4*n) + vals := make([]float64, 0, 4*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 4) + if i+1 < n { + add(i, i+1, -1) + add(i+1, i, -1+0.5) // asymmetric neighbour coupling + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + bv := make([]float64, n) + for i := range n { + bv[i] = float64(i+1) / float64(n+1) + } + b := mustFloats(t, bv) + x, err := SpSolveBiCGSTAB(coo, b, 1e-10, 0) + if err != nil { + t.Fatalf("SpSolveBiCGSTAB: %v", err) + } + // Residual check against the sparse product path. + ax, err := cooToCSRForTest(t, coo).MatVec(x) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + res := 0.0 + for i := range n { + d := b.FloatAt(i) - ax.FloatAt(i) + res += d * d + } + if math.Sqrt(res) > 1e-9*math.Sqrt(dotF64(bv, bv)) { + t.Fatalf("relative residual %g exceeds 1e-9", math.Sqrt(res)/math.Sqrt(dotF64(bv, bv))) + } +} + +// cooToCSRForTest converts a COO to the public CSR for residual checks. +func cooToCSRForTest(t *testing.T, coo *core.SparseCOO) *SparseCSR { + t.Helper() + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + return csr +} + +// TestSpSolveBiCGSTABConvergenceFailure reports an honest failure when +// the budget cannot reach the tolerance. +func TestSpSolveBiCGSTABConvergenceFailure(t *testing.T) { + const n = 30 + idx := make([]int64, 0, 3*n) + vals := make([]float64, 0, 3*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 4) + if i+1 < n { + add(i, i+1, -1) + add(i+1, i, -0.5) + } + } + indices, _ := core.FromInts(idx, len(vals), 2) + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + b := mustFloats(t, make([]float64, n)) + for i := range n { + b.RawFloats()[i] = 1 + } + if _, err := SpSolveBiCGSTAB(coo, b, 1e-14, 2); err == nil { + t.Fatal("tight tolerance with a two-step budget: want an error") + } +} + +// TestSpSolveBiCGSTABErrors pins the validation surface. +func TestSpSolveBiCGSTABErrors(t *testing.T) { + // Build the COO inputs directly for the error paths. + mkCOO := func(t *testing.T, rows, cols int, entries [][3]float64) *core.SparseCOO { + t.Helper() + idx := make([]int64, 0, len(entries)*2) + vals := make([]float64, 0, len(entries)) + for _, e := range entries { + idx = append(idx, int64(e[0]), int64(e[1])) + vals = append(vals, e[2]) + } + indices, err := core.FromInts(idx, len(entries), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(entries)}), []int{rows, cols}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo + } + b2 := mustFloats(t, []float64{1, 1}) + if _, err := SpSolveBiCGSTAB(mkCOO(t, 2, 3, [][3]float64{{0, 0, 2}}), b2, 0, 0); err == nil { + t.Fatal("non-square matrix: want an error") + } + // A missing diagonal entry defeats the Jacobi preconditioner. + gap := mkCOO(t, 2, 2, [][3]float64{{0, 1, 1}, {1, 1, 3}}) + if _, err := SpSolveBiCGSTAB(gap, b2, 0, 0); err == nil { + t.Fatal("missing diagonal: want an error") + } + bad := mkCOO(t, 2, 2, [][3]float64{{0, 0, 2}, {1, 1, 3}}) + if _, err := SpSolveBiCGSTAB(bad, mustFloats(t, []float64{1, 2, 3}), 0, 0); err == nil { + t.Fatal("wrong right-hand side length: want an error") + } + complexVals := core.New(core.Complex, 2) + cx := mkCOO(t, 2, 2, [][3]float64{{0, 0, 1}, {1, 1, 1}}) + cx.Values = complexVals + if _, err := SpSolveBiCGSTAB(cx, b2, 0, 0); err == nil { + t.Fatal("complex values: want an error") + } +} diff --git a/linalg/sparsecholesky.go b/linalg/sparsecholesky.go new file mode 100644 index 0000000..456f2c8 --- /dev/null +++ b/linalg/sparsecholesky.go @@ -0,0 +1,524 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "cmp" + "math" + "slices" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The sparse Cholesky factorisation: the elimination tree (Liu's +// construction with path compression) discovers the fill of every row +// cheaply, and the row-oriented elimination that follows is the +// up-looking form of the direct-methods standard (George and Liu; +// Davis). The factor exists only for symmetric positive definite +// matrices, where it is the direct solver of choice: no pivoting, no +// fill beyond the pattern the ordering decides. + +// SparseOrdering selects the fill-reducing permutation a sparse +// factorisation applies before it eliminates. +type SparseOrdering int + +const ( + // SparseOrderingNatural eliminates rows and columns in stored + // order. Only right for matrices whose pattern is already dense + // near the diagonal, and the reference point every ordering is + // measured against. + SparseOrderingNatural SparseOrdering = iota + // SparseOrderingReverseCuthillMcKee orders every component in + // reverse breadth first order from a pseudo-peripheral start: the + // classic choice for mesh-shaped patterns, where it keeps the + // factor banded and the fill small. + SparseOrderingReverseCuthillMcKee + // SparseOrderingMinimumDegree eliminates the vertex with the + // fewest remaining neighbours at every step, absorbing each + // eliminated pattern into its neighbours: the stronger choice on + // irregular patterns, where a bandwidth order leaves the factor + // dense in the middle. + SparseOrderingMinimumDegree +) + +// rowValuePair carries a factor entry through a column sort: the row +// index travels with its value. +type rowValuePair struct { + row int + val float64 +} + +// SparseCholesky carries the factorisation of a symmetric positive +// definite matrix: L with A permuted to the elimination order, so +// that P·A·Pᵀ = L·Lᵀ. One factorisation solves any number of right +// hand sides. +type SparseCholesky struct { + // perm[k] is the original index of the row eliminated k-th; + // inversePerm maps an original index back to its position. + perm []int + inversePerm []int + n int + // L in CSC form: column j holds the off-diagonal rows i > j in + // ascending order, and diag carries L[j][j]. + colStart []int + rowIdx []int + values []float64 + diag []float64 +} + +// NewSparseCholesky factors a symmetric positive definite matrix. The +// lower triangle defines the matrix: a stored upper triangle entry is +// refused unless its lower counterpart is stored with exactly the +// same value, so a silently asymmetric input cannot be factored. The +// ordering argument picks the fill-reducing permutation; see +// SparseOrdering. +func NewSparseCholesky(a *core.SparseCOO, ordering SparseOrdering) (*SparseCholesky, error) { + const name = "NewSparseCholesky" + if a.Values.Dtype() == core.Complex { + return nil, base.Errf("%s: complex sparse matrices are not supported", name) + } + if len(a.Shape) != 2 || a.Shape[0] != a.Shape[1] { + return nil, base.Errf("%s: needs a square 2-D sparse matrix, got shape %v", name, a.Shape) + } + c, err := CSCFromCOO(a) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + i := c.RowIdx[p] + v := c.Values[p] + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: entry [%d,%d] is not finite", name, i, j) + } + // The lower triangle is the authority: an upper entry + // must mirror a stored lower entry bit for bit. + if i < j { + q, ok := findEntry(c, j, i) + if !ok || c.Values[q] != v { + return nil, base.Errf("%s: entry [%d,%d] = %g has no matching lower counterpart [%d,%d]", + name, i, j, v, j, i) + } + } + } + } + var perm []int + switch ordering { + case SparseOrderingNatural: + perm = make([]int, c.Cols) + for i := range perm { + perm[i] = i + } + case SparseOrderingReverseCuthillMcKee: + perm, err = reverseCuthillMcKee(c) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + case SparseOrderingMinimumDegree: + perm, err = minimumDegree(c) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + default: + return nil, base.Errf("%s: unknown ordering %d", name, int(ordering)) + } + f := &SparseCholesky{perm: perm, n: c.Cols} + f.inversePerm = make([]int, c.Cols) + for k, i := range perm { + f.inversePerm[i] = k + } + lower := permutedLower(c, perm) + rows := lowerRows(lower) + if err := f.factor(lower, rows); err != nil { + return nil, base.Errf("%s: %w", name, err) + } + return f, nil +} + +// lowerRows returns the row view of a lower triangular matrix: row k +// holds its columns in ascending order, diagonal included. The full +// transpose would carry the upper entries into the rows, which the +// elimination must not see. +func lowerRows(c *SparseCSC) *SparseCSR { + rows := &SparseCSR{Rows: c.Cols, Cols: c.Cols, RowStart: make([]int, c.Cols+1)} + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + if c.RowIdx[p] >= j { + rows.RowStart[c.RowIdx[p]+1]++ + } + } + } + for i := range c.Cols { + rows.RowStart[i+1] += rows.RowStart[i] + } + rows.ColIdx = make([]int, rows.RowStart[c.Cols]) + rows.Values = make([]float64, rows.RowStart[c.Cols]) + next := make([]int, c.Cols) + copy(next, rows.RowStart[:c.Cols]) + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + i := c.RowIdx[p] + if i < j { + continue + } + q := next[i] + rows.ColIdx[q] = j + rows.Values[q] = c.Values[p] + next[i] = q + 1 + } + } + return rows +} + +// findEntry looks a stored entry (i, j) up by binary search; the +// canonical form keeps every column's row indices sorted. +func findEntry(c *SparseCSC, i, j int) (int, bool) { + at, found := slices.BinarySearch(c.RowIdx[c.ColStart[j]:c.ColStart[j+1]], i) + if !found { + return 0, false + } + return c.ColStart[j] + at, true +} + +// permutedLower builds the lower triangle of P·A·Pᵀ in CSC form: new +// position k carries the original row perm[k]. Permuting does not +// keep the triangle entry for entry, so every stored lower entry (i, +// j) is placed at the unordered pair of new positions: row max of +// inversePerm[i] and inversePerm[j], column min. The symmetry check +// in NewSparseCholesky has already pinned each upper entry to its +// lower counterpart, so walking the lower triangle alone visits +// every pair once. +func permutedLower(c *SparseCSC, perm []int) *SparseCSC { + inversePerm := make([]int, c.Cols) + for k, i := range perm { + inversePerm[i] = k + } + counts := make([]int, c.Cols+1) + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + i := c.RowIdx[p] + if i < j { + continue + } + // The entry lands below the new diagonal: row max of the + // two new positions, column min. + b := min(inversePerm[j], inversePerm[i]) + counts[b+1]++ + } + } + for i := range c.Cols { + counts[i+1] += counts[i] + } + out := &SparseCSC{Rows: c.Cols, Cols: c.Cols, ColStart: counts} + out.RowIdx = make([]int, counts[c.Cols]) + out.Values = make([]float64, counts[c.Cols]) + next := make([]int, c.Cols) + copy(next, counts[:c.Cols]) + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + i := c.RowIdx[p] + if i < j { + continue + } + ki, kj := inversePerm[i], inversePerm[j] + b := ki + r := kj + if kj < ki { + b = kj + r = ki + } + q := next[b] + out.RowIdx[q] = r + out.Values[q] = c.Values[p] + next[b] = q + 1 + } + } + // Each new column receives entries from several old columns, so + // the rows must be sorted into the canonical ascending order, + // values travelling with their rows. One scratch buffer serves + // every column: it grows to the longest unsorted column once + // instead of allocating per column. The row indices within a + // column are unique, so the sorted placement is unique and the + // scratch cannot move a value. + pairsScratch := make([]rowValuePair, 0, c.Cols) + for b := range c.Cols { + segment := out.RowIdx[counts[b]:counts[b+1]] + if slices.IsSorted(segment) { + continue + } + pairs := pairsScratch[:0] + for p := counts[b]; p < counts[b+1]; p++ { + pairs = append(pairs, rowValuePair{out.RowIdx[p], out.Values[p]}) + } + slices.SortFunc(pairs, func(x, y rowValuePair) int { return cmp.Compare(x.row, y.row) }) + for p, e := range pairs { + out.RowIdx[counts[b]+p] = e.row + out.Values[counts[b]+p] = e.val + } + pairsScratch = pairs + } + return out +} + +// eliminationTree returns the parent of every column in the +// elimination tree of the lower triangular pattern, Liu's +// construction with path compression: for every stored entry (k, j), +// j < k, the root of j's tree is hung under k, so parent[i] > i and +// the chain from any j whose row reaches k ends exactly at k. +// Unparented columns carry −1. +func eliminationTree(n int, rows *SparseCSR) []int { + parent := make([]int, n) + ancestor := make([]int, n) + for i := range ancestor { + ancestor[i] = -1 + } + for k := range n { + parent[k] = -1 + for p := rows.RowStart[k]; p < rows.RowStart[k+1]; p++ { + j := rows.ColIdx[p] + if j >= k { + continue + } + r := j + for r != -1 && r != k { + next := ancestor[r] + ancestor[r] = k + if next == -1 { + parent[r] = k + break + } + r = next + } + } + } + return parent +} + +// rowReach collects the update columns of row k: every stored entry +// (k, j) walked up its parent chain until the walk reaches k. The +// chain can also run through positions whose working entry stays +// zero, which subtracts nothing later, so the superset costs +// arithmetic, never correctness. +// +// The result is sorted ascending, and the walk keeps it that way as +// it grows: a chain climbs, so its next vertex usually lands past the +// tail and appends, and one that lands below inserts into its sorted +// place. The set's order is part of the elimination's arithmetic: the +// columns below are processed in it, and two of them can push into +// one working entry. +func rowReach(rows *SparseCSR, parent []int, k int, mark []bool, set []int) []int { + set = rowReachUnsorted(rows, parent, k, mark, set) + if !slices.IsSorted(set) { + slices.Sort(set) + } + return set +} + +// rowReachUnsorted collects the same set without ordering it. The +// counting pass reads the set's length alone, where the walk order +// costs nothing, so half the reach walks pay no sort. +func rowReachUnsorted(rows *SparseCSR, parent []int, k int, mark []bool, set []int) []int { + set = set[:0] + for p := rows.RowStart[k]; p < rows.RowStart[k+1]; p++ { + j := rows.ColIdx[p] + for j != k && j != -1 && !mark[j] { + mark[j] = true + set = append(set, j) + j = parent[j] + } + } + for _, j := range set { + mark[j] = false + } + return set +} + +// cholColumnCounts returns the exact number of entries every column of +// the factor receives, by running the same reach walk the numeric +// elimination runs and counting the columns it visits. The numeric pass +// appends an entry to column j for every j in the reach of a row k +// except where a working entry has cancelled to zero, which only makes +// the real count smaller, and both passes call the same pure walk on +// the same elimination tree. The counts are therefore an exact upper +// bound: the arena below needs no growth and leaves no gap between +// columns. The walk runs unsorted here: the counts read the set's +// size, which no order changes. +// +// The result is returned in prefix-sum form: column j's entries start +// at counts[j], and counts[n] is the factor's off-diagonal entry count. +func cholColumnCounts(n int, rows *SparseCSR, parent []int, mark []bool, set []int) []int { + counts := make([]int, n+1) + for k := range n { + set = rowReachUnsorted(rows, parent, k, mark, set) + for _, j := range set { + counts[j+1]++ + } + } + for j := range n { + counts[j+1] += counts[j] + } + return counts +} + +// factor runs the row-oriented elimination. Row k of L is its +// working row divided through the square-rooted diagonal, the update +// columns come from row k's stored entries walked through the +// elimination tree, and each settled entry L[k,j] pushes its +// contribution down column j of the factor, the update the left-looking +// recurrence demands: L[k,i] loses L[k,j]·L[i,j] for every earlier +// column j, so the columns are walked ascending and every working +// entry has settled by the time its turn comes. The elimination +// tree's guarantee makes the walk safe: every stored entry (k, j) +// sits on j's parent chain, and a chain position whose working entry +// stays zero simply subtracts nothing. +// +// The factor is written straight into one flat arena, a contiguous +// block per column with a cursor, so the whole numeric factor costs two +// allocations and the scatter into the working row reads one base array +// instead of a slice header per column. A column's entries land in +// ascending k, the order the solver reads. +func (f *SparseCholesky) factor(lower *SparseCSC, rows *SparseCSR) error { + const name = "NewSparseCholesky" + n := lower.Cols + parent := eliminationTree(n, rows) + f.diag = make([]float64, n) + x := make([]float64, n) + set := make([]int, 0, n) + mark := make([]bool, n) + f.colStart = cholColumnCounts(n, rows, parent, mark, set) + arenaRows := make([]int, f.colStart[n]) + arenaVals := make([]float64, f.colStart[n]) + // next[j] is column j's cursor: where its next entry lands. It + // starts at the column's block and ends one past its last written + // entry, so next[j]−colStart[j] is the column's actual length. + next := make([]int, n) + copy(next, f.colStart[:n]) + for k := range n { + // x holds row k of the working matrix: the stored entries + // minus the updates below. + rs, re := rows.RowStart[k], rows.RowStart[k+1] + ci := rows.ColIdx[rs:re:re] + cv := rows.Values[rs:re:re] + for i, c := range ci { + x[c] = cv[i] + } + set = rowReach(rows, parent, k, mark, set) + for _, j := range set { + dj := f.diag[j] + lkj := x[j] / dj + if lkj != 0 { + at := next[j] + arenaRows[at] = k + arenaVals[at] = lkj + next[j] = at + 1 + } + // Push the entry down column j: the diagonal term zeroes + // x[j], the rows below carry the fill, and the k term is + // the diagonal sum of this very row. The column is walked + // through slices cut to its written range, so the loop + // carries no bound checks and no reload of the cursor + // next[j] the writes above could be thought to alias. + x[j] -= lkj * dj + cs, ce := f.colStart[j], next[j] + rw := arenaRows[cs:ce:ce] + rv := arenaVals[cs:ce:ce] + for i, r := range rw { + x[r] -= lkj * rv[i] + } + } + d := x[k] + if math.IsInf(d, 0) { + return base.Errf("%s: the pivot at [%d,%d] overflowed to %g; the elimination has left the float64 range", + name, f.perm[k], f.perm[k], d) + } + if math.IsNaN(d) || d <= 0 { + return base.Errf("%s: the pivot at [%d,%d] is %g; the matrix is not positive definite", + name, f.perm[k], f.perm[k], d) + } + f.diag[k] = math.Sqrt(d) + // Restore every x position the row touched: the update + // columns, their factor entries, and the diagonal. + for _, j := range set { + x[j] = 0 + cs, ce := f.colStart[j], next[j] + for _, r := range arenaRows[cs:ce:ce] { + x[r] = 0 + } + } + x[k] = 0 + } + // Pack the columns: a column whose working entry cancelled along + // the way sits shorter than its block, so only the written range + // moves, and a column's entries keep their ascending-k order. + total := 0 + for j := range n { + total += next[j] - f.colStart[j] + } + colStart := make([]int, n+1) + f.rowIdx = make([]int, total) + f.values = make([]float64, total) + for j := range n { + b, w := f.colStart[j], next[j] + colStart[j+1] = colStart[j] + (w - b) + copy(f.rowIdx[colStart[j]:], arenaRows[b:w]) + copy(f.values[colStart[j]:], arenaVals[b:w]) + } + f.colStart = colStart + return nil +} + +// Solve computes x = A⁻¹·b for a dense vector b: permute, forward +// substitution with L, backward substitution with Lᵀ, undo the +// permutation. +func (f *SparseCholesky) Solve(b *core.Array) (*core.Array, error) { + const name = "Solve" + if b.NDim() != 1 { + return nil, base.Errf("%s: the right hand side must be rank 1", name) + } + if b.Dtype() == core.Complex { + return nil, base.Errf("%s: complex right hand sides are not supported", name) + } + if b.Len() != f.n { + return nil, base.Errf("%s: right hand side length %d does not match %d rows", name, b.Len(), f.n) + } + x := core.New(core.Float, []int{f.n}...) + xf := x.RawFloats() + for k := range f.n { + xf[k] = b.FloatAt(f.perm[k]) + } + for j := range f.n { + xj := xf[j] / f.diag[j] + xf[j] = xj + for p := f.colStart[j]; p < f.colStart[j+1]; p++ { + xf[f.rowIdx[p]] -= f.values[p] * xj + } + } + for j := f.n - 1; j >= 0; j-- { + sum := xf[j] + for p := f.colStart[j]; p < f.colStart[j+1]; p++ { + sum -= f.values[p] * xf[f.rowIdx[p]] + } + xf[j] = sum / f.diag[j] + } + // z solves P·A·Pᵀ·z = P·b, so z[k] is the answer at original + // position perm[k]: the scatter walks the same side as the + // gather. + out := core.New(core.Float, []int{f.n}...) + of := out.RawFloats() + for k := range f.n { + of[f.perm[k]] = xf[k] + } + return out, nil +} + +// Permutation returns the elimination order: position k holds the +// original index factored k-th, with P·A·Pᵀ = L·Lᵀ. The slice is a +// copy, so the caller cannot move the factor's state. +func (f *SparseCholesky) Permutation() []int { + return slices.Clone(f.perm) +} + +// NNZ returns the count of stored non-zeros in the factor, diagonal +// included. +func (f *SparseCholesky) NNZ() int { return len(f.values) + f.n } diff --git a/linalg/sparsecholesky_test.go b/linalg/sparsecholesky_test.go new file mode 100644 index 0000000..ab99fda --- /dev/null +++ b/linalg/sparsecholesky_test.go @@ -0,0 +1,535 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// gridLaplacianCOO builds the 5-point Laplacian on a w by h grid: the +// standard sparse positive definite test matrix, symmetric with a +// dominant diagonal. When shuffle is true the grid vertices are +// relabelled by a fixed seeded permutation first, so the stored order +// carries no bandwidth for the natural ordering to lean on. +func gridLaplacianCOOShuffled(t *testing.T, w, h int, shuffle bool) *core.SparseCOO { + t.Helper() + idx := make([]int64, 0, 5*w*h) + vals := make([]float64, 0, 5*w*h) + label := func(x, y int) int { return y*w + x } + if shuffle { + g := core.NewGenerator(5) + perm := make([]int, w*h) + for i := range perm { + perm[i] = i + } + // Fisher-Yates with the seeded generator, a fixed relabelling. + for i := len(perm) - 1; i > 0; i-- { + j := int(g.Next() % uint64(i+1)) + perm[i], perm[j] = perm[j], perm[i] + } + label = func(x, y int) int { return perm[y*w+x] } + } + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + at := func(x, y int) int { return label(x, y) } + for y := range h { + for x := range w { + add(at(x, y), at(x, y), 4) + if x+1 < w { + add(at(x, y), at(x+1, y), -1) + add(at(x+1, y), at(x, y), -1) + } + if y+1 < h { + add(at(x, y), at(x, y+1), -1) + add(at(x, y+1), at(x, y), -1) + } + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{w * h, w * h}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// gridLaplacianCOO builds the plain row-major grid. +func gridLaplacianCOO(t *testing.T, w, h int) *core.SparseCOO { + t.Helper() + return gridLaplacianCOOShuffled(t, w, h, false) +} + +func TestSparseCholeskySolvesLaplacian(t *testing.T) { + const w, h = 12, 10 + coo := gridLaplacianCOO(t, w, h) + n := w * h + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + // The truth is constructive: xTrue is fixed, b comes from the + // independently tested CSR MatVec, and the factor must invert it. + xTrue := core.New(core.Float, n) + for i := range n { + xTrue.RawFloats()[i] = math.Sin(float64(i)) + float64(i%7)*0.1 + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + for _, ordering := range []SparseOrdering{SparseOrderingNatural, SparseOrderingReverseCuthillMcKee} { + f, err := NewSparseCholesky(coo, ordering) + if err != nil { + t.Fatalf("NewSparseCholesky(%d): %v", ordering, err) + } + x, err := f.Solve(b) + if err != nil { + t.Fatalf("Solve(%d): %v", ordering, err) + } + worst := 0.0 + for i := range n { + if d := math.Abs(x.FloatAt(i) - xTrue.FloatAt(i)); d > worst { + worst = d + } + } + if worst > 1e-9 { + t.Fatalf("ordering %d: worst solution error %.3g, want under 1e-9", ordering, worst) + } + // The residual through the original matrix closes the loop: + // A·x must give b back. + ax, err := csr.MatVec(x) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + res := 0.0 + for i := range n { + if d := math.Abs(ax.FloatAt(i) - b.FloatAt(i)); d > res { + res = d + } + } + if res > 1e-8 { + t.Fatalf("ordering %d: residual %.3g, want under 1e-8", ordering, res) + } + } +} + +// irregularSPDCOO builds a diagonally dominant symmetric matrix with +// long-range random couplings: the pattern no bandwidth ordering can +// tame, the case the minimum degree order exists for. +func irregularSPDCOO(t *testing.T, seed int64, n int) *core.SparseCOO { + t.Helper() + g := core.NewGenerator(seed) + deg := make([]float64, n) + type entry struct { + r, c int + } + var entries []entry + for range 700 { + i := int(g.Next() % uint64(n)) + j := int(g.Next() % uint64(n)) + if i == j { + continue + } + entries = append(entries, entry{i, j}) + deg[i]++ + deg[j]++ + } + idx := make([]int64, 0, 2*len(entries)+2*n) + vals := make([]float64, 0, 2*len(entries)+n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for _, e := range entries { + add(e.r, e.c, -1) + add(e.c, e.r, -1) + } + for i := range n { + add(i, i, float64(deg[i])+2) + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// TestSparseCholeskyMinimumDegreeOrdering pins the minimum degree +// order's property: on both a shuffled mesh and a long-range +// irregular pattern its factor stores fewer non-zeros than the +// reverse Cuthill-McKee factor, and the solve stays exact. The +// measured numbers sit far inside the bounds and land in the log. +func TestSparseCholeskyMinimumDegreeOrdering(t *testing.T) { + coo := gridLaplacianCOOShuffled(t, 15, 15, true) + nat, err := NewSparseCholesky(coo, SparseOrderingNatural) + if err != nil { + t.Fatalf("NewSparseCholesky(natural): %v", err) + } + rcm, err := NewSparseCholesky(coo, SparseOrderingReverseCuthillMcKee) + if err != nil { + t.Fatalf("NewSparseCholesky(rcm): %v", err) + } + md, err := NewSparseCholesky(coo, SparseOrderingMinimumDegree) + if err != nil { + t.Fatalf("NewSparseCholesky(md): %v", err) + } + t.Logf("grid n=225: natural %d, rcm %d, md %d", nat.NNZ(), rcm.NNZ(), md.NNZ()) + if md.NNZ() >= rcm.NNZ() { + t.Fatalf("md fill %d is not below rcm fill %d on the grid", md.NNZ(), rcm.NNZ()) + } + // Long-range irregular pattern: both orders must still solve, and + // the minimum degree order must keep its edge. + irr := irregularSPDCOO(t, 11, 300) + csr, err := CSRFromCOO(irr) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + xTrue := core.New(core.Float, 300) + for i := range 300 { + xTrue.RawFloats()[i] = math.Cos(0.3*float64(i)) + float64(i%7) + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + rcmI, err := NewSparseCholesky(irr, SparseOrderingReverseCuthillMcKee) + if err != nil { + t.Fatalf("NewSparseCholesky(rcm): %v", err) + } + mdI, err := NewSparseCholesky(irr, SparseOrderingMinimumDegree) + if err != nil { + t.Fatalf("NewSparseCholesky(md): %v", err) + } + t.Logf("irregular n=300: rcm %d, md %d", rcmI.NNZ(), mdI.NNZ()) + if mdI.NNZ() >= rcmI.NNZ() { + t.Fatalf("md fill %d is not below rcm fill %d on the irregular pattern", mdI.NNZ(), rcmI.NNZ()) + } + for name, f := range map[string]*SparseCholesky{"rcm": rcmI, "md": mdI} { + x, err := f.Solve(b) + if err != nil { + t.Fatalf("%s solve: %v", name, err) + } + for i := range 300 { + if math.Abs(x.FloatAt(i)-xTrue.FloatAt(i)) > 1e-9 { + t.Fatalf("%s: entry %d error %.3g", name, i, math.Abs(x.FloatAt(i)-xTrue.FloatAt(i))) + } + } + } + perm := md.Permutation() + sorted := slices.Clone(perm) + slices.Sort(sorted) + for i := range sorted { + if sorted[i] != i { + t.Fatalf("permutation entry %d holds %d; not a permutation", i, sorted[i]) + } + } +} + +// TestSparseCholeskyMatchesDense factors the same grid Laplacian as a +// dense matrix and requires the two solvers to agree: the sparse +// factorisation is a different algorithm for the same A⁻¹. +func TestSparseCholeskyMatchesDense(t *testing.T) { + const w, h = 8, 8 + coo := gridLaplacianCOO(t, w, h) + n := w * h + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + xTrue := core.New(core.Float, n) + for i := range n { + xTrue.RawFloats()[i] = math.Cos(0.3*float64(i)) + float64(i%5) + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + f, err := NewSparseCholesky(coo, SparseOrderingReverseCuthillMcKee) + if err != nil { + t.Fatalf("NewSparseCholesky: %v", err) + } + sparse, err := f.Solve(b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + denseVals := make([]float64, n*n) + nnz := coo.Indices.Shape()[0] + for i := range nnz { + r := int(coo.Indices.RawInts()[i*2]) + c := int(coo.Indices.RawInts()[i*2+1]) + denseVals[r*n+c] = coo.Values.FloatAt(i) + } + dense, err := Solve(floatsToArray(denseVals, []int{n, n}), b) + if err != nil { + t.Fatalf("dense Solve: %v", err) + } + scale := 0.0 + for i := range n { + if d := math.Abs(dense.FloatAt(i)); d > scale { + scale = d + } + } + for i := range n { + if math.Abs(sparse.FloatAt(i)-dense.FloatAt(i)) > 1e-9*scale { + t.Fatalf("entry %d: sparse %.12g vs dense %.12g", i, sparse.FloatAt(i), dense.FloatAt(i)) + } + } +} + +// TestSparseCholeskyFillMeasures pins the property the orderings +// exist for: on a shuffled mesh pattern the reverse Cuthill-McKee +// factor must store markedly fewer non-zeros than the natural order +// factor. The bounds carry slack on purpose; the measured numbers +// land far inside them and the log line records them, so a regression +// in the ordering shows up as a test failure, not as a slow solver. +func TestSparseCholeskyFillMeasures(t *testing.T) { + coo := gridLaplacianCOOShuffled(t, 15, 15, true) + natural, err := NewSparseCholesky(coo, SparseOrderingNatural) + if err != nil { + t.Fatalf("NewSparseCholesky(natural): %v", err) + } + rcm, err := NewSparseCholesky(coo, SparseOrderingReverseCuthillMcKee) + if err != nil { + t.Fatalf("NewSparseCholesky(rcm): %v", err) + } + t.Logf("n=225 shuffled: natural L nnz %d, rcm L nnz %d", natural.NNZ(), rcm.NNZ()) + if rcm.NNZ() >= natural.NNZ() { + t.Fatalf("rcm fill %d is not below natural fill %d", rcm.NNZ(), natural.NNZ()) + } + if rcm.NNZ() > natural.NNZ()/2 { + t.Fatalf("rcm fill %d did not at least halve natural fill %d", rcm.NNZ(), natural.NNZ()) + } + // Both factors must still solve: the ordering changes the fill, + // never the answer. + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + xTrue := core.New(core.Float, 225) + for i := range xTrue.Len() { + xTrue.RawFloats()[i] = math.Cos(0.7*float64(i)) + float64(i%11)*0.2 + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + for name, f := range map[string]*SparseCholesky{"natural": natural, "rcm": rcm} { + x, err := f.Solve(b) + if err != nil { + t.Fatalf("%s solve: %v", name, err) + } + for i := range xTrue.Len() { + if math.Abs(x.FloatAt(i)-xTrue.FloatAt(i)) > 1e-9 { + t.Fatalf("%s: entry %d error %.3g", name, i, math.Abs(x.FloatAt(i)-xTrue.FloatAt(i))) + } + } + } + perm := rcm.Permutation() + sorted := slices.Clone(perm) + slices.Sort(sorted) + for i := range sorted { + if sorted[i] != i { + t.Fatalf("permutation entry %d holds %d; not a permutation", i, sorted[i]) + } + } + // The permutation is a copy: moving the caller's slice must not + // move the factor's. + perm[0] = -1 + if rcm.Permutation()[0] == -1 { + t.Fatal("Permutation exposed the factor's internal slice") + } +} + +func TestSparseCholeskyRefusals(t *testing.T) { + good := gridLaplacianCOO(t, 4, 4) + if _, err := NewSparseCholesky(good, SparseOrdering(7)); err == nil { + t.Fatal("an unknown ordering was accepted") + } + // A stored upper entry without its lower counterpart is a silent + // asymmetry the factor refuses to inherit. Entries: (0,0)=4, + // (0,1)=1, (1,1)=4, (1,2)=1, (2,2)=4, (2,3)=1, (3,3)=4: the upper + // (0,1) has no lower (1,0). + oneWay, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3}, 7, 2), + floatsToArray([]float64{4, 1, 4, 1, 4, 1, 4}, []int{7}), + []int{4, 4}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseCholesky(oneWay, SparseOrderingNatural); err == nil || !strings.Contains(err.Error(), "counterpart") { + t.Fatalf("one-sided upper entry: %v", err) + } + // A conflicting pair refuses too. Entries: (0,0)=4, (0,1)=1, + // (1,1)=4, (1,2)=2, (2,2)=4, (1,3)=3, (3,3)=4, (2,1)=5: the upper + // (1,2)=2 and the lower (2,1)=5 disagree. + conflicting, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 1, 3, 3, 3, 2, 1}, 8, 2), + floatsToArray([]float64{4, 1, 4, 2, 4, 3, 4, 5}, []int{8}), + []int{4, 4}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseCholesky(conflicting, SparseOrderingNatural); err == nil || !strings.Contains(err.Error(), "counterpart") { + t.Fatalf("conflicting upper entry: %v", err) + } + // Not positive definite: the zero diagonal has no square root. + sing := triDiagCOO(t, 4, 1, 0, 1) + if _, err := NewSparseCholesky(sing, SparseOrderingNatural); err == nil || !strings.Contains(err.Error(), "positive definite") { + t.Fatalf("zero diagonal: %v", err) + } + // Non-finite stored value. + bad, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0, 1, 1}, 2, 2), + floatsToArray([]float64{4, math.NaN()}, []int{2}), + []int{2, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseCholesky(bad, SparseOrderingNatural); err == nil || !strings.Contains(err.Error(), "finite") { + t.Fatalf("NaN entry: %v", err) + } + // Rectangular input. + rect, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0}, 1, 2), + floatsToArray([]float64{1}, []int{1}), + []int{1, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseCholesky(rect, SparseOrderingNatural); err == nil { + t.Fatal("a rectangular matrix was accepted") + } + // Solve-side refusals. + f, err := NewSparseCholesky(good, SparseOrderingNatural) + if err != nil { + t.Fatalf("NewSparseCholesky: %v", err) + } + if _, err := f.Solve(core.New(core.Float, 3, 3)); err == nil { + t.Fatal("a rank 2 right hand side was accepted") + } + if _, err := f.Solve(floatsToArray([]float64{1, 2, 3}, []int{3})); err == nil { + t.Fatal("a short right hand side was accepted") + } +} + +// TestSparseCholeskyIsDeterministic factors the same matrix twice and +// requires the factor's stored values to come out bit for bit equal, +// the contract every Tensor entry point carries. +func TestSparseCholeskyIsDeterministic(t *testing.T) { + coo := gridLaplacianCOO(t, 9, 9) + f1, err := NewSparseCholesky(coo, SparseOrderingReverseCuthillMcKee) + if err != nil { + t.Fatalf("first factorisation: %v", err) + } + f2, err := NewSparseCholesky(coo, SparseOrderingReverseCuthillMcKee) + if err != nil { + t.Fatalf("second factorisation: %v", err) + } + if f1.NNZ() != f2.NNZ() { + t.Fatalf("factor sizes differ: %d vs %d", f1.NNZ(), f2.NNZ()) + } + for i := range f1.values { + if f1.values[i] != f2.values[i] { + t.Fatalf("value %d differs: %.17g vs %.17g", i, f1.values[i], f2.values[i]) + } + } + if slices.Compare(f1.perm, f2.perm) != 0 { + t.Fatal("permutations differ") + } +} + +// TestCSCCanonicalisation checks CSCFromCOO's contract beside +// CSRFromCOO's: duplicates sum, explicit zeros drop, every column's +// row indices are sorted and unique, and the transpose round trips +// agree entry for entry with the direct conversions. +func TestCSCCanonicalisation(t *testing.T) { + // Entries: (2,0)=3, (2,0)=1 duplicate, (1,1)=5, (0,2)=7 explicit + // zero, (0,2)=2. + coo, err := core.NewSparseCOO( + mustInts(t, []int64{2, 0, 2, 0, 1, 1, 0, 2, 0, 2}, 5, 2), + floatsToArray([]float64{3, 1, 5, 0, 2}, []int{5}), + []int{3, 3}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + csc, err := CSCFromCOO(coo) + if err != nil { + t.Fatalf("CSCFromCOO: %v", err) + } + if csc.NNZ() != 3 { + t.Fatalf("nnz %d after merging and zero drop, want 3", csc.NNZ()) + } + if csc.Values[0] != 4 { + t.Fatalf("duplicates did not sum: column 0 first value %g", csc.Values[0]) + } + for j := range csc.Cols { + rows := csc.RowIdx[csc.ColStart[j]:csc.ColStart[j+1]] + if !slices.IsSorted(rows) { + t.Fatalf("column %d rows not sorted: %v", j, rows) + } + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + round, err := csc.ToCSR() + if err != nil { + t.Fatalf("ToCSR: %v", err) + } + if slices.Compare(csr.RowStart, round.RowStart) != 0 || + slices.Compare(csr.ColIdx, round.ColIdx) != 0 || + slices.Compare(csr.Values, round.Values) != 0 { + t.Fatal("CSC to CSR round trip disagrees with the direct conversion") + } + back, err := csr.ToCSC() + if err != nil { + t.Fatalf("ToCSC: %v", err) + } + if slices.Compare(csc.ColStart, back.ColStart) != 0 || + slices.Compare(csc.RowIdx, back.RowIdx) != 0 || + slices.Compare(csc.Values, back.Values) != 0 { + t.Fatal("CSR to CSC round trip disagrees with the direct conversion") + } + // MatVec over the CSC must answer what the CSR answers. + x := floatsToArray([]float64{1, -2, 3}, []int{3}) + ycsr, err := csr.MatVec(x) + if err != nil { + t.Fatalf("CSR MatVec: %v", err) + } + ycsc, err := csc.MatVec(x) + if err != nil { + t.Fatalf("CSC MatVec: %v", err) + } + for i := range ycsr.Len() { + if ycsr.FloatAt(i) != ycsc.FloatAt(i) { + t.Fatalf("MatVec row %d: csr %.17g vs csc %.17g", i, ycsr.FloatAt(i), ycsc.FloatAt(i)) + } + } + if _, err := csc.MatVec(floatsToArray([]float64{1, 2}, []int{2})); err == nil { + t.Fatal("a wrong-length vector was accepted") + } +} + +func mustInts(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 +} diff --git a/linalg/sparsecholupdate.go b/linalg/sparsecholupdate.go new file mode 100644 index 0000000..d11077e --- /dev/null +++ b/linalg/sparsecholupdate.go @@ -0,0 +1,160 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Rank-one modification of the elimination-tree sparse Cholesky +// factor, the sparse counterpart of `CholeskyUpdate` and +// `CholeskyDowndate`. Refactorising from scratch reruns the whole +// ordering and elimination; a rank-one change of the matrix needs one +// sweep of (hyperbolic for the downdate) rotations down the stored +// columns, the sweep the dense contract applies to a full n×n factor. +// +// The honest contract of the sparse form is narrower than the dense +// one. A rank-one change of A generally changes the elimination +// pattern, so the updated factor generally does not live on the old +// one's fill; no rotation sweep can invent the missing entries. What +// these methods do is transform the stored NUMERIC factor on the +// EXISTING pattern: the sweep refuses, with a clear error, the moment +// the update would need an entry the pattern does not hold, and it +// leaves the factor untouched when it does. Where the pattern +// suffices, and it does for the bread-and-butter cases of a supported +// vector on a banded or otherwise already-rich pattern, the result is +// the exact factor of the modified matrix; otherwise the fallback is a +// fresh `NewSparseCholesky` of the explicitly modified matrix. + +// Update applies the rank-one update A ← A + x·xᵀ to the factor in +// place: afterwards the factor is the Cholesky factor of the modified +// matrix, on the same pattern and in the same elimination order. x is +// a vector in the original coordinates, the coordinates `Solve` takes, +// and is gathered through the permutation internally. +// +// The sweep runs the dense rotation column by column, touching only +// the stored entries of each column. If the rotation would write a +// value into a position the pattern does not store, the update is +// refused and the factor keeps its original numbers, mirroring the +// dense contract where the input factor is left intact on failure; a +// fresh factorisation of the modified matrix is the fallback. A +// non-finite entry, a complex vector or a wrong length is refused +// outright. +func (f *SparseCholesky) Update(x *core.Array) error { + return f.cholRankOneUpdate(x, "SparseCholesky.Update", +1) +} + +// Downdate applies the rank-one downdate A ← A − x·xᵀ to the factor in +// place, with Update's pattern contract. Each hyperbolic rotation +// squares the remaining diagonal against the working vector: when a +// diagonal can no longer dominate, the modified matrix has left the +// positive definite cone and the downdate is an error, exactly the +// dense refusal; when the square leaves the float64 range the downdate +// refuses rather than poisoning the factor. +func (f *SparseCholesky) Downdate(x *core.Array) error { + return f.cholRankOneUpdate(x, "SparseCholesky.Downdate", -1) +} + +// cholRankOneUpdate carries the shared sweep on a copy of the numeric +// factor, committing only when every column has settled, so a refusal +// anywhere leaves the receiver untouched. +func (f *SparseCholesky) cholRankOneUpdate(x *core.Array, name string, sign int) error { + if x.Dtype() == core.Complex { + return base.Errf("%s: complex inputs are not supported", name) + } + if x.NDim() != 1 || x.Len() != f.n { + return base.Errf("%s: the vector must have length %d, got shape %s", + name, f.n, base.ShapeText(x.Shape())) + } + // The factor lives in the elimination order: the working vector is + // x gathered through the permutation, the same gather `Solve` + // applies to a right-hand side. + v := make([]float64, f.n) + for k := range f.n { + z := x.FloatAt(f.perm[k]) + if math.IsNaN(z) || math.IsInf(z, 0) { + return base.Errf("%s: entry %d is not finite", name, f.perm[k]) + } + v[k] = z + } + vals := slices.Clone(f.values) + diag := slices.Clone(f.diag) + // touched holds every position the sweep has seen non-zero, so the + // fill check scans the support of the working vector instead of all + // n positions of every rotated column. The support only grows along + // stored entries, which keeps the check proportional to what the + // sweep actually touches. + touched := make([]bool, f.n) + nz := make([]int, 0, f.n) + for k := range f.n { + if v[k] != 0 { + touched[k] = true + nz = append(nz, k) + } + } + for k := range f.n { + if v[k] == 0 { + // The rotation on a zero pivot is the identity, c = 1 and + // s = 0, needs no fill and writes nothing: the skip is + // bit-exact against the dense sweep. + continue + } + // The fill check. The rotation writes ±s·v[i] into every + // position (i, k) with v[i] ≠ 0, so one such position outside + // the stored column would create fill the pattern does not + // hold. The column's rows are sorted, so the membership test is + // a binary search. + for _, i := range nz { + if i <= k || v[i] == 0 { + continue + } + if _, ok := slices.BinarySearch(f.rowIdx[f.colStart[k]:f.colStart[k+1]], i); ok { + continue + } + return base.Errf("%s: the modification needs a factor entry at permuted position [%d,%d] (original [%d,%d]), which the stored pattern does not hold; factorise the modified matrix afresh instead", + name, i, k, f.perm[i], f.perm[k]) + } + d := diag[k] + var r, c, s float64 + if sign > 0 { + r = math.Hypot(d, v[k]) + if !finiteF64(r) { + return base.Errf("%s: the update at column %d (original %d) overflows the float64 range", name, k, f.perm[k]) + } + c, s = d/r, v[k]/r + } else { + r2 := d*d - v[k]*v[k] + if math.IsNaN(r2) || math.IsInf(r2, 0) { + return base.Errf("%s: the downdate at column %d (original %d) overflows the float64 range", name, k, f.perm[k]) + } + if r2 <= 0 { + return base.Errf("%s: the modified matrix is not positive definite at column %d (original %d)", name, k, f.perm[k]) + } + r = math.Sqrt(r2) + c, s = d/r, v[k]/r + } + diag[k] = r + for p := f.colStart[k]; p < f.colStart[k+1]; p++ { + i := f.rowIdx[p] + lik, xi := vals[p], v[i] + if sign > 0 { + vals[p] = c*lik + s*xi + } else { + vals[p] = c*lik - s*xi + } + v[i] = c*xi - s*lik + if !touched[i] { + touched[i] = true + nz = append(nz, i) + } + } + } + f.values = vals + f.diag = diag + return nil +} diff --git a/linalg/sparsecholupdate_test.go b/linalg/sparsecholupdate_test.go new file mode 100644 index 0000000..cf08384 --- /dev/null +++ b/linalg/sparsecholupdate_test.go @@ -0,0 +1,314 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// unitVectorCOO builds a coordinate vector with a single non-zero. +func unitVector(t *testing.T, n, pos int, val float64) *core.Array { + t.Helper() + vals := make([]float64, n) + vals[pos] = val + return floatsToArray(vals, []int{n}) +} + +// denseSparseNorm returns the largest absolute difference of two +// factor states: the diagonal and the stored values in order. +func factorDelta(f, g *SparseCholesky) float64 { + worst := 0.0 + for i := range f.n { + if d := math.Abs(f.diag[i] - g.diag[i]); d > worst { + worst = d + } + } + for i := range f.values { + if d := math.Abs(f.values[i] - g.values[i]); d > worst { + worst = d + } + } + return worst +} + +// addRankOneCOO returns the coordinate matrix of A + x·xᵀ, x given in +// the stored coordinates, by summing the modified entries on top of +// the original ones. +func addRankOneCOO(t *testing.T, a *core.SparseCOO, x *core.Array) *core.SparseCOO { + t.Helper() + nnz := a.Indices.Shape()[0] + entries := make(map[[2]int]float64) + for i := range nnz { + r := int(a.Indices.RawInts()[i*2]) + c := int(a.Indices.RawInts()[i*2+1]) + entries[[2]int{r, c}] += a.Values.FloatAt(i) + } + for i := range x.Len() { + if x.FloatAt(i) == 0 { + continue + } + for j := range x.Len() { + if x.FloatAt(j) == 0 { + continue + } + entries[[2]int{i, j}] += x.FloatAt(i) * x.FloatAt(j) + } + } + idx := make([]int64, 0, 2*len(entries)) + vals := make([]float64, 0, len(entries)) + for k, v := range entries { + idx = append(idx, int64(k[0]), int64(k[1])) + vals = append(vals, v) + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), a.Shape) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// TestSparseCholeskyUpdateDowndateRoundTrip pins the round trip on the +// banded case the contract names: a path-graph Laplacian, whose +// natural and reverse Cuthill-McKee factors are both banded, with +// supported updates whose sweeps stay on the stored pattern. Update +// then downdate must restore the factor, and the updated factor must +// solve like a fresh factorisation of the modified matrix. +func TestSparseCholeskyUpdateDowndateRoundTrip(t *testing.T) { + coo := gridLaplacianCOO(t, 12, 1) + n := 12 + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + b, err := csr.MatVec(floatsToArray([]float64{1, -2, 3, -4, 5, -1, 2, -3, 4, -5, 1, 0}, []int{n})) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + for _, ordering := range []SparseOrdering{SparseOrderingNatural, SparseOrderingReverseCuthillMcKee} { + for name, x := range map[string]*core.Array{ + "single": unitVector(t, n, 3, 0.5), + "adjacent": floatsToArray([]float64{0, 0, 0, 0, 0, 0, 0, 0.5, 0.25, 0, 0, 0}, []int{n}), + } { + f, err := NewSparseCholesky(coo, ordering) + if err != nil { + t.Fatalf("NewSparseCholesky: %v", err) + } + diagBefore := append([]float64(nil), f.diag...) + valsBefore := append([]float64(nil), f.values...) + if err := f.Update(x); err != nil { + t.Fatalf("%s ordering %d: update: %v", name, ordering, err) + } + // The updated factor solves like the refactored matrix. + refactor, err := NewSparseCholesky(addRankOneCOO(t, coo, x), ordering) + if err != nil { + t.Fatalf("refactor: %v", err) + } + sModified, err := f.Solve(b) + if err != nil { + t.Fatalf("solve on the updated factor: %v", err) + } + sRefactor, err := refactor.Solve(b) + if err != nil { + t.Fatalf("solve on the refactored matrix: %v", err) + } + scale := 0.0 + for i := range n { + if v := math.Abs(sRefactor.FloatAt(i)); v > scale { + scale = v + } + } + for i := range n { + if math.Abs(sModified.FloatAt(i)-sRefactor.FloatAt(i)) > 1e-9*scale { + t.Fatalf("%s ordering %d: updated solve disagrees with the refactored one at %d", + name, ordering, i) + } + } + // Downdate restores the factor to round-off. + if err := f.Downdate(x); err != nil { + t.Fatalf("%s ordering %d: downdate: %v", name, ordering, err) + } + diagAfter := append([]float64(nil), f.diag...) + valsAfter := append([]float64(nil), f.values...) + worst := 0.0 + for i := range diagBefore { + worst = math.Max(worst, math.Abs(diagAfter[i]-diagBefore[i])) + } + for i := range valsBefore { + worst = math.Max(worst, math.Abs(valsAfter[i]-valsBefore[i])) + } + if worst > 1e-12 { + t.Fatalf("%s ordering %d: round trip lost %.3g on the stored factor", name, ordering, worst) + } + } + } +} + +// TestSparseCholeskyUpdateMatchesRefactorisation pins the numeric +// truth of the update: on the two-dimensional grid with a +// single-coordinate update the sweep stays on the pattern, and every +// stored entry must equal the fresh factorisation of the explicitly +// modified matrix, whose pattern the diagonal change cannot alter. +func TestSparseCholeskyUpdateMatchesRefactorisation(t *testing.T) { + coo := gridLaplacianCOO(t, 8, 8) + n := 64 + x := unitVector(t, n, 30, 0.5) + for _, ordering := range []SparseOrdering{SparseOrderingNatural, SparseOrderingReverseCuthillMcKee} { + f, err := NewSparseCholesky(coo, ordering) + if err != nil { + t.Fatalf("NewSparseCholesky: %v", err) + } + if err := f.Update(x); err != nil { + t.Fatalf("update: %v", err) + } + refactor, err := NewSparseCholesky(addRankOneCOO(t, coo, x), ordering) + if err != nil { + t.Fatalf("refactor: %v", err) + } + if f.NNZ() != refactor.NNZ() { + t.Fatalf("the update changed the fill: %d vs %d", f.NNZ(), refactor.NNZ()) + } + if d := factorDelta(f, refactor); d > 1e-9 { + t.Fatalf("the updated factor differs from the refactored one by %.3g", d) + } + } +} + +// TestSparseCholeskyUpdateDowndateExact pins bit-for-bit restoration +// on the engineered pattern-stable example: a banded matrix whose +// stored factor carries the integers 3 and 5, updated by the +// Pythagorean vector 4·e₀, where hypot(3, 4) = 5 and the rotations +// round exactly. Update turns every 3 into a 5 and every 5 into a 3; +// the downdate turns them back without losing a bit. +func TestSparseCholeskyUpdateDowndateExact(t *testing.T) { + const n = 6 + idx := make([]int64, 0, 3*n) + vals := make([]float64, 0, 3*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + if i == 0 { + add(0, 0, 9) + } else { + add(i, i, 34) + } + if i+1 < n { + add(i, i+1, 15) + add(i+1, i, 15) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + f, err := NewSparseCholesky(coo, SparseOrderingNatural) + if err != nil { + t.Fatalf("NewSparseCholesky: %v", err) + } + diagBefore := append([]float64(nil), f.diag...) + valsBefore := append([]float64(nil), f.values...) + for i := range diagBefore { + if diagBefore[i] != 3 { + t.Fatalf("the engineered factor's diagonal is %.4g, want 3", diagBefore[i]) + } + } + x := unitVector(t, n, 0, 4) + if err := f.Update(x); err != nil { + t.Fatalf("update: %v", err) + } + for i := range f.diag { + if f.diag[i] != 5 { + t.Fatalf("after the update diag[%d] = %.17g, want exactly 5", i, f.diag[i]) + } + } + for i := range f.values { + if f.values[i] != 3 { + t.Fatalf("after the update values[%d] = %.17g, want exactly 3", i, f.values[i]) + } + } + if err := f.Downdate(x); err != nil { + t.Fatalf("downdate: %v", err) + } + for i := range f.diag { + if f.diag[i] != diagBefore[i] { + t.Fatalf("the downdate lost diag[%d]: %.17g vs %.17g", i, f.diag[i], diagBefore[i]) + } + } + for i := range f.values { + if f.values[i] != valsBefore[i] { + t.Fatalf("the downdate lost values[%d]: %.17g vs %.17g", i, f.values[i], valsBefore[i]) + } + } +} + +func TestSparseCholeskyUpdateRefusals(t *testing.T) { + coo := gridLaplacianCOO(t, 12, 1) + // The update whose support spreads beyond a stored column needs + // fill the pattern does not hold: refused, factor untouched. + f, err := NewSparseCholesky(coo, SparseOrderingNatural) + if err != nil { + t.Fatalf("NewSparseCholesky: %v", err) + } + before := append([]float64(nil), f.values...) + spread := unitVector(t, 12, 0, 1) + spread.RawFloats()[5] = 1 + err = f.Update(spread) + if err == nil || !strings.Contains(err.Error(), "pattern does not hold") { + t.Fatalf("a pattern-violating update was accepted: %v", err) + } + for i := range f.values { + if f.values[i] != before[i] { + t.Fatalf("the refused update moved values[%d]", i) + } + } + // The downdate that loses positive dominance: 3·e₀ against a unit + // diagonal. + if err := f.Downdate(unitVector(t, 12, 0, 3)); err == nil || !strings.Contains(err.Error(), "positive definite") { + t.Fatalf("a downdate outside the cone was accepted: %v", err) + } + // A downdate whose square leaves the float64 range. + huge, err := core.NewSparseCOO(mustInts(t, []int64{0, 0}, 1, 2), floatsToArray([]float64{1e308}, []int{1}), []int{1, 1}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + hf, err := NewSparseCholesky(huge, SparseOrderingNatural) + if err != nil { + t.Fatalf("NewSparseCholesky: %v", err) + } + if err := hf.Update(unitVector(t, 1, 0, 1e154)); err != nil { + t.Fatalf("huge update: %v", err) + } + if err := hf.Downdate(unitVector(t, 1, 0, 1e154)); err == nil || !strings.Contains(err.Error(), "range") { + t.Fatalf("an out-of-range downdate was accepted: %v", err) + } + // Validation. + if err := f.Update(unitVector(t, 11, 0, 1)); err == nil { + t.Fatal("a short vector was accepted") + } + if err := f.Downdate(core.New(core.Float, 3, 4)); err == nil { + t.Fatal("a rank-2 vector was accepted") + } + if err := f.Update(core.New(core.Complex, 12)); err == nil { + t.Fatal("a complex vector was accepted") + } + // A non-finite entry. + nan := unitVector(t, 12, 2, 1) + nan.RawFloats()[2] = math.NaN() + if err := f.Update(nan); err == nil || !strings.Contains(err.Error(), "finite") { + t.Fatalf("a NaN entry was accepted: %v", err) + } +} diff --git a/linalg/sparsecomplex.go b/linalg/sparsecomplex.go new file mode 100644 index 0000000..da02cad --- /dev/null +++ b/linalg/sparsecomplex.go @@ -0,0 +1,745 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Complex sparse linear algebra. The real sparse surface +// stores float64 payloads; the complex side mirrors it with a complex +// CSR whose values are complex128, built from a core.SparseCOO holding +// complex values (NewSparseCOO accepts them; SparseFrom does not, so +// callers assemble the COO directly). The solvers are complex +// conjugate gradient for Hermitian positive-definite systems, complex +// BiCGSTAB for general ones, and a Hermitian Lanczos eigensolver; the +// electromagnetics Helmholtz problems that motivated this surface hit +// all three. + +// complexCSR is the complex128 compressed sparse row form. +type complexCSR struct { + rowStart []int + colIdx []int + vals []complex128 + n int +} + +// cEntry is one coordinate entry on the way into CSR form. +type cEntry struct { + row, col int + val complex128 +} + +// stableOrderComplexEntries orders entries by (row, col) and returns +// the slice holding the result. Two stable counting passes reach the +// order a stable comparison sort by (row, col) reaches, in linear time: +// the first places by column, the second by row, so entries that share +// both coordinates keep their input order and their values accumulate +// in it. +func stableOrderComplexEntries(entries []cEntry, rows, cols int) []cEntry { + buf := make([]cEntry, len(entries)) + count := make([]int, max(rows, cols)+1) + countingPlace(buf, entries, count, func(e cEntry) int { return e.col }) + countingPlace(entries, buf, count, func(e cEntry) int { return e.row }) + return entries +} + +// cooToComplexCSR converts a complex-valued SparseCOO to CSR form, +// merging duplicate coordinates and dropping explicit zeros. The +// values must be complex128; anything else is an error naming the +// dtype. +func cooToComplexCSR(s *core.SparseCOO, name string) (*complexCSR, error) { + if s.Values.Dtype() != core.Complex { + return nil, base.Errf("%s: needs complex128 values, got %s", name, s.Values.Dtype()) + } + if len(s.Shape) != 2 { + return nil, base.Errf("%s: needs a 2-D matrix, got rank %d", name, len(s.Shape)) + } + rows, cols := s.Shape[0], s.Shape[1] + if rows != cols { + return nil, base.Errf("%s: needs a square matrix, got %d×%d", name, rows, cols) + } + nnz := s.Indices.Shape()[0] + idx := s.Indices.RawInts() + entries := make([]cEntry, nnz) + for i := range nnz { + r := int(idx[i*2]) + c := int(idx[i*2+1]) + if r < 0 || r >= rows || c < 0 || c >= cols { + return nil, base.Errf("%s: index [%d,%d] out of range for %d×%d", name, r, c, rows, cols) + } + entries[i] = cEntry{r, c, s.Values.ComplexAt(i)} + } + // Sort by (row, col) so duplicates merge and rows are contiguous, + // keeping equal coordinates in COO order. + sorted := stableOrderComplexEntries(entries, rows, cols) + cm := &complexCSR{ + rowStart: make([]int, rows+1), + colIdx: make([]int, 0, nnz), + vals: make([]complex128, 0, nnz), + n: rows, + } + // The duplicates are adjacent after the sort, so each run merges into + // one accumulated value on the way into the compressed form. + for p := 0; p < len(sorted); { + e := sorted[p] + v := e.val + p++ + for p < len(sorted) && sorted[p].row == e.row && sorted[p].col == e.col { + v += sorted[p].val + p++ + } + if v == 0 { + continue + } + cm.colIdx = append(cm.colIdx, e.col) + cm.vals = append(cm.vals, v) + cm.rowStart[e.row+1]++ + } + for i := range cm.n { + cm.rowStart[i+1] += cm.rowStart[i] + } + return cm, nil +} + +// checkHermitian verifies every stored entry has a conjugate mirror, +// the property the complex CG solver and the Hermitian Lanczos rely +// on. The tolerance is purely relative to the largest magnitude present, +// deliberately without an absolute floor: a floor would approve a matrix +// of small scale whose asymmetry is a large fraction of it, and the +// Jacobi and Lanczos routes would then answer for a matrix that is +// neither A nor Aᴴ. Relative-only keeps the approval and the arithmetic +// in agreement. A non-finite entry is refused in its own right: NaN +// compares unequal to everything, so the mirror test alone would wave +// it through as Hermitian. The wording of the refusal is unchanged. +func (c *complexCSR) checkHermitian(name string) error { + scale := 0.0 + for _, v := range c.vals { + if a := cmplxAbs(v); a > scale { + scale = a + } + } + tol := 1e-12 * scale + for i := range c.n { + for p := c.rowStart[i]; p < c.rowStart[i+1]; p++ { + j := c.colIdx[p] + v := c.vals[p] + if math.IsNaN(real(v)) || math.IsNaN(imag(v)) || + math.IsInf(real(v), 0) || math.IsInf(imag(v), 0) { + return base.Errf("%s: entry [%d,%d] is not finite", name, i, j) + } + mirror, ok := c.at(j, i) + if !ok || cmplxAbs(mirror-complexConj(v)) > tol { + return base.Errf("%s: matrix is not Hermitian within 1e-12 tolerance", name) + } + } + } + return nil +} + +func complexConj(z complex128) complex128 { return complex(real(z), -imag(z)) } + +func cmplxAbs(z complex128) float64 { return math.Hypot(real(z), imag(z)) } + +// at reads the entry (row, col), reporting whether it is stored. The +// row's column indices are sorted, so the lookup is a binary search: +// the Hermitian check calls it once per stored entry, and a linear +// scan made that check quadratic in the row length. +func (c *complexCSR) at(row, col int) (complex128, bool) { + lo, hi := c.rowStart[row], c.rowStart[row+1] + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if c.colIdx[mid] < col { + lo = mid + 1 + } else { + hi = mid + } + } + if lo < c.rowStart[row+1] && c.colIdx[lo] == col { + return c.vals[lo], true + } + return 0, false +} + +// complexMatVecRange computes y[i] = A[i,:]·x for the rows in [s, e) of +// a matrix in compressed row form. It is the whole body of the product, +// so the serial path calls it directly and the parallel path calls it +// per chunk: taking the structure slices as arguments rather than as +// captured variables keeps the serial path free of the closure a row +// split needs, which the complex solvers call once per iteration. +func complexMatVecRange(s, e int, rowStart, colIdx []int, values, x, y []complex128) { + for i := s; i < e; i++ { + sum := 0i + for p := rowStart[i]; p < rowStart[i+1]; p++ { + sum += values[p] * x[colIdx[p]] + } + y[i] = sum + } +} + +// matVec computes y = A·x over complex vectors. Output rows are +// independent, so the row range splits across workers once a worker's +// share of the stored entries pays for the split, and runs on the +// calling goroutine below that. +func (c *complexCSR) matVec(x, y []complex128) { + // The three structure slices are taken once: every worker's rows walk + // them per stored entry, and the row loop holds no other state. + rowStart := c.rowStart + colIdx := c.colIdx + vals := c.vals + if sparseMatVecSplit(c.n, len(vals)) { + engine.ParallelMin(c.n, 1, func(s, e int) { + complexMatVecRange(s, e, rowStart, colIdx, vals, x, y) + }) + return + } + complexMatVecRange(0, c.n, rowStart, colIdx, vals, x, y) +} + +// complexVecF64 flattens a rank-1 complex array, rejecting anything +// else loudly. +func complexVecF64(b *core.Array, n int, name string) ([]complex128, error) { + if b.NDim() != 1 || b.Len() != n { + return nil, base.Errf("%s: b must be a rank-1 vector of length %d", name, n) + } + if b.Dtype() != core.Complex { + return nil, base.Errf("%s: b must be complex128, got %s", name, b.Dtype()) + } + out := make([]complex128, n) + if !b.Strided() { + copy(out, b.RawComplexes()[:n]) + return out, nil + } + for i := range n { + out[i] = b.ComplexAt(i) + } + return out, nil +} + +// complexFromArray builds the result array from a complex vector. +func complexFromArray(v []complex128) (*core.Array, error) { + out, err := zeros(core.Complex, []int{len(v)}) + if err != nil { + return nil, err + } + copy(out.RawComplexes(), v) + return out, nil +} + +// SpSolveComplexCG returns x solving A·x = b for a Hermitian +// positive-definite sparse A with complex entries, by conjugate +// gradient with a complex Jacobi (diagonal) preconditioner. The +// stopping rule mirrors SpSolve: ‖b − A·x‖₂ ≤ tol·‖b‖₂, tol ≤ 0 uses +// 1e-10, maxIter ≤ 0 uses n, and an unconverged solve is an error +// naming the achieved residual rather than a silent approximation. A +// zero or missing diagonal entry is refused, exactly as the real +// solver refuses it. +func SpSolveComplexCG(a *core.SparseCOO, b *core.Array, tol float64, maxIter int) (*core.Array, error) { + const name = "SpSolveComplexCG" + c, err := cooToComplexCSR(a, name) + if err != nil { + return nil, err + } + if err := c.checkHermitian(name); err != nil { + return nil, err + } + n := c.n + bv, err := complexVecF64(b, n, name) + if err != nil { + return nil, err + } + if tol <= 0 { + tol = spSolveTol + } + if maxIter <= 0 { + maxIter = n + } + diag := make([]complex128, n) + for i := range n { + d, ok := c.at(i, i) + if !ok || d == 0 { + return nil, base.Errf("%s: zero diagonal entry at row %d", name, i) + } + diag[i] = d + } + // The system is solved in scaled units, which leaves the solution x + // untouched: A·x = b and (f·A)·x = (f·b) have the same x. The factor + // is 1 for any ordinary magnitude, so the recurrence stays + // bit-identical there, and outside that range it is what keeps the + // dot products and the squared magnitudes in the recurrence from + // overflowing or underflowing to nothing. + ws := scaleSystem(c, bv, diag) + bNorm := normC(bv) + if bNorm == 0 { + return complexFromArray(make([]complex128, n)) + } + x := make([]complex128, n) + r := append([]complex128(nil), bv...) + p := make([]complex128, n) + z := make([]complex128, n) + Ap := make([]complex128, n) + for i := range n { + z[i] = r[i] / diag[i] + } + copy(p, z) + rz := dotC(r, z) + for iter := range maxIter { + c.matVec(p, Ap) + pAp := dotC(p, Ap) + if !cFinite(pAp) || pAp == 0 { + return nil, base.Errf("%s: breakdown at step %d (p·A·p vanished)", name, iter+1) + } + alpha := rz / pAp + for i := range n { + x[i] += alpha * p[i] + r[i] -= alpha * Ap[i] + } + if normC(r) <= tol*bNorm { + return complexFromArray(x) + } + for i := range n { + z[i] = r[i] / diag[i] + } + rzNext := dotC(r, z) + if rzNext == 0 { + return nil, base.Errf("%s: breakdown at step %d (r·z vanished)", name, iter+1) + } + beta := rzNext / rz + for i := range n { + p[i] = z[i] + beta*p[i] + } + rz = rzNext + } + // The residual and tolerance are reported in the caller's units, not + // in the scaled units the recurrence ran in (ws is 1 unless the + // system was extreme). + return nil, base.Errf("%s: no convergence in %d steps, residual %.3g (tolerance %.3g)", + name, maxIter, normC(r)/ws, tol*bNorm/ws) +} + +// dotC is the complex inner product aᴴ·b (conjugating the left side, +// the convention every complex Krylov method uses). +func dotC(a, b []complex128) complex128 { + s := 0i + for i := range a { + s += complexConj(a[i]) * b[i] + } + return s +} + +// normC is the Euclidean norm of a complex vector. The squares are +// summed directly while the sum stays finite and non-zero, which is what +// the stopping rules of the solvers were tuned on and is bit-for-bit the +// magnitude of every ordinary vector. A sum that overflows to +Inf or +// underflows to 0 (the whole vector below about 1e-162) carries no +// magnitude at all, so it falls back to a max-scaled accumulation that +// does: the magnitude has to exist for a relative convergence test to +// mean anything. +func normC(a []complex128) float64 { + s := 0.0 + for _, v := range a { + s += real(v)*real(v) + imag(v)*imag(v) + } + if s == 0 || math.IsInf(s, 1) { + return normCScaled(a) + } + return math.Sqrt(s) +} + +// normCScaled sums the squares relative to the largest magnitude, so no +// intermediate leaves the normal float64 range. +func normCScaled(a []complex128) float64 { + maxAbs := 0.0 + for _, v := range a { + if m := cmplxAbs(v); m > maxAbs { + maxAbs = m + } + } + if maxAbs == 0 { + return 0 + } + s := 0.0 + for _, v := range a { + re := real(v) / maxAbs + im := imag(v) / maxAbs + s += re*re + im*im + } + return maxAbs * math.Sqrt(s) +} + +// scaleSystem moves a system's matrix, right-hand side and preconditioner +// into the magnitude window the Krylov recurrences are safe in, by one +// shared exact power of two, and reports that factor. A·x = b and +// (f·A)·x = (f·b) have the same solution x, so the factor cancels: it +// only has to come off the residual and tolerance the error paths print +// back. Inside the window the dot products, the squared magnitudes and +// the stopping rule's tol*bNorm all mean what they claim; outside it a +// huge b makes bNorm +Inf (every affine check passes) and a tiny one +// makes it 0 (the zero solution is returned as exact). +func scaleSystem(c *complexCSR, bv, diag []complex128) float64 { + ws := windowScale(math.Max(maxMagComplex(c.vals), maxMagComplex(bv))) + if ws != 1 { + scaleComplexes(c.vals, ws) + scaleComplexes(bv, ws) + scaleComplexes(diag, ws) + } + return ws +} + +// SpSolveComplexBiCGSTAB returns x solving A·x = b for a general +// (nonsymmetric, non-Hermitian) sparse A with complex entries, by +// BiCGSTAB with a complex Jacobi preconditioner. The contract mirrors +// SpSolveBiCGSTAB, including the affirmative convergence checks a NaN +// residual can never pass and the breakdown errors that name the step. +func SpSolveComplexBiCGSTAB(a *core.SparseCOO, b *core.Array, tol float64, maxIter int) (*core.Array, error) { + const name = "SpSolveComplexBiCGSTAB" + c, err := cooToComplexCSR(a, name) + if err != nil { + return nil, err + } + n := c.n + bv, err := complexVecF64(b, n, name) + if err != nil { + return nil, err + } + if tol <= 0 { + tol = spSolveTol + } + if maxIter <= 0 { + maxIter = n + } + diag := make([]complex128, n) + for i := range n { + d, ok := c.at(i, i) + if !ok || d == 0 { + return nil, base.Errf("%s: zero diagonal entry at row %d", name, i) + } + diag[i] = d + } + // The system is solved in scaled units, which leaves the solution x + // untouched, exactly as in SpSolveComplexCG above. + ws := scaleSystem(c, bv, diag) + bNorm := normC(bv) + if bNorm == 0 { + return complexFromArray(make([]complex128, n)) + } + r := append([]complex128(nil), bv...) + rHat := append([]complex128(nil), r...) + x := make([]complex128, n) + v := make([]complex128, n) + p := make([]complex128, n) + s := make([]complex128, n) + t := make([]complex128, n) + precond := func(dst, src []complex128) { + for i := range n { + dst[i] = src[i] / diag[i] + } + } + rho := 1 + 0i + alpha := 1 + 0i + omega := 1 + 0i + z := make([]complex128, n) + y := make([]complex128, n) + mt := make([]complex128, n) + for iter := range maxIter { + rhoNext := dotC(rHat, r) + if !cFinite(rhoNext) || rhoNext == 0 { + return nil, base.Errf("%s: breakdown at step %d (rho vanished)", name, iter+1) + } + if iter > 0 { + beta := rhoNext / rho * (alpha / omega) + for i := range n { + p[i] = r[i] + beta*(p[i]-omega*v[i]) + } + } else { + copy(p, r) + } + precond(z, p) + c.matVec(z, v) + rhatV := dotC(rHat, v) + if rhatV == 0 { + return nil, base.Errf("%s: breakdown at step %d (direction orthogonal to shadow)", name, iter+1) + } + alpha = rhoNext / rhatV + for i := range n { + s[i] = r[i] - alpha*v[i] + } + if !vecFiniteC(s) { + return nil, base.Errf("%s: breakdown at step %d (non-finite residual)", name, iter+1) + } + // The half step may already be exact. The check is affirmative + // so a NaN norm (which fails every comparison) can never read + // as converged. + if normC(s) > tol*bNorm { + precond(y, s) + c.matVec(y, t) + precond(mt, t) + mtS := dotC(mt, s) + mtT := dotC(mt, t) + if mtT == 0 { + return nil, base.Errf("%s: breakdown at step %d (stabiliser vanished)", name, iter+1) + } + omega = mtS / mtT + if !cFinite(omega) { + return nil, base.Errf("%s: breakdown at step %d (stabiliser vanished)", name, iter+1) + } + for i := range n { + x[i] += alpha*z[i] + omega*y[i] + r[i] = s[i] - omega*t[i] + } + if !vecFiniteC(r) { + return nil, base.Errf("%s: breakdown at step %d (non-finite residual)", name, iter+1) + } + if normC(r) <= tol*bNorm { + return complexFromArray(x) + } + if omega == 0 { + return nil, base.Errf("%s: stagnation at step %d (omega zero)", name, iter+1) + } + } else { + for i := range n { + x[i] += alpha * z[i] + } + return complexFromArray(x) + } + rho = rhoNext + } + // The residual and tolerance are reported in the caller's units, not + // in the scaled units the recurrence ran in (ws is 1 unless the + // system was extreme). + return nil, base.Errf("%s: no convergence in %d steps, residual %.3g (tolerance %.3g)", + name, maxIter, normC(r)/ws, tol*bNorm/ws) +} + +func cFinite(z complex128) bool { + return !math.IsNaN(real(z)) && !math.IsInf(real(z), 0) && + !math.IsNaN(imag(z)) && !math.IsInf(imag(z), 0) +} + +func vecFiniteC(v []complex128) bool { + for _, z := range v { + if !cFinite(z) { + return false + } + } + return true +} + +// SpEigenComplex returns the k eigenvalues of largest magnitude of a +// Hermitian sparse matrix with complex entries, each with its unit +// eigenvector, by the Lanczos recurrence over complex vectors. The +// contract mirrors SpEigen: values ordered by descending magnitude, +// eigenvectors as columns of an (n, k) complex array, approximate +// Ritz pairs whose accuracy improves with the iteration budget. A +// non-Hermitian matrix is refused; the general complex eigenproblem +// has no short recurrence and belongs to dense methods. +func SpEigenComplex(s *core.SparseCOO, k int, gen *core.Generator) (values, vectors *core.Array, err error) { + const name = "SpEigenComplex" + if len(s.Shape) != 2 || s.Shape[0] != s.Shape[1] { + return nil, nil, base.Errf("%s: needs a square 2-D sparse matrix, got shape %v", name, s.Shape) + } + n := s.Shape[0] + if n == 0 { + return nil, nil, base.Errf("%s: zero-sized matrix, got shape %v", name, s.Shape) + } + if k < 1 || k > n { + return nil, nil, base.Errf("%s: k must be in [1, %d], got %d", name, n, k) + } + c, err := cooToComplexCSR(s, name) + if err != nil { + return nil, nil, err + } + if err := c.checkHermitian(name); err != nil { + return nil, nil, err + } + // The recurrence is projected in scaled units: the alphas and betas + // carry the matrix's magnitude, the basis vectors do not, and a + // power-of-two factor multiplies the first two exactly while leaving + // the third bit-identical, so only the Ritz values have to be moved + // back. Without it the projected tridiagonal's squared magnitudes + // overflow at the top of the range and the recurrence collapses to + // NaN, and at the bottom the deflation floor sees zeros. + ws := windowScale(maxMagComplex(c.vals)) + if ws != 1 { + scaleComplexes(c.vals, ws) + } + if gen == nil { + gen = core.NewGenerator(spEigenSeed) + } + alphas, betas, basis := c.lanczosComplex(k, gen) + r := len(alphas) + + tMat := make([]float64, r*r) + for i := range r { + tMat[i*r+i] = alphas[i] + if i+1 < r { + tMat[i*r+i+1] = betas[i] + tMat[(i+1)*r+i] = betas[i] + } + } + tVec := eye(r) + if err := symmetricQr(tMat, tVec, r); err != nil { + return nil, nil, base.Errf("SpEigenComplex: %w", err) + } + + idx := make([]int, r) + for i := range r { + idx[i] = i + } + // Rank by descending magnitude, the same comparator SpEigen uses. + for i := 1; i < r; i++ { + for j := i; j > 0; j-- { + a, b := idx[j-1], idx[j] + da, db := math.Abs(tMat[a*r+a]), math.Abs(tMat[b*r+b]) + if da > db || (da == db && a <= b) { + break + } + idx[j-1], idx[j] = idx[j], idx[j-1] + } + } + sel := idx[:k] + + outVals := make([]float64, k) + outVecs := make([]complex128, n*k) + tmp := make([]complex128, n) + for j, si := range sel { + outVals[j] = tMat[si*r+si] + // Ritz vector: lift the real tridiagonal eigenvector through + // the complex Lanczos basis, v = Q·y. + clear(tmp) + for p := range r { + y := tVec[p*r+si] + if y == 0 { + continue + } + row := basis[p*n : (p+1)*n] + for i := range n { + tmp[i] += complex(y, 0) * row[i] + } + } + norm := normC(tmp) + if norm > 0 { + for i := range n { + tmp[i] /= complex(norm, 0) + } + } + for i := range n { + outVecs[i*k+j] = tmp[i] + } + } + if ws != 1 { + unscaleFloats(outVals, ws) + } + vecs, err := complexFromArray2D(outVecs, n, k) + if err != nil { + return nil, nil, err + } + return floatsToArray(outVals, []int{k}), vecs, nil +} + +// complexFromArray2D builds an (n, k) complex array from row-major +// values. +func complexFromArray2D(v []complex128, n, k int) (*core.Array, error) { + out, err := zeros(core.Complex, []int{n, k}) + if err != nil { + return nil, err + } + copy(out.RawComplexes(), v) + return out, nil +} + +// lanczosComplex runs the Hermitian Lanczos recurrence over complex +// vectors: alphas stay real (qᴴAq of a Hermitian A), betas stay real +// (vector norms), so the projected problem reuses the real symmetric +// tridiagonal eigensolver unchanged. Full reorthogonalisation and the +// purely scale-relative deflation floor mirror the real lanczos. +func (c *complexCSR) lanczosComplex(k int, gen *core.Generator) (alphas, betas []float64, basis []complex128) { + n := c.n + steps := min(n, max(2*k, k+spEigenBlock)) + w := make([]complex128, n) + q := make([]complex128, n) + prev := make([]complex128, n) + // The budget bounds the recurrence: one basis row and at most one + // coefficient per step, so all three are sized once rather than grown. + alphas = make([]float64, 0, steps) + betas = make([]float64, 0, steps) + basis = make([]complex128, 0, steps*n) + scale := 0.0 + addScale := func(v float64) { + if v > scale { + scale = v + } + } + start := func() { + for i := range n { + q[i] = complex(gen.NormalUnit(), gen.NormalUnit()) + } + for p := range len(alphas) { + row := basis[p*n : (p+1)*n] + d := dotC(row, q) + for i := range n { + q[i] -= d * row[i] + } + } + if norm := normC(q); norm > 0 { + for i := range n { + q[i] /= complex(norm, 0) + } + } + } + start() + for range steps { + basis = append(basis, q...) + c.matVec(q, w) + alpha := real(dotC(q, w)) + alphas = append(alphas, alpha) + addScale(math.Abs(alpha)) + a := complex(alpha, 0) + for i := range n { + w[i] -= a * q[i] + } + if len(alphas) > 1 { + b := complex(betas[len(betas)-1], 0) + for i := range n { + w[i] -= b * prev[i] + } + } + // Full reorthogonalisation, twice: one pass removes the + // accumulated loss, the second what the first reintroduces. + for range 2 { + for p := range len(alphas) { + row := basis[p*n : (p+1)*n] + d := dotC(row, w) + for i := range n { + w[i] -= d * row[i] + } + } + } + beta := normC(w) + // A purely scale-relative exhaustion threshold; an absolute + // floor would collapse the recurrence for tiny-norm matrices. + if beta <= float64(n)*base.EpsF*scale { + if len(alphas) >= n { + break + } + betas = append(betas, 0) + start() + continue + } + betas = append(betas, beta) + addScale(beta) + copy(prev, q) + for i := range n { + q[i] = w[i] / complex(beta, 0) + } + } + if len(betas) >= len(alphas) { + betas = betas[:len(alphas)-1] + } + return alphas, betas, basis +} diff --git a/linalg/sparsecomplex_test.go b/linalg/sparsecomplex_test.go new file mode 100644 index 0000000..a233160 --- /dev/null +++ b/linalg/sparsecomplex_test.go @@ -0,0 +1,207 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// complexCOO builds a SparseCOO from (row, col, val) triples. +func complexCOO(t *testing.T, n int, entries []complex128) *core.SparseCOO { + t.Helper() + idx := make([]int64, 0, 2*len(entries)/3) + vals := make([]complex128, 0, len(entries)/3) + for i := 0; i+2 < len(entries); i += 3 { + idx = append(idx, int64(real(entries[i])), int64(real(entries[i+1]))) + vals = append(vals, entries[i+2]) + } + idxArr, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + valArr, err := core.FromComplexes(vals, len(vals)) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + coo, err := core.NewSparseCOO(idxArr, valArr, []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// TestSpSolveComplexCG pins the Hermitian solver: the diagonal +// complex matrix with Hermitian couplings must reproduce the known +// solution to tolerance. +func TestSpSolveComplexCG(t *testing.T) { + // A = diag(2, 3) with coupling i on both sides: Hermitian. + sp := complexCOO(t, 2, []complex128{ + 0, 0, 2, + 1, 1, 3, + 0, 1, 1i, + 1, 0, -1i, + }) + b, err := core.FromComplexes([]complex128{1 + 2i, -1}, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + x, err := SpSolveComplexCG(sp, b, 1e-12, 50) + if err != nil { + t.Fatalf("SpSolveComplexCG: %v", err) + } + // Verify A·x = b directly. + cm, err := cooToComplexCSR(sp, "test") + if err != nil { + t.Fatalf("cooToComplexCSR: %v", err) + } + xv := make([]complex128, 2) + for i := range 2 { + xv[i] = x.ComplexAt(i) + } + ax := make([]complex128, 2) + cm.matVec(xv, ax) + for i := range 2 { + if cmplxAbs(ax[i]-b.ComplexAt(i)) > 1e-10 { + t.Fatalf("residual[%d] = %v, want 0", i, ax[i]-b.ComplexAt(i)) + } + } +} + +// TestSpSolveComplexCGRejectsNonHermitian pins the symmetry gate. +func TestSpSolveComplexCGRejectsNonHermitian(t *testing.T) { + sp := complexCOO(t, 2, []complex128{ + 0, 0, 2, + 1, 1, 3, + 0, 1, 1i, + 1, 0, 1i, // same sign: not conjugate mirrors + }) + b, _ := core.FromComplexes([]complex128{1, 1}, 2) + if _, err := SpSolveComplexCG(sp, b, 1e-12, 50); err == nil { + t.Fatal("SpSolveComplexCG accepted a non-Hermitian matrix") + } +} + +// TestSpSolveComplexBiCGSTAB pins the general solver on a genuinely +// non-Hermitian complex system. +func TestSpSolveComplexBiCGSTAB(t *testing.T) { + sp := complexCOO(t, 3, []complex128{ + 0, 0, 2 + 1i, + 1, 1, 3, + 2, 2, 1 - 1i, + 0, 1, 1i, + 1, 0, 2, // A[1][0] ≠ conj(A[0][1]): non-Hermitian + 1, 2, 0.5, + }) + b, err := core.FromComplexes([]complex128{1, 2 - 1i, 3}, 3) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + x, err := SpSolveComplexBiCGSTAB(sp, b, 1e-12, 100) + if err != nil { + t.Fatalf("SpSolveComplexBiCGSTAB: %v", err) + } + cm, err := cooToComplexCSR(sp, "test") + if err != nil { + t.Fatalf("cooToComplexCSR: %v", err) + } + xv := make([]complex128, 3) + for i := range 3 { + xv[i] = x.ComplexAt(i) + } + ax := make([]complex128, 3) + cm.matVec(xv, ax) + for i := range 3 { + if cmplxAbs(ax[i]-b.ComplexAt(i)) > 1e-9 { + t.Fatalf("residual[%d] = %v, want 0", i, ax[i]-b.ComplexAt(i)) + } + } +} + +// TestSpSolveComplexBiCGSTABNaN pins the NaN guards: a poisoned entry +// must be a breakdown error, never a converged garbage answer. +func TestSpSolveComplexBiCGSTABNaN(t *testing.T) { + sp := complexCOO(t, 2, []complex128{ + 0, 0, 1, + 1, 1, complex(math.NaN(), 0), + }) + b, _ := core.FromComplexes([]complex128{1, 1}, 2) + if _, err := SpSolveComplexBiCGSTAB(sp, b, 1e-10, 50); err == nil { + t.Fatal("BiCGSTAB reported convergence on a NaN matrix") + } +} + +// TestSpEigenComplex pins the Hermitian Lanczos against a matrix with +// known spectrum: eigenvalues and the residual ‖A·v − λ·v‖. +func TestSpEigenComplex(t *testing.T) { + // [[2, i], [−i, 2]] has eigenvalues 3 and 1 with eigenvectors + // (1, −i)/√2 (for 3) and (1, i)/√2 (for 1). + sp := complexCOO(t, 2, []complex128{ + 0, 0, 2, + 1, 1, 2, + 0, 1, 1i, + 1, 0, -1i, + }) + vals, vecs, err := SpEigenComplex(sp, 2, core.NewGenerator(7)) + if err != nil { + t.Fatalf("SpEigenComplex: %v", err) + } + hi, lo := vals.FloatAt(0), vals.FloatAt(1) + if math.Abs(hi-3) > 1e-10 || math.Abs(lo-1) > 1e-10 { + t.Fatalf("eigenvalues (%g, %g), want (3, 1)", hi, lo) + } + cm, _ := cooToComplexCSR(sp, "test") + for j := range 2 { + v := make([]complex128, 2) + for i := range 2 { + v[i] = vecs.ComplexAt(i*2 + j) + } + if n := normC(v); math.Abs(n-1) > 1e-10 { + t.Fatalf("eigenvector %d norm %g, want 1", j, n) + } + av := make([]complex128, 2) + cm.matVec(v, av) + lam := complex(vals.FloatAt(j), 0) + for i := range 2 { + if cmplxAbs(av[i]-lam*v[i]) > 1e-9 { + t.Fatalf("eigenpair %d residual[%d] = %v", j, i, av[i]-lam*v[i]) + } + } + } +} + +// TestSpEigenComplexRejectsNonHermitian pins the symmetry gate. +func TestSpEigenComplexRejectsNonHermitian(t *testing.T) { + sp := complexCOO(t, 2, []complex128{ + 0, 0, 2, + 1, 1, 2, + 0, 1, 1i, + 1, 0, 1i, + }) + if _, _, err := SpEigenComplex(sp, 1, nil); err == nil { + t.Fatal("SpEigenComplex accepted a non-Hermitian matrix") + } +} + +// TestComplexCSRFromCOOMergesDuplicates pins the COO semantics: a +// coordinate repeated twice sums. +func TestComplexCSRFromCOOMergesDuplicates(t *testing.T) { + sp := complexCOO(t, 2, []complex128{ + 0, 0, 1 + 1i, + 0, 0, 2 - 1i, + 1, 1, 4, + }) + cm, err := cooToComplexCSR(sp, "test") + if err != nil { + t.Fatalf("cooToComplexCSR: %v", err) + } + v, ok := cm.at(0, 0) + if !ok || v != 3 { + t.Fatalf("merged (0,0) = %v ok=%v, want 3+0i", v, ok) + } + if got := len(cm.vals); got != 2 { + t.Fatalf("stored %d entries, want 2", got) + } +} diff --git a/linalg/sparsecsc.go b/linalg/sparsecsc.go new file mode 100644 index 0000000..e0fa58c --- /dev/null +++ b/linalg/sparsecsc.go @@ -0,0 +1,309 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "slices" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The compressed sparse column view of a COO matrix: the format direct +// factorisations run on, where a column at a time is eliminated and +// the fill of one column extends the entries below the diagonal. +type SparseCSC struct { + ColStart []int + RowIdx []int + Values []float64 + Rows int + Cols int +} + +// stableOrderEntriesByCol orders entries by (col, row) and returns the +// slice holding the result, with the same stability the row-major order +// keeps: one counting pass fixes the major key, the other the minor one, +// and entries that share both coordinates keep their input order and +// accumulate in it. +func stableOrderEntriesByCol(entries []cooEntry, rows, cols int) []cooEntry { + buf := make([]cooEntry, len(entries)) + count := make([]int, max(rows, cols)+1) + countingPlace(buf, entries, count, func(e cooEntry) int { return e.row }) + countingPlace(entries, buf, count, func(e cooEntry) int { return e.col }) + return entries +} + +// CSCFromCOO converts a core.SparseCOO to CSC form with the same +// canonicalisation CSRFromCOO applies: duplicate coordinates sum, the +// way COO semantics require, explicit zeros drop, and every column's +// row indices end up sorted and unique. core.Complex values are +// refused. +func CSCFromCOO(s *core.SparseCOO) (*SparseCSC, error) { + if s.Values.Dtype() == core.Complex { + return nil, base.Errf("CSCFromCOO: complex sparse matrices are not supported") + } + if len(s.Shape) != 2 { + return nil, base.Errf("CSCFromCOO: needs a 2-D matrix, got rank %d", len(s.Shape)) + } + rows, cols := s.Shape[0], s.Shape[1] + nnz := s.Indices.Shape()[0] + idx := s.Indices.RawInts() + vals := s.Values.RawFloats() + payload := s.Values.Dtype() == core.Float && !s.Values.Strided() + entries := make([]cooEntry, nnz) + for i := range nnz { + r := int(idx[i*2]) + c := int(idx[i*2+1]) + if r < 0 || r >= rows || c < 0 || c >= cols { + return nil, base.Errf("CSCFromCOO: index [%d,%d] out of range for %d×%d", r, c, rows, cols) + } + v := 0.0 + if payload { + v = vals[i] + } else { + v = s.Values.FloatAt(i) + } + entries[i] = cooEntry{r, c, v} + } + // Sort by (col, row) so duplicates merge and columns are + // contiguous, keeping equal coordinates in COO order. + sorted := stableOrderEntriesByCol(entries, rows, cols) + csc := &SparseCSC{Rows: rows, Cols: cols} + csc.ColStart = make([]int, cols+1) + csc.RowIdx = make([]int, 0, nnz) + csc.Values = make([]float64, 0, nnz) + // Merge duplicates before the zero drop, exactly as CSRFromCOO + // argues: dropping first would let a later duplicate of a dropped + // coordinate accumulate onto an unrelated slot. + for p := 0; p < len(sorted); { + e := sorted[p] + v := e.val + p++ + for p < len(sorted) && sorted[p].row == e.row && sorted[p].col == e.col { + v += sorted[p].val + p++ + } + if v == 0 { + continue + } + csc.RowIdx = append(csc.RowIdx, e.row) + csc.Values = append(csc.Values, v) + csc.ColStart[e.col+1]++ + } + for i := range cols { + csc.ColStart[i+1] += csc.ColStart[i] + } + return csc, nil +} + +// ToCSR returns the same matrix in CSR form, built by one counting +// pass over the stored entries. (The Cholesky lower-triangle row view +// is lowerRows, not this conversion.) +func (c *SparseCSC) ToCSR() (*SparseCSR, error) { + csr := &SparseCSR{Rows: c.Rows, Cols: c.Cols, RowStart: make([]int, c.Rows+1)} + csr.ColIdx = make([]int, len(c.Values)) + csr.Values = make([]float64, len(c.Values)) + for _, r := range c.RowIdx { + if r < 0 || r >= c.Rows { + return nil, base.Errf("ToCSR: row index %d out of range for %d rows", r, c.Rows) + } + csr.RowStart[r+1]++ + } + for i := range c.Rows { + csr.RowStart[i+1] += csr.RowStart[i] + } + // next holds the insertion point of each row while the columns + // are walked in order, which leaves every row's entries sorted by + // column. + next := make([]int, c.Rows) + copy(next, csr.RowStart[:c.Rows]) + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + r := c.RowIdx[p] + q := next[r] + csr.ColIdx[q] = j + csr.Values[q] = c.Values[p] + next[r] = q + 1 + } + } + return csr, nil +} + +// ToCSC returns the same matrix in CSC form, built by one counting +// pass over the stored entries. +func (c *SparseCSR) ToCSC() (*SparseCSC, error) { + csc := &SparseCSC{Rows: c.Rows, Cols: c.Cols, ColStart: make([]int, c.Cols+1)} + csc.RowIdx = make([]int, len(c.Values)) + csc.Values = make([]float64, len(c.Values)) + for _, j := range c.ColIdx { + if j < 0 || j >= c.Cols { + return nil, base.Errf("ToCSC: column index %d out of range for %d columns", j, c.Cols) + } + csc.ColStart[j+1]++ + } + for i := range c.Cols { + csc.ColStart[i+1] += csc.ColStart[i] + } + // next holds the insertion point of each column while the rows + // are walked in order, which leaves every column's entries sorted + // by row. + next := make([]int, c.Cols) + copy(next, csc.ColStart[:c.Cols]) + for i := range c.Rows { + for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ { + j := c.ColIdx[p] + q := next[j] + csc.RowIdx[q] = i + csc.Values[q] = c.Values[p] + next[j] = q + 1 + } + } + return csc, nil +} + +// cscMatVecColumns accumulates y += A[:,s:e]·x for the columns in +// [s, e). It is the whole body of the product on the calling +// goroutine: the scatter walks the columns in ascending order, which +// is the reduction order every row of the output sees. +func cscMatVecColumns(s, e int, colStart, rowIdx []int, values, x, y []float64) { + for j := s; j < e; j++ { + xj := x[j] + p0, p1 := colStart[j], colStart[j+1] + rv := rowIdx[p0:p1:p1] + vv := values[p0:p1:p1] + for i, r := range rv { + y[r] += vv[i] * xj + } + } +} + +// cscMatVecRowBlock computes y[s:e] = (A·x)[s:e] by walking every +// column and storing only the entries whose row lands in [s, e). The +// canonical form keeps a column's rows ascending, so one pair of end +// rows decides whether the column can reach the block at all, and a +// column that cannot costs two loads instead of its stored entries. The +// reads are paid once per block, but a row's contributions still arrive +// in ascending column order, exactly the order the serial walk gives +// them, so the result is bit-identical to the serial product. +func cscMatVecRowBlock(s, e int, colStart, rowIdx []int, values, x, y []float64) { + for j := range len(colStart) - 1 { + p0, p1 := colStart[j], colStart[j+1] + if p0 == p1 { + continue + } + if rowIdx[p1-1] < s || rowIdx[p0] >= e { + continue + } + // The column's rows ascend, so the block's share of the column + // is one contiguous span; the two binary searches find it and + // the walk below carries no per-entry window test. + seg := rowIdx[p0:p1:p1] + lo, _ := slices.BinarySearch(seg, s) + hi, _ := slices.BinarySearch(seg, e) + xj := x[j] + rv := rowIdx[p0+lo : p0+hi : p0+hi] + vv := values[p0+lo : p0+hi : p0+hi] + for i, r := range rv { + y[r] += vv[i] * xj + } + } +} + +// cscMatVecMaxBlocks caps how many row blocks the split uses. The +// split's cost per block is a full sweep over the column boundaries +// whatever the block's share of entries, so blocks past four pay more +// in sweeps than they win in parallel entries: measured on the banded +// benchmark at 131,072 rows, four blocks hold the best time and +// thirty-two regress past the serial product. +const cscMatVecMaxBlocks = 4 + +// cscMatVecMaxWidth is the stored entries per column past which the +// row split stays on the calling goroutine. A block's sweep reads every +// column's boundaries, so a wide column's structure traffic repeats per +// block while only its in-block entries turn into work: past about +// sixteen entries per column the sweeps outweigh the parallel entries, +// measured between a thirteen-wide band, which still wins, and a +// seventeen-wide one, which does not. +const cscMatVecMaxWidth = 16 + +// cscMatVecBlocks returns the row-block count the CSC matrix-vector +// product splits into, one for the calling goroutine. The per-worker +// share of stored entries must reach the work budget the CSR product +// applies, the columns must stay narrow enough for the boundary sweeps +// to pay, and more than one block must be available to split into. +func cscMatVecBlocks(rows, cols, nnz int) int { + w := min(engine.WorkersFor(rows), cscMatVecMaxBlocks) + if w > 1 && nnz/w >= sparseMatVecWorkBudget && nnz <= cscMatVecMaxWidth*cols { + return w + } + return 1 +} + +// cscRowsAscending reports whether every column's row indices ascend, +// the property the block walk's column skip and early exit rely on. +// The constructors produce it, but the fields are exported and can +// carry anything, so the split verifies it with one pass. +func cscRowsAscending(colStart, rowIdx []int) bool { + for j := range len(colStart) - 1 { + for p := colStart[j] + 1; p < colStart[j+1]; p++ { + if rowIdx[p] < rowIdx[p-1] { + return false + } + } + } + return true +} + +// MatVec computes y = A·x for a dense vector x of length Cols. Rows of +// the output are independent, so the row range splits across workers +// once the split pays, and the block walk visits a row's contributions +// in the same ascending column order the column walk gives them, which +// keeps the split bit-identical to the serial product. The split also +// verifies the canonical form with one pass, because its column skip +// and early exit read a column's first and last row as its bounds: a +// matrix whose columns do not ascend falls back to the serial product, +// which is order-blind and answers it exactly as the previous release +// did. Below the split the product runs on the calling goroutine, +// column by column in a fixed order, which keeps the reduction order +// deterministic. The constructors (CSCFromCOO, ToCSC, the factorisation +// patterns) all produce the canonical form. +func (c *SparseCSC) MatVec(x *core.Array) (*core.Array, error) { + if x.Dtype() == core.Complex { + return nil, base.Errf("MatVec: complex vectors are not supported") + } + if x.Len() != c.Cols { + return nil, base.Errf("MatVec: vector length %d does not match %d columns", x.Len(), c.Cols) + } + out := core.New(core.Float, []int{c.Rows}...) + xf := contiguousF64(x) + yv := out.RawFloats() + // The three structure slices are taken once: the scatter walks them + // per stored entry, and the loop below holds no other state. + colStart := c.ColStart + rowIdx := c.RowIdx + values := c.Values + w := cscMatVecBlocks(c.Rows, c.Cols, len(values)) + if w > 1 && !cscRowsAscending(colStart, rowIdx) { + w = 1 + } + if w > 1 { + chunk := (c.Rows + w - 1) / w + var wg sync.WaitGroup + for s := 0; s < c.Rows; s += chunk { + e := min(s+chunk, c.Rows) + wg.Go(func() { + cscMatVecRowBlock(s, e, colStart, rowIdx, values, xf, yv) + }) + } + wg.Wait() + return out, nil + } + cscMatVecColumns(0, c.Cols, colStart, rowIdx, values, xf, yv) + return out, nil +} + +// NNZ returns the count of stored non-zeros. +func (c *SparseCSC) NNZ() int { return len(c.Values) } diff --git a/linalg/sparsecsc_split_test.go b/linalg/sparsecsc_split_test.go new file mode 100644 index 0000000..b4e5dfa --- /dev/null +++ b/linalg/sparsecsc_split_test.go @@ -0,0 +1,123 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// sparseCSCMatVecReference accumulates the product from the stored +// entries directly. The test matrices carry small integer values, so +// every sum is exact whatever order the terms arrive in and the +// reference's answer pins MatVec bit for bit. +func sparseCSCMatVecReference(c *SparseCSC, x []float64) []float64 { + y := make([]float64, c.Rows) + for j, v := range x { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + y[c.RowIdx[p]] += c.Values[p] * v + } + } + return y +} + +// TestSparseCSCMatVecSplitMatchesReference drives the split path of +// the CSC product: a banded 131072-row matrix clears the split's work +// and width budgets, and its answer must match the exact reference. +// A band is the shape the split exists for, three entries per column. +func TestSparseCSCMatVecSplitMatchesReference(t *testing.T) { + // The split is a decision of the worker policy, whose default + // follows the machine's CPU count, so the policy is pinned here: + // the guard below states the matrix's shape, and the split is + // driven on a single-core machine exactly as on a many-core one. + prev := engine.SetNumWorkers(4) + defer engine.SetNumWorkers(prev) + + const n = 131072 + c := &SparseCSC{Rows: n, Cols: n, ColStart: make([]int, n+1)} + for j := range n { + c.ColStart[j] = len(c.RowIdx) + for _, r := range []int{j - 1, j, j + 1} { + if r < 0 || r >= n { + continue + } + c.RowIdx = append(c.RowIdx, r) + c.Values = append(c.Values, float64(r%7+1)) + } + } + c.ColStart[n] = len(c.RowIdx) + if cscMatVecBlocks(c.Rows, c.Cols, len(c.Values)) <= 1 { + t.Fatal("the banded matrix does not reach the split, the test would not drive it") + } + xv := make([]float64, n) + for i := range xv { + xv[i] = float64(i%11 + 1) + } + out, err := c.MatVec(mustFloats(t, xv, n)) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + want := sparseCSCMatVecReference(c, xv) + got := out.RawFloats() + for i := range want { + if got[i] != want[i] { + t.Fatalf("element %d = %v, want %v", i, got[i], want[i]) + } + } +} + +// TestSparseCSCMatVecNonCanonicalFallsBack pins the canonicality +// check: a matrix above the split's budgets whose rows descend inside +// a column must answer the exact product anyway, through the serial +// column walk the split falls back to. The exported fields make such +// a matrix constructible by any caller, and the serial walk answered +// it before the split existed. +func TestSparseCSCMatVecNonCanonicalFallsBack(t *testing.T) { + // The same pinned policy as above: the guard needs the split + // reachable, and only then does the canonicality check hand the + // matrix to the fallback. + prev := engine.SetNumWorkers(4) + defer engine.SetNumWorkers(prev) + + const n = 8192 + c := &SparseCSC{Rows: n, Cols: n, ColStart: make([]int, n+1)} + for j := range n { + c.ColStart[j] = len(c.RowIdx) + if j == 100 { + // Column 100 stores its two entries in descending row + // order, and the two rows land in different blocks: the + // block walk's column skip reads the first entry's row as + // the column's lower bound, drops the pair from the lower + // block, and the early exit drops it from the upper one. + c.RowIdx = append(c.RowIdx, 5000, 100) + c.Values = append(c.Values, 1, 3) + continue + } + c.RowIdx = append(c.RowIdx, j) + c.Values = append(c.Values, float64(j%5+1)) + } + c.ColStart[n] = len(c.RowIdx) + if cscMatVecBlocks(c.Rows, c.Cols, len(c.Values)) <= 1 { + t.Fatal("the matrix does not reach the split, the fallback would not be exercised") + } + if cscRowsAscending(c.ColStart, c.RowIdx) { + t.Fatal("the matrix reads as canonical, the fallback would not be exercised") + } + xv := make([]float64, n) + for i := range xv { + xv[i] = float64(i%9 + 1) + } + out, err := c.MatVec(mustFloats(t, xv, n)) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + want := sparseCSCMatVecReference(c, xv) + got := out.RawFloats() + for i := range want { + if got[i] != want[i] { + t.Fatalf("element %d = %v, want %v", i, got[i], want[i]) + } + } +} diff --git a/linalg/sparsecsr.go b/linalg/sparsecsr.go new file mode 100644 index 0000000..30302b7 --- /dev/null +++ b/linalg/sparsecsr.go @@ -0,0 +1,211 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The compressed sparse row view of a COO matrix: the format the +// iterative solvers and the Lanczos eigensolver run on. +type SparseCSR struct { + RowStart []int + ColIdx []int + Values []float64 + Rows int + Cols int +} + +// contiguousF64 returns an array's values as a float64 slice with no +// dtype switch and no strides test per element: the payload itself when +// the array is a contiguous float64 one, and a widening copy otherwise. +// Either way the values are the ones FloatAt returns. +func contiguousF64(a *core.Array) []float64 { + if a.Dtype() == core.Float && !a.Strided() { + return a.RawFloats()[:a.Len()] + } + out := make([]float64, a.Len()) + for i := range out { + out[i] = a.FloatAt(i) + } + return out +} + +// sparseMatVecWorkBudget is the number of stored multiply-adds one +// worker must carry before a row split pays for the goroutine that +// carries it. A stored entry costs a gathered load, a multiply and an +// add, which measures a little under a nanosecond on the banded sweep, +// so a worker's start-up is worth about this many of them. Measured on +// that sweep: 12,286 stored entries split across the worker count +// measure 10 µs against the calling goroutine's 10 µs, 24,574 measure +// 20 µs against 28 µs, and 393,214 measure 126 µs against 441 µs. +const sparseMatVecWorkBudget = 1 << 9 + +// sparseMatVecSplit reports whether a sparse matrix-vector product with +// the given shape spreads its rows across the workers. The fan-out is +// worth it only once a worker's share of the stored entries reaches the +// work budget, and the worker count bounds that share, so a matrix whose +// product misses the budget stays on the calling goroutine. +func sparseMatVecSplit(rows, nnz int) bool { + w := engine.WorkersFor(rows) + return w > 1 && nnz/w >= sparseMatVecWorkBudget +} + +// csrMatVecRange computes y[i] = A[i,:]·x for the rows in [s, e) of a +// matrix in compressed row form. It is the whole body of the product, so +// the serial path calls it directly and the parallel path calls it per +// chunk: taking the structure slices as arguments rather than as +// captured variables keeps the serial path free of the closure a row +// split needs, which the solvers call once per iteration. +func csrMatVecRange(s, e int, rowStart, colIdx []int, values, x, y []float64) { + for i := s; i < e; i++ { + rs, re := rowStart[i], rowStart[i+1] + vs := values[rs:re:re] + cs := colIdx[rs:re:re] + sum := 0.0 + for p := range vs { + sum += vs[p] * x[cs[p]] + } + y[i] = sum + } +} + +// cooEntry is one coordinate entry on the way from COO into a +// compressed form. +type cooEntry struct { + row, col int + val float64 +} + +// countingPlace stably places src into dst ordered by key: the key +// function returns the sort key of an entry, count is scratch holding at +// least the largest key plus two entries, and entries that share a key +// keep the order they had in src, because each one is appended to the +// end of its key's run as the input is walked. +func countingPlace[E any](dst, src []E, count []int, key func(E) int) { + clear(count) + for i := range src { + count[key(src[i])+1]++ + } + for i := range len(count) - 1 { + count[i+1] += count[i] + } + for i := range src { + k := key(src[i]) + dst[count[k]] = src[i] + count[k]++ + } +} + +// stableOrderEntries orders entries by (row, col) and returns the slice +// holding the result. Two stable counting passes reach the order a +// stable comparison sort by (row, col) reaches, in linear time: the +// first places by column, the second by row, so entries that share both +// coordinates keep their input order and their values accumulate in it. +func stableOrderEntries(entries []cooEntry, rows, cols int) []cooEntry { + buf := make([]cooEntry, len(entries)) + count := make([]int, max(rows, cols)+1) + countingPlace(buf, entries, count, func(e cooEntry) int { return e.col }) + countingPlace(entries, buf, count, func(e cooEntry) int { return e.row }) + return entries +} + +// CSRFromCOO converts a core.SparseCOO to CSR form, summing duplicate +// coordinates the way COO semantics require and dropping explicit +// zeros. core.Complex values are refused. +func CSRFromCOO(s *core.SparseCOO) (*SparseCSR, error) { + if s.Values.Dtype() == core.Complex { + return nil, base.Errf("CSRFromCOO: complex sparse matrices are not supported") + } + if len(s.Shape) != 2 { + return nil, base.Errf("CSRFromCOO: needs a 2-D matrix, got rank %d", len(s.Shape)) + } + rows, cols := s.Shape[0], s.Shape[1] + nnz := s.Indices.Shape()[0] + idx := s.Indices.RawInts() + vals := s.Values.RawFloats() + payload := s.Values.Dtype() == core.Float && !s.Values.Strided() + entries := make([]cooEntry, nnz) + for i := range nnz { + r := int(idx[i*2]) + c := int(idx[i*2+1]) + if r < 0 || r >= rows || c < 0 || c >= cols { + return nil, base.Errf("CSRFromCOO: index [%d,%d] out of range for %d×%d", r, c, rows, cols) + } + v := 0.0 + if payload { + v = vals[i] + } else { + v = s.Values.FloatAt(i) + } + entries[i] = cooEntry{r, c, v} + } + // Sort by (row, col) so duplicates merge and rows are contiguous, + // keeping equal coordinates in COO order. + sorted := stableOrderEntries(entries, rows, cols) + csr := &SparseCSR{Rows: rows, Cols: cols} + csr.RowStart = make([]int, rows+1) + csr.ColIdx = make([]int, 0, nnz) + csr.Values = make([]float64, 0, nnz) + // Merge duplicates first (already adjacent after the sort) so the CSR + // is canonical: one entry per (row, col), NNZ counts unique + // coordinates, and downstream consumers may assume sorted, unique + // column indices per row. Merging comes before the zero drop on + // purpose: dropping first would let a later duplicate of a dropped + // coordinate accumulate onto whatever slot happens to sit last in + // Values, which is a different coordinate, or no slot at all. + for p := 0; p < len(sorted); { + e := sorted[p] + v := e.val + p++ + for p < len(sorted) && sorted[p].row == e.row && sorted[p].col == e.col { + v += sorted[p].val + p++ + } + if v == 0 { + continue + } + csr.ColIdx = append(csr.ColIdx, e.col) + csr.Values = append(csr.Values, v) + csr.RowStart[e.row+1]++ + } + for i := range rows { + csr.RowStart[i+1] += csr.RowStart[i] + } + return csr, nil +} + +// MatVec computes y = A·x for a dense vector x of length Cols. +// Output rows are independent, so the row range splits across +// workers once a worker's share of the stored entries pays for the +// split, and runs on the calling goroutine below that. +func (c *SparseCSR) MatVec(x *core.Array) (*core.Array, error) { + if x.Dtype() == core.Complex { + return nil, base.Errf("MatVec: complex vectors are not supported") + } + if x.Len() != c.Cols { + return nil, base.Errf("MatVec: vector length %d does not match %d columns", x.Len(), c.Cols) + } + out := core.New(core.Float, []int{c.Rows}...) + xf := contiguousF64(x) + yv := out.RawFloats() + // The three structure slices are taken once: every worker's rows walk + // them per stored entry, and the row loop holds no other state. + rowStart := c.RowStart + colIdx := c.ColIdx + values := c.Values + if sparseMatVecSplit(c.Rows, len(values)) { + engine.ParallelMin(c.Rows, 1, func(s, e int) { + csrMatVecRange(s, e, rowStart, colIdx, values, xf, yv) + }) + return out, nil + } + csrMatVecRange(0, c.Rows, rowStart, colIdx, values, xf, yv) + return out, nil +} + +// NNZ returns the count of stored non-zeros. +func (c *SparseCSR) NNZ() int { return len(c.Values) } diff --git a/linalg/sparsefill.go b/linalg/sparsefill.go new file mode 100644 index 0000000..0be5470 --- /dev/null +++ b/linalg/sparsefill.go @@ -0,0 +1,493 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "cmp" + "slices" + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// Fill-reducing orderings for the sparse factorisations. The +// symmetrised adjacency builder and the orderings live here; the +// factorisation consumes nothing but a permutation, so a new ordering +// slots in without touching the numeric code. + +// symmetrisedAdjacency returns the adjacency lists of the pattern of +// A + Aᵀ with the diagonal dropped: for every stored entry (i, j), +// i ≠ j, i sits on j's list and j on i's. Every list is sorted by +// ascending degree and then index, which fixes the visiting order the +// breadth first search uses. The lists slice out of one flat arena +// sized by a counting pass, so the construction allocates twice +// whatever the vertex count; the degree snapshot counts the entries a +// list received, duplicates included, exactly as the incremental build +// measured them, and the degree-then-index sort is a total order, so +// the fill order inside a list cannot reach the result. +func symmetrisedAdjacency(c *SparseCSC) ([][]int, error) { + adj, counts, err := symmetrisedAdjacencyBuild(c) + if err != nil { + return nil, err + } + for i := range adj { + slices.SortStableFunc(adj[i], func(x, y int) int { + if d := cmp.Compare(counts[x], counts[y]); d != 0 { + return d + } + return cmp.Compare(x, y) + }) + } + return adj, nil +} + +// adjacencyIndexOrdered returns the same adjacency lists in plain +// ascending index order, the order the minimum-degree absorption +// merges: it skips the degree-then-index sort the breadth first +// search needs and no index-ordered consumer does. +func adjacencyIndexOrdered(c *SparseCSC) ([][]int, error) { + adj, _, err := symmetrisedAdjacencyBuild(c) + return adj, err +} + +// symmetrisedAdjacencyBuild builds the adjacency lists in ascending +// index order: one flat arena sized by a counting pass, every list +// sorted and compacted. The counts it also returns are the raw +// duplicate-inclusive neighbour counts the build measured, the values +// the degree-then-index ordering sorts by. +func symmetrisedAdjacencyBuild(c *SparseCSC) ([][]int, []int, error) { + counts := make([]int, c.Cols) + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + i := c.RowIdx[p] + if i == j { + continue + } + if i < 0 || i >= c.Cols { + return nil, nil, base.Errf("adjacency: row index %d out of range for %d columns", i, c.Cols) + } + counts[j]++ + counts[i]++ + } + } + total := 0 + for i := range c.Cols { + total += counts[i] + } + arena := make([]int, total) + adj := make([][]int, c.Cols) + next := make([]int, c.Cols) + off := 0 + for i := range c.Cols { + adj[i] = arena[off : off+counts[i]] + next[i] = off + off += counts[i] + } + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + i := c.RowIdx[p] + if i == j { + continue + } + arena[next[j]] = i + next[j]++ + arena[next[i]] = j + next[i]++ + } + } + for i := range c.Cols { + slices.Sort(adj[i]) + adj[i] = slices.Compact(adj[i]) + } + return adj, counts, nil +} + +// breadthFirstSearch walks the component of start and returns the +// visit order, labelling level for every vertex it reaches. It +// initialises level to the unvisited state itself, so the caller can +// pass a scratch slice. +func breadthFirstSearch(adj [][]int, level []int, start int) []int { + for i := range level { + level[i] = -1 + } + order := []int{start} + level[start] = 0 + for head := 0; head < len(order); head++ { + v := order[head] + for _, u := range adj[v] { + if level[u] < 0 { + level[u] = level[v] + 1 + order = append(order, u) + } + } + } + return order +} + +// deepestLevel returns the vertices of the largest level in visit +// order. +func deepestLevel(order []int, level []int) []int { + max := 0 + for _, v := range order { + if level[v] > max { + max = level[v] + } + } + last := make([]int, 0, len(order)) + for _, v := range order { + if level[v] == max { + last = append(last, v) + } + } + return last +} + +// pseudoPeripheralStart finds a start vertex whose eccentricity is +// close to the component's diameter: two sweeps of taking the +// shallowest-degree vertex of the deepest level, the standard +// pseudo-peripheral heuristic. Every choice is deterministic, so the +// ordering is a pure function of the pattern. +func pseudoPeripheralStart(adj [][]int, level []int, guess int) int { + start := guess + for range 2 { + order := breadthFirstSearch(adj, level, start) + last := deepestLevel(order, level) + best := last[0] + for _, v := range last[1:] { + if len(adj[v]) < len(adj[best]) { + best = v + } + } + start = best + } + return start +} + +// reverseCuthillMcKee returns the elimination order of the symmetric +// pattern: the vertices of every component in reverse breadth first +// order from a pseudo-peripheral start. The reverse of the search +// order is what shrinks the bandwidth, and with it the fill a +// triangular factorisation produces. Components are taken in +// ascending order of their smallest unreached vertex, so the result +// is a pure function of the pattern: position k holds the original +// index of the row eliminated k-th. +func reverseCuthillMcKee(c *SparseCSC) ([]int, error) { + adj, err := symmetrisedAdjacency(c) + if err != nil { + return nil, err + } + level := make([]int, c.Cols) + visited := make([]bool, c.Cols) + order := make([]int, 0, c.Cols) + for s := range c.Cols { + if visited[s] { + continue + } + start := pseudoPeripheralStart(adj, level, s) + component := breadthFirstSearch(adj, level, start) + for _, v := range component { + visited[v] = true + } + for _, c := range slices.Backward(component) { + order = append(order, c) + } + } + return order, nil +} + +// minimumDegree returns the elimination order of the symmetric +// pattern by minimum degree: at every step the uneliminated vertex +// with the fewest remaining neighbours is eliminated, its neighbours +// absorb its pattern (the element the elimination creates), and +// degree ties break to the smallest index, so the order is a pure +// function of the pattern. On irregular patterns it shrinks the fill +// well below what a bandwidth ordering reaches; on structured meshes +// the reverse Cuthill-McKee order is its match. +// +// The selection runs on a lazy frontier instead of a scan over the +// surviving vertices: a binary heap of (degree, vertex) pairs where a +// vertex enters when its degree is first known and its pair is rewritten +// in place whenever the degree moves, so the heap holds every surviving +// vertex exactly once and hands out the smallest (degree, index) among +// them, which is the first vertex of the smallest degree the scan would +// pick, at a logarithmic price per entry instead of a full sweep per +// elimination. +func minimumDegree(c *SparseCSC) ([]int, error) { + adj, err := adjacencyIndexOrdered(c) + if err != nil { + return nil, err + } + n := c.Cols + eliminated := make([]bool, n) + degree := make([]int, n) + frontier := °reeFrontier{entries: make([]degreeEntry, 0, n)} + arena := &intArena{} + for i := range n { + degree[i] = len(adj[i]) + frontier.push(degreeEntry{degree[i], i}) + } + order := make([]int, 0, n) + for range n { + if frontier.len() == 0 { + return nil, base.Errf("minimumDegree: no vertex left to eliminate") + } + p := frontier.pop().vertex + order = append(order, p) + eliminated[p] = true + absorbElement(adj, degree, eliminated, p, frontier, arena) + } + return order, nil +} + +// degreeEntry pairs a vertex with the degree it carried when it +// entered the frontier. +type degreeEntry struct { + degree int + vertex int +} + +// before orders the pairs by degree, then index: the heap's top is +// the vertex the selection would take. +func (e degreeEntry) before(other degreeEntry) bool { + if e.degree != other.degree { + return e.degree < other.degree + } + return e.vertex < other.vertex +} + +// degreeFrontier is a binary min-heap of degreeEntry pairs holding +// every vertex at most once: push rewrites the vertex's pair in place +// when one is already queued, so the heap never carries a stale pair +// and its top is always the vertex the selection takes. The absence of +// stale pairs keeps the heap at the vertex count instead of one entry +// per degree movement, which is what a sift-down per pop would +// otherwise charge. +type degreeFrontier struct { + entries []degreeEntry + pos []int +} + +// len reports how many pairs the frontier holds. +func (f *degreeFrontier) len() int { return len(f.entries) } + +// push records the vertex's current degree: the pair of a vertex +// already queued moves to its new place, a vertex without a pair +// enters at the end. pos grows to cover the vertex and fills its new +// slots with -1, so an absent vertex reads -1 and a queued one reads +// its heap index. +func (f *degreeFrontier) push(e degreeEntry) { + for len(f.pos) <= e.vertex { + f.pos = append(f.pos, -1) + } + if at := f.pos[e.vertex]; at != -1 { + f.entries[at] = e + f.fix(at) + return + } + f.pos[e.vertex] = len(f.entries) + f.entries = append(f.entries, e) + f.up(len(f.entries) - 1) +} + +// pop removes and returns the smallest pair. It is the caller's job to +// check len first; the returned pair is never stale. +func (f *degreeFrontier) pop() degreeEntry { + entries := f.entries + top := entries[0] + f.pos[top.vertex] = -1 + last := len(entries) - 1 + entries[0] = entries[last] + f.entries = entries[:last] + if last > 0 { + f.pos[entries[0].vertex] = 0 + f.down(0) + } + return top +} + +// fix restores the heap order around i after the pair at i changed. +func (f *degreeFrontier) fix(i int) { + if !f.up(i) { + f.down(i) + } +} + +// up lifts the pair at i until its parent stops outranking it, and +// reports whether it moved at all. +func (f *degreeFrontier) up(i int) bool { + entries := f.entries + moved := false + for i > 0 { + parent := (i - 1) / 2 + if !entries[i].before(entries[parent]) { + break + } + entries[i], entries[parent] = entries[parent], entries[i] + f.pos[entries[i].vertex] = i + i = parent + moved = true + } + f.pos[entries[i].vertex] = i + return moved +} + +// down drops the pair at i until both children stop outranking it. +func (f *degreeFrontier) down(i int) { + entries := f.entries + for { + child := 2*i + 1 + if child >= len(entries) { + break + } + if right := child + 1; right < len(entries) && entries[right].before(entries[child]) { + child = right + } + if !entries[child].before(entries[i]) { + break + } + entries[i], entries[child] = entries[child], entries[i] + f.pos[entries[i].vertex] = i + i = child + } + f.pos[entries[i].vertex] = i +} + +// intArena hands out int slices from one bump-allocated backing block, +// so the minimum-degree absorption merges cost an allocation per block +// instead of one per merge. A returned slice owns its range: the merge +// appends into a zero-length view whose capacity is its exact upper +// bound, and the arena never hands the same range out twice. The +// capacity a merge leaves over returns to the bump space through +// release, which must reach the arena before its next alloc, while the +// released slice is still the most recent one. Backing blocks are kept +// alive by the adjacency lists that point into them, not by the arena. +type intArena struct { + buf []int + // bump is the offset of the next free entry in buf; the block is + // replaced once the space behind bump no longer fits a request. + bump int +} + +// arenaBlockInts is the smallest backing block: big enough that the +// small early merges share one allocation, small enough that a block +// full of dead lists retires without pinning much memory. +const arenaBlockInts = 16384 + +// alloc returns a zero-length slice with capacity n. The arena works +// from an empty buffer as well as a partially consumed one: the bump +// offset starts at zero and the block check reads it against the block +// length, so a zero-value arena allocates rather than panics. +func (a *intArena) alloc(n int) []int { + if len(a.buf)-a.bump < n { + a.buf = make([]int, max(arenaBlockInts, n)) + a.bump = 0 + } + out := a.buf[a.bump : a.bump : a.bump+n] + a.bump += n + return out +} + +// release folds the capacity a merge left over back into the bump +// space, so the next alloc reuses it instead of moving past it. The +// slice must be the arena's most recent allocation and must not be +// released twice. +func (a *intArena) release(out []int) { + a.bump -= cap(out) - len(out) +} + +// absorbElement merges the element the eliminated vertex p leaves +// behind into every surviving neighbour's adjacency list, the fill +// the elimination creates: adj(j) becomes adj(j) ∪ adj(p) minus j and +// minus every eliminated vertex. The set difference matters: p sits on +// j's list and j on p's, so the plain union hands j its own index +// back, and a self entry both inflates the degree the selection reads +// and spreads to other lists through later absorptions. A neighbour's +// degree is the count of surviving entries in the merged list, the +// number a scan would recount at the same step, kept here where the +// list changes; a neighbour whose degree moved re-enters the frontier. +func absorbElement(adj [][]int, degree []int, eliminated []bool, p int, frontier *degreeFrontier, arena *intArena) { + // The element list is compacted once, in place, before any merge: + // p is eliminated and its list is dead, no vertex becomes + // eliminated while the absorb runs, and every neighbour's merge + // walks the same list, so one filter pass here replaces one per + // merged element. The compaction is stable, so the merge order is + // the list's own. + b := adj[p] + w := 0 + for _, u := range b { + if !eliminated[u] { + b[w] = u + w++ + } + } + b = b[:w] + for _, j := range b { + before := degree[j] + adj[j] = mergedElement(adj[j], b, j, eliminated, degree, arena) + if degree[j] != before { + frontier.push(degreeEntry{degree[j], j}) + } + } +} + +// mergedElement merges two sorted unique slices into one sorted unique +// slice without the self entry and without the eliminated entries, and +// stores the surviving count into degree[self]. The b side arrives +// already free of eliminated entries, so only the entries that come +// from a (and the merged equal pairs) pay the filter; b still carries +// self exactly once, which its own check drops. The degree the +// selection reads is the count of surviving entries either way, so the +// elimination order the frontier produces cannot move. The result lands +// in an arena slice capped at len(a)+len(b), which the merge cannot +// exceed: every entry of both lists is emitted at most once and the +// self entry never is. +func mergedElement(a, b []int, self int, eliminated []bool, degree []int, arena *intArena) []int { + n := len(a) + len(b) + out := arena.alloc(n) + // The arena reserved exactly n entries, so the merge writes through + // the stretched view and cuts the result back at the end. + out = out[:n:n] + w := 0 + i, j := 0, 0 + for i < len(a) && j < len(b) { + u := a[i] + switch v := b[j]; { + case u < v: + i++ + if u != self && !eliminated[u] { + out[w] = u + w++ + } + case v < u: + j++ + if v != self { + out[w] = v + w++ + } + default: + i++ + j++ + if u != self && !eliminated[u] { + out[w] = u + w++ + } + } + } + for ; i < len(a); i++ { + if u := a[i]; u != self && !eliminated[u] { + out[w] = u + w++ + } + } + for ; j < len(b); j++ { + if v := b[j]; v != self { + out[w] = v + w++ + } + } + out = out[:w] + degree[self] = w + // The merge filled w of the capacity it reserved: the tail goes + // back to the arena before the next merge carves past it. + arena.release(out) + return out +} diff --git a/linalg/sparsefill_bench_test.go b/linalg/sparsefill_bench_test.go new file mode 100644 index 0000000..f69c311 --- /dev/null +++ b/linalg/sparsefill_bench_test.go @@ -0,0 +1,161 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the minimum degree ordering: the frontier selection +// against the reference scan beside it, on the mesh patterns the +// direct solvers are measured on. Both variants run in one binary, and +// the pair runs in both orders across the two parent benchmarks, so a +// drift of the machine between the two halves of a round cannot dress +// itself up as a difference between the variants. The small grid is +// there so a regression the large grids would drown stays visible. + +// gridLaplacian3D builds the 7-point Laplacian on a w×h×d grid in +// row-major order: symmetric positive definite, and the 3-D mesh where +// an ordering's fill decisions cost the most. +func gridLaplacian3D(b *testing.B, w, h, d int) *core.SparseCOO { + b.Helper() + n := w * h * d + idx := make([]int64, 0, 7*n) + vals := make([]float64, 0, 7*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + at := func(x, y, z int) int { return (z*h+y)*w + x } + for z := range d { + for y := range h { + for x := range w { + add(at(x, y, z), at(x, y, z), 6) + if x+1 < w { + add(at(x, y, z), at(x+1, y, z), -1) + add(at(x+1, y, z), at(x, y, z), -1) + } + if y+1 < h { + add(at(x, y, z), at(x, y+1, z), -1) + add(at(x, y+1, z), at(x, y, z), -1) + } + if z+1 < d { + add(at(x, y, z), at(x, y, z+1), -1) + add(at(x, y, z+1), at(x, y, z), -1) + } + } + } + } + return sparseCOOFrom(b, n, idx, vals) +} + +// orderingGrid is one mesh the ordering benchmarks run on. +type orderingGrid struct { + name string + w, h, d int +} + +// orderingGrids are the meshes: the small grid every regression would +// show on, the two 2-D meshes, and the 3-D mesh where the quadratic +// scan costs most. +var orderingGrids = []orderingGrid{ + {"2d-361", 19, 19, 1}, + {"2d-4096", 64, 64, 1}, + {"2d-16384", 128, 128, 1}, + {"3d-32768", 32, 32, 32}, +} + +// orderingInput builds one grid's matrix, with the CSC form the +// orderings consume beside the COO the constructor takes, outside +// every timed region. +func orderingInput(b *testing.B, g orderingGrid) (*core.SparseCOO, *SparseCSC) { + b.Helper() + var coo *core.SparseCOO + if g.d == 1 { + coo = gridLaplacian(b, g.w, g.h) + } else { + coo = gridLaplacian3D(b, g.w, g.h, g.d) + } + c, err := CSCFromCOO(coo) + if err != nil { + b.Fatal(err) + } + return coo, c +} + +// orderingVariants are the two selections, named for the sub-benchmark +// that times them: the reference scan and the production frontier. +func orderingVariants() []struct { + name string + run func(*SparseCSC) error +} { + return []struct { + name string + run func(*SparseCSC) error + }{ + {"scan", func(c *SparseCSC) error { + _, err := minimumDegreeScan(c) + return err + }}, + {"frontier", func(c *SparseCSC) error { + _, err := minimumDegree(c) + return err + }}, + } +} + +// orderingBattery runs the scan and the frontier against the same +// pattern, in the order the caller picks. +func orderingBattery(b *testing.B, reverse bool) { + for _, g := range orderingGrids { + _, c := orderingInput(b, g) + variants := orderingVariants() + if reverse { + variants[0], variants[1] = variants[1], variants[0] + } + for _, v := range variants { + variant := v + b.Run(g.name+"/"+variant.name, func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if err := variant.run(c); err != nil { + b.Fatal(err) + } + } + }) + } + } +} + +// BenchmarkMinimumDegreeOrdering measures the ordering alone, scan +// first. +func BenchmarkMinimumDegreeOrdering(b *testing.B) { + orderingBattery(b, false) +} + +// BenchmarkMinimumDegreeOrderingRev measures the ordering alone, +// frontier first: the mirror of the other parent, so the pair's two +// halves alternate which variant pays for the position. +func BenchmarkMinimumDegreeOrderingRev(b *testing.B) { + orderingBattery(b, true) +} + +// BenchmarkSparseCholeskyMinimumDegreeFactor measures NewSparseCholesky +// end to end with the minimum degree ordering: the ordering sits inside +// the construction, so its cost is part of the number. +func BenchmarkSparseCholeskyMinimumDegreeFactor(b *testing.B) { + for _, g := range orderingGrids { + coo, _ := orderingInput(b, g) + b.Run(g.name, func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := NewSparseCholesky(coo, SparseOrderingMinimumDegree); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/linalg/sparsefill_test.go b/linalg/sparsefill_test.go new file mode 100644 index 0000000..189667b --- /dev/null +++ b/linalg/sparsefill_test.go @@ -0,0 +1,337 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math/rand/v2" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The reference side of the minimum degree tests: the selection scan +// the frontier replaced, kept verbatim so the equivalence test can +// hold the frontier to the exact sequence of choices the scan makes, +// tie breaks included, and the ordering benchmark can put a number on +// the difference. + +// minimumDegreeScan returns the elimination order the original scan +// selects: every step walks the surviving vertices, counts each one's +// remaining neighbours and keeps the first vertex of the smallest +// count, so ties break to the smallest index. +func minimumDegreeScan(c *SparseCSC) ([]int, error) { + adj, err := symmetrisedAdjacency(c) + if err != nil { + return nil, err + } + // The absorption below merges the lists as plain ascending index + // sets, so the breadth first search's (degree, index) order has to + // go: re-sort by index, exactly as the production order does. + for i := range adj { + slices.Sort(adj[i]) + } + n := c.Cols + eliminated := make([]bool, n) + order := make([]int, 0, n) + for range n { + p := -1 + best := 0 + for i := range n { + if eliminated[i] { + continue + } + d := 0 + for _, u := range adj[i] { + if !eliminated[u] { + d++ + } + } + if p == -1 || d < best { + p = i + best = d + } + } + if p == -1 { + return nil, base.Errf("minimumDegree: no vertex left to eliminate") + } + order = append(order, p) + eliminated[p] = true + absorbElementScan(adj, eliminated, p) + } + return order, nil +} + +// absorbElementScan merges the element the eliminated vertex p leaves +// behind into every surviving neighbour's adjacency list, the scan's +// form of the absorption: adj(j) becomes adj(j) ∪ adj(p) \ {j}. +func absorbElementScan(adj [][]int, eliminated []bool, p int) { + for _, j := range adj[p] { + if eliminated[j] { + continue + } + adj[j] = sortedSetUnionScan(adj[j], adj[p]) + if at, found := slices.BinarySearch(adj[j], j); found { + adj[j] = slices.Delete(adj[j], at, at+1) + } + } +} + +// sortedSetUnionScan merges two sorted unique slices into one sorted +// unique slice. +func sortedSetUnionScan(a, b []int) []int { + out := make([]int, 0, len(a)+len(b)) + i, j := 0, 0 + for i < len(a) && j < len(b) { + switch { + case a[i] < b[j]: + out = append(out, a[i]) + i++ + case b[j] < a[i]: + out = append(out, b[j]) + j++ + default: + out = append(out, a[i]) + i++ + j++ + } + } + out = append(out, a[i:]...) + out = append(out, b[j:]...) + return out +} + +// patternCOO assembles a symmetric matrix from off-diagonal edges plus +// a diagonal heavy enough to keep the matrix positive definite, though +// the ordering reads the pattern alone. +func patternCOO(t *testing.T, n int, edges [][2]int) *core.SparseCOO { + t.Helper() + idx := make([]int64, 0, 2*len(edges)+2*n) + vals := make([]float64, 0, 2*len(edges)+2*n) + add := func(r, c int) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, 1) + } + for _, e := range edges { + add(e[0], e[1]) + add(e[1], e[0]) + } + for i := range n { + idx = append(idx, int64(i), int64(i)) + vals = append(vals, float64(n)+2) + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// gridEdges returns the edges of the w×h grid graph. +func gridEdges(w, h int) [][2]int { + at := func(x, y int) int { return y*w + x } + edges := make([][2]int, 0, 2*w*h) + for y := range h { + for x := range w { + if x+1 < w { + edges = append(edges, [2]int{at(x, y), at(x+1, y)}) + } + if y+1 < h { + edges = append(edges, [2]int{at(x, y), at(x, y+1)}) + } + } + } + return edges +} + +// cubeEdges returns the edges of the w×h×d grid graph. +func cubeEdges(w, h, d int) [][2]int { + at := func(x, y, z int) int { return (z*h+y)*w + x } + edges := make([][2]int, 0, 3*w*h*d) + for z := range d { + for y := range h { + for x := range w { + if x+1 < w { + edges = append(edges, [2]int{at(x, y, z), at(x+1, y, z)}) + } + if y+1 < h { + edges = append(edges, [2]int{at(x, y, z), at(x, y+1, z)}) + } + if z+1 < d { + edges = append(edges, [2]int{at(x, y, z), at(x, y, z+1)}) + } + } + } + } + return edges +} + +// starEdges returns the edges of the star graph: a centre joined to +// every other vertex, the shape whose eliminations hand the centre its +// degree one neighbour at a time. +func starEdges(n int) [][2]int { + edges := make([][2]int, 0, n-1) + for i := 1; i < n; i++ { + edges = append(edges, [2]int{0, i}) + } + return edges +} + +// completeEdges returns the edges of the complete graph on n vertices, +// the shape whose first elimination fills everything. +func completeEdges(n int) [][2]int { + edges := make([][2]int, 0, n*(n-1)/2) + for i := range n { + for j := i + 1; j < n; j++ { + edges = append(edges, [2]int{i, j}) + } + } + return edges +} + +// relabelledEdges renames every vertex through a random permutation: +// the same pattern under an index order chosen to scatter the degree +// ties the tie break has to survive. +func relabelledEdges(rng *rand.Rand, edges [][2]int) [][2]int { + label := make(map[int]int) + for _, e := range edges { + label[e[0]] = 0 + label[e[1]] = 0 + } + names := make([]int, 0, len(label)) + for v := range label { + names = append(names, v) + } + slices.Sort(names) + order := rng.Perm(len(names)) + for i, v := range names { + label[v] = order[i] + } + out := make([][2]int, len(edges)) + for i, e := range edges { + out[i] = [2]int{label[e[0]], label[e[1]]} + } + return out +} + +// randomEdges draws m distinct off-diagonal edges of an n-vertex +// graph, the irregular patterns the ordering exists for. A random +// clique rides along every few calls to force the fill the absorptions +// have to keep up with. +func randomEdges(rng *rand.Rand, n, m int) [][2]int { + if n < 2 { + return nil + } + largest := n * (n - 1) / 2 + m = min(m, largest) + seen := make(map[[2]int]bool) + edges := make([][2]int, 0, m) + for len(edges) < m { + i := rng.IntN(n) + j := rng.IntN(n) + if i == j { + continue + } + e := [2]int{min(i, j), max(i, j)} + if seen[e] { + continue + } + seen[e] = true + edges = append(edges, e) + } + if rng.IntN(4) == 0 { + clique := min(n, 2+rng.IntN(n/4+1)) + start := rng.IntN(n - clique + 1) + for i := start; i < start+clique; i++ { + for j := i + 1; j < start+clique; j++ { + e := [2]int{i, j} + if !seen[e] { + seen[e] = true + edges = append(edges, e) + } + } + } + } + return edges +} + +// checkOrderMatchesScan factors the pattern's adjacency once for each +// side and requires the frontier order to equal the scan order entry +// for entry, and to be a permutation at all. +func checkOrderMatchesScan(t *testing.T, name string, coo *core.SparseCOO) { + t.Helper() + c, err := CSCFromCOO(coo) + if err != nil { + t.Fatalf("%s: CSCFromCOO: %v", name, err) + } + want, err := minimumDegreeScan(c) + if err != nil { + t.Fatalf("%s: scan: %v", name, err) + } + got, err := minimumDegree(c) + if err != nil { + t.Fatalf("%s: frontier: %v", name, err) + } + for i := range min(len(got), len(want)) { + if got[i] != want[i] { + t.Fatalf("%s: step %d eliminates %d, the scan eliminates %d", name, i, got[i], want[i]) + } + } + if len(got) != len(want) { + t.Fatalf("%s: order lengths differ: %d vs %d", name, len(got), len(want)) + } + sorted := slices.Clone(got) + slices.Sort(sorted) + for i := range sorted { + if sorted[i] != i { + t.Fatalf("%s: order entry %d holds %d; not a permutation", name, i, sorted[i]) + } + } +} + +// TestMinimumDegreeMatchesScan holds the frontier ordering to the +// reference scan: on every pattern below, both must return the same +// elimination order, tie breaks included. The permutation decides the +// factorisation's fill, so a single differing choice is a failure. +func TestMinimumDegreeMatchesScan(t *testing.T) { + rng := rand.New(rand.NewPCG(2026, 9)) + structured := []struct { + name string + n int + edges [][2]int + }{ + {"empty", 12, nil}, + {"path-50", 50, gridEdges(50, 1)}, + {"star-50", 50, starEdges(50)}, + {"complete-30", 30, completeEdges(30)}, + {"grid-12x12", 144, gridEdges(12, 12)}, + {"grid-15x15-shuffled", 225, relabelledEdges(rng, gridEdges(15, 15))}, + {"grid-64x64", 4096, gridEdges(64, 64)}, + {"cube-6x6x6", 216, cubeEdges(6, 6, 6)}, + {"two-grids", 64 + 36, append(gridEdges(8, 8), func() [][2]int { + shifted := gridEdges(6, 6) + for i := range shifted { + shifted[i][0] += 64 + shifted[i][1] += 64 + } + return shifted + }()...)}, + } + for _, tc := range structured { + checkOrderMatchesScan(t, tc.name, patternCOO(t, tc.n, tc.edges)) + } + for range 400 { + n := 1 + rng.IntN(90) + edges := randomEdges(rng, n, rng.IntN(3*n+1)) + if rng.IntN(2) == 0 { + edges = relabelledEdges(rng, edges) + } + checkOrderMatchesScan(t, "random", patternCOO(t, n, edges)) + } +} diff --git a/linalg/sparsegeneral.go b/linalg/sparsegeneral.go new file mode 100644 index 0000000..f716024 --- /dev/null +++ b/linalg/sparsegeneral.go @@ -0,0 +1,327 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The general (non-Hermitian) sparse eigenproblem. Lanczos +// needs symmetry: its three-term recurrence exists because the Krylov +// space of a symmetric operator comes with an orthogonal basis for +// free. A general operator needs the full Arnoldi recurrence, one +// explicit orthogonalisation per column, and its projected matrix is +// upper Hessenberg rather than tridiagonal. Resonance problems in +// electromagnetics and damped oscillation systems produce exactly +// these operators, complex and non-Hermitian included. +// +// Like SpEigen, this is a one-shot Krylov method: the basis runs to +// min(n, max(2k, k+40)) columns, the small Hessenberg eigenproblem is +// solved exactly by the dense EigenGeneral, and the returned pairs are +// Ritz approximations whose accuracy improves with the budget. The +// restart-on-breakdown scheme mirrors lanczos: an exhausted direction +// closes its block with a zero coupling and a fresh orthogonal random +// direction opens the next. + +// SpEigenGeneral returns the k eigenvalues of largest magnitude of a +// general real sparse matrix, each with its unit eigenvector: values +// is a complex128 vector ordered by descending magnitude (a real +// matrix may carry complex conjugate pairs), vectors a complex128 +// (n, k) array whose column j belongs to values[j]. The matrix must be +// square and real; complex input belongs to SpEigenGeneralComplex, +// symmetric input gets a cheaper answer from SpEigen. gen seeds the +// start vector (nil uses a fixed seed). Like the dense EigenGeneral, +// and unlike the symmetry-checking SpEigen, entries are not screened +// for finiteness: a non-finite entry answers NaN Ritz pairs. +func SpEigenGeneral(s *core.SparseCOO, k int, gen *core.Generator) (values, vectors *core.Array, err error) { + const name = "SpEigenGeneral" + if s.Values.Dtype() == core.Complex { + return nil, nil, base.Errf("%s: complex matrices belong to SpEigenGeneralComplex", name) + } + // The same square 2-D contract SpEigen enforces, and for the same + // reason: the Krylov product indexes x by the stored column, so a + // column beyond the row count reads out of range, and the read + // happens inside a worker goroutine the caller cannot recover. + // Ranks above two are refused too: the index walk would silently + // ignore every dimension past the first two. + if len(s.Shape) != 2 || s.Shape[0] != s.Shape[1] { + return nil, nil, base.Errf("%s: needs a square 2-D sparse matrix, got shape %v", name, s.Shape) + } + c, err := cooToCSR(s, name) + if err != nil { + return nil, nil, err + } + return arnoldiEigen(c.matVec, c.n, k, gen, name, krylovF64) +} + +// SpEigenGeneralComplex is SpEigenGeneral for complex128 sparse +// matrices: non-Hermitian operators included, which is the +// electromagnetics case. +func SpEigenGeneralComplex(s *core.SparseCOO, k int, gen *core.Generator) (values, vectors *core.Array, err error) { + const name = "SpEigenGeneralComplex" + c, err := cooToComplexCSR(s, name) + if err != nil { + return nil, nil, err + } + return arnoldiEigen(c.matVec, c.n, k, gen, name, krylovC128) +} + +// arnoldiEigen runs the shared pipeline: Arnoldi basis, dense +// eigensolve of the projected Hessenberg, Ritz lift. The matvec +// closure hides whether the operator is real or complex, and kern +// supplies the element-type-specific primitives the basis arithmetic +// needs. +func arnoldiEigen[T scalar](matvec func(x, y []T), n, k int, gen *core.Generator, name string, kern krylovKernel[T]) (*core.Array, *core.Array, error) { + if n == 0 { + return nil, nil, base.Errf("%s: zero-sized matrix", name) + } + if k < 1 || k > n { + return nil, nil, base.Errf("%s: k must be in [1, %d], got %d", name, n, k) + } + if gen == nil { + gen = core.NewGenerator(spEigenSeed) + } + m := min(n, max(2*k, k+spEigenBlock)) + h, v := arnoldi(matvec, n, m, gen, kern) + + // The projected problem: the leading m×m block of the Hessenberg, + // dense and small enough for the general eigensolver. + hArr, err := zeros(core.Complex, []int{m, m}) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + hc := hArr.RawComplexes() + for j := range m { + for i := 0; i <= j+1 && i < m; i++ { + hc[i*m+j] = toComplex(h[i*m+j]) + } + } + vals, vecs, err := EigenGeneral(hArr) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + if k > vals.Len() { + k = vals.Len() + } + outVals, err := zeros(core.Complex, []int{k}) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + vc := outVals.RawComplexes() + outVecs := make([]complex128, n*k) + tmp := make([]complex128, n) + for j := range k { + vc[j] = vals.ComplexAt(j) + // Ritz vector: lift the Hessenberg eigenvector through the + // Arnoldi basis, u = V·y. + clear(tmp) + for p := range m { + y := vecs.ComplexAt(p*m + j) + if y == 0 { + continue + } + kern.widenAxpy(tmp, v[p*n:(p+1)*n], y) + } + norm := 0.0 + for _, z := range tmp { + norm += real(z)*real(z) + imag(z)*imag(z) + } + norm = math.Sqrt(norm) + if norm > 0 { + for i := range n { + tmp[i] /= complex(norm, 0) + } + } + for i := range n { + outVecs[i*k+j] = tmp[i] + } + } + ritzVecs, err := complexFromArray2D(outVecs, n, k) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + return outVals, ritzVecs, nil +} + +// toComplex widens a scalar kernel value for the complex lift. +func toComplex[T scalar](v T) complex128 { + switch x := any(v).(type) { + case float64: + return complex(x, 0) + case complex128: + return x + } + return 0 +} + +// arnoldi builds the Krylov basis v (columns of length n, m+1 of +// them) and the upper Hessenberg h ((m+1)×m, row-major) of the +// operator behind matvec. Full reorthogonalisation, twice per column, +// keeps the basis orthonormal to rounding level; a collapsed direction +// closes its block with a zero coupling and restarts in a fresh +// orthogonal random direction. +func arnoldi[T scalar](matvec func(x, y []T), n, m int, gen *core.Generator, kern krylovKernel[T]) (h, v []T) { + // The basis never exceeds the space it spans: m above n would let + // the truncation branch below hand back a shorter h than the m×m + // read loop expects. + if m > n { + panic("arnoldi: the block size exceeds the matrix order") + } + h = make([]T, (m+1)*m) + v = make([]T, n*(m+1)) + w := make([]T, n) + scale := 0.0 + addScale := func(z T) { + if a := absOf(z); a > scale { + scale = a + } + } + // normT is the one accumulation in the recurrence that squares its + // inputs. A column whose entries are above about 1.3e154 overflows + // the raw sum of squares to +Inf (the projected coupling becomes + // infinite and the Ritz values come back wrong without a word), and + // one below about 1.5e-162 underflows it to zero (the exhaustion + // test reads that as a collapsed Krylov block and closes the block + // on a direction the operator never exhausted). The sum therefore + // runs raw while it lands in the normal range, which keeps the + // ordinary-scale arithmetic bit-for-bit, and falls back to a + // max-scaled accumulation outside it, the way norm2F64 and normC + // do. NaN is carried through the raw path so a poisoned input stays + // poisoned rather than reading as an exhausted block. + normT := func(a []T) float64 { + if s := absOf(kern.dot(a, a)); s <= math.MaxFloat64 && (s >= 0x1p-1022 || math.IsNaN(s)) { + return math.Sqrt(s) + } + maxAbs := 0.0 + for _, z := range a { + if m := absOf(z); m > maxAbs { + maxAbs = m + } + } + if maxAbs == 0 { + return 0 + } + s := 0.0 + for _, z := range a { + r := absOf(z) / maxAbs + s += r * r + } + return maxAbs * math.Sqrt(s) + } + start := func(col int) { + q := v[col*n : (col+1)*n] + kern.fillRand(q, gen) + for range 2 { + for p := range col { + row := v[p*n : (p+1)*n] + d := kern.dot(row, q) + for i := range n { + q[i] -= d * row[i] + } + } + } + if norm := normT(q); norm > 0 { + kern.scaleInto(q, q, 1/norm) + } + } + start(0) + for j := range m { + q := v[j*n : (j+1)*n] + matvec(q, w) + for i := 0; i <= j; i++ { + row := v[i*n : (i+1)*n] + d := kern.dot(row, w) + h[i*m+j] = d + for t := range n { + w[t] -= d * row[t] + } + } + // Full reorthogonalisation, twice: the first pass removes the + // accumulated loss, the second what the first reintroduces. + for range 2 { + for i := 0; i <= j; i++ { + row := v[i*n : (i+1)*n] + d := kern.dot(row, w) + h[i*m+j] += d + for t := range n { + w[t] -= d * row[t] + } + } + } + beta := normT(w) + addScale(h[j*m+j]) + if j+1 < m { + if beta <= float64(n)*base.EpsF*scale { + // The Krylov space of this block is exhausted; close it + // with a zero coupling and restart in a fresh + // direction. m never exceeds n (asserted at the top), + // so the basis always has room for the restart column. + start(j + 1) + continue + } + h[(j+1)*m+j] = realT[T](beta) + kern.scaleInto(v[(j+1)*n:(j+2)*n], w, 1/beta) + } + } + return h, v +} + +// krylovKernel holds the element-type-specific primitives the shared +// Arnoldi recurrence calls once per vector pass: the conjugated inner +// product, the real scaling of a vector, the start-vector fill and the +// widening lift of a real basis row into the complex Ritz vector. +// Passing them in keeps one copy of the recurrence's arithmetic while +// the element type is dispatched on once per pass instead of once per +// element. +type krylovKernel[T scalar] struct { + dot func(a, b []T) T + scaleInto func(dst, src []T, s float64) + fillRand func(dst []T, gen *core.Generator) + widenAxpy func(dst []complex128, src []T, y complex128) +} + +// krylovF64 is the kernel of a real operator. +var krylovF64 = krylovKernel[float64]{ + dot: dotF64, + scaleInto: func(dst, src []float64, s float64) { + for i := range dst { + dst[i] = src[i] * s + } + }, + fillRand: func(dst []float64, gen *core.Generator) { + for i := range dst { + dst[i] = gen.NormalUnit() + } + }, + widenAxpy: func(dst []complex128, src []float64, y complex128) { + for i, v := range src { + dst[i] += y * complex(v, 0) + } + }, +} + +// krylovC128 is the kernel of a complex operator: the inner product +// conjugates its left operand, and the start vector draws a real and an +// imaginary part from the generator, in that order. +var krylovC128 = krylovKernel[complex128]{ + dot: dotC, + scaleInto: func(dst, src []complex128, s float64) { + for i := range dst { + dst[i] = src[i] * complex(s, 0) + } + }, + fillRand: func(dst []complex128, gen *core.Generator) { + for i := range dst { + dst[i] = complex(gen.NormalUnit(), gen.NormalUnit()) + } + }, + widenAxpy: func(dst []complex128, src []complex128, y complex128) { + for i, v := range src { + dst[i] += y * v + } + }, +} diff --git a/linalg/sparsegeneral_test.go b/linalg/sparsegeneral_test.go new file mode 100644 index 0000000..56e8373 --- /dev/null +++ b/linalg/sparsegeneral_test.go @@ -0,0 +1,233 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestSpEigenGeneralRotationBlocks pins the real general solver: block +// rotations give complex conjugate eigenvalue pairs a symmetric-only +// method could never reach, and the answer must agree with the dense +// EigenGeneral on the same matrix. +func TestSpEigenGeneralRotationBlocks(t *testing.T) { + // Block diagonal: rotations by 5 and 2 plus one 7: spectrum + // {7, ±5i, ±2i}. + dense := []float64{ + 0, -5, 0, 0, 0, + 5, 0, 0, 0, 0, + 0, 0, 0, -2, 0, + 0, 0, 2, 0, 0, + 0, 0, 0, 0, 7, + } + d, err := core.FromFloats(dense, 5, 5) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.SparseFrom(d) + if err != nil { + t.Fatalf("SparseFrom: %v", err) + } + vals, vecs, err := SpEigenGeneral(sp, 3, core.NewGenerator(31)) + if err != nil { + t.Fatalf("SpEigenGeneral: %v", err) + } + want, _, err := EigenGeneral(d) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + // The spectrum is compared as a multiset against the dense one: a + // conjugate pair shares one magnitude, so which of the two a Krylov + // run reports at which index is an arbitrary tie-break of its own + // rounding, not a property of the matrix, and pinning it made the + // test depend on the last bit of a magnitude. Every computed value + // must still match some dense value, and the residual loop below + // pins each value to its own vector. + used := make([]bool, 3) + for j := range 3 { + got := vals.ComplexAt(j) + best, bestDist := -1, math.Inf(1) + for i := range 3 { + if used[i] { + continue + } + w := want.ComplexAt(i) + if dist := math.Hypot(real(got)-real(w), imag(got)-imag(w)); dist < bestDist { + best, bestDist = i, dist + } + } + if best < 0 || bestDist > 1e-8 { + t.Fatalf("value[%d] = %v matches no dense value (closest %v at distance %.3g)", + j, got, want.ComplexAt(best), bestDist) + } + used[best] = true + } + // Residuals ‖A·v − λ·v‖ against the original sparse operator. The + // eigenvectors may carry any complex phase, so the full complex + // vector enters the check. + for j := range 3 { + v := make([]complex128, 5) + for i := range 5 { + v[i] = vecs.ComplexAt(i*3 + j) + } + av := make([]complex128, 5) + for i := range 5 { + for p := range 5 { + av[i] += complex(dense[i*5+p], 0) * v[p] + } + } + lam := vals.ComplexAt(j) + for i := range 5 { + res := av[i] - lam*v[i] + if math.Hypot(real(res), imag(res)) > 1e-7 { + t.Fatalf("residual[%d][%d] = %v", j, i, res) + } + } + } +} + +// TestSpEigenGeneralComplexTriangular pins the complex general solver: +// a non-Hermitian triangular operator whose spectrum is its diagonal. +func TestSpEigenGeneralComplexTriangular(t *testing.T) { + // Upper triangular with distinct diagonal; the off-diagonal + // couplings make it genuinely non-normal. + entries := []complex128{ + 0, 0, 2 + 3i, + 1, 1, -1 + 1i, + 2, 2, 0.5 - 2i, + 0, 1, 0.7 + 0.3i, + 1, 2, -0.4, + } + idx := make([]int64, 0, 10) + valsIn := make([]complex128, 0, 5) + for i := 0; i+2 < len(entries); i += 3 { + idx = append(idx, int64(real(entries[i])), int64(real(entries[i+1]))) + valsIn = append(valsIn, entries[i+2]) + } + idxArr, err := core.FromInts(idx, len(valsIn), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + valArr, err := core.FromComplexes(valsIn, len(valsIn)) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + sp, err := core.NewSparseCOO(idxArr, valArr, []int{3, 3}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + vals, vecs, err := SpEigenGeneralComplex(sp, 2, core.NewGenerator(5)) + if err != nil { + t.Fatalf("SpEigenGeneralComplex: %v", err) + } + // |2+3i| ≈ 3.606 > |0.5−2i| ≈ 2.062 > |−1+i| ≈ 1.414. + wantTop := complex(2, 3) + wantSecond := complex(0.5, -2) + for j, want := range []complex128{wantTop, wantSecond} { + got := vals.ComplexAt(j) + if math.Abs(real(got)-real(want)) > 1e-9 || math.Abs(imag(got)-imag(want)) > 1e-9 { + t.Fatalf("value[%d] = %v, want %v", j, got, want) + } + } + // Residual against the dense operator. + dense := make([]complex128, 9) + for i := 0; i+2 < len(entries); i += 3 { + dense[int(real(entries[i]))*3+int(real(entries[i+1]))] = entries[i+2] + } + for j := range 2 { + v := make([]complex128, 3) + for i := range 3 { + v[i] = vecs.ComplexAt(i*2 + j) + } + var n2 float64 + for _, z := range v { + n2 += real(z)*real(z) + imag(z)*imag(z) + } + if math.Abs(n2-1) > 1e-8 { + t.Fatalf("vector %d norm² = %g, want 1", j, n2) + } + av := make([]complex128, 3) + for i := range 3 { + for p := range 3 { + av[i] += dense[i*3+p] * v[p] + } + } + lam := vals.ComplexAt(j) + for i := range 3 { + res := av[i] - lam*v[i] + if math.Hypot(real(res), imag(res)) > 1e-8 { + t.Fatalf("residual[%d][%d] = %v", j, i, res) + } + } + } +} + +// TestSpEigenGeneralLargerMatrix pins convergence on a bigger +// nonsymmetric operator: the top eigenvalues must match the dense +// reference. +func TestSpEigenGeneralLargerMatrix(t *testing.T) { + const n = 40 + g := core.NewGenerator(77) + dense := make([]float64, n*n) + for i := range n * n { + dense[i] = g.NormalUnit() + } + // A sprinkle of larger entries decides the spectrum's top end. + for i := range n { + dense[i*n+i] += 6 * float64(n-i) / float64(n) + } + d, err := core.FromFloats(dense, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.SparseFrom(d) + if err != nil { + t.Fatalf("SparseFrom: %v", err) + } + vals, _, err := SpEigenGeneral(sp, 2, core.NewGenerator(9)) + if err != nil { + t.Fatalf("SpEigenGeneral: %v", err) + } + want, _, err := EigenGeneral(d) + if err != nil { + t.Fatalf("EigenGeneral: %v", err) + } + // A real matrix carries conjugate twins of equal magnitude, so + // the k-th slot may hold either member: match against the set. + for j := range 2 { + got := vals.ComplexAt(j) + ok := cmplxAbs(got-want.ComplexAt(j)) <= 1e-6*cmplxAbs(got) || + cmplxAbs(got-complexConj(want.ComplexAt(j))) <= 1e-6*cmplxAbs(got) + if !ok { + t.Fatalf("value[%d] = %v, dense says %v (or its conjugate)", j, got, want.ComplexAt(j)) + } + } +} + +// TestSpEigenGeneralErrors pins the routing and validation. +func TestSpEigenGeneralErrors(t *testing.T) { + d, _ := core.FromFloats([]float64{0, -1, 1, 0}, 2, 2) + sp, _ := core.SparseFrom(d) + if _, _, err := SpEigenGeneral(sp, 0, nil); err == nil { + t.Error("k = 0 accepted") + } + if _, _, err := SpEigenGeneral(sp, 3, nil); err == nil { + t.Error("k > n accepted") + } + // Complex input routes to the complex entry point. + idx, _ := core.FromInts([]int64{0, 0, 1, 1}, 2, 2) + cv, _ := core.FromComplexes([]complex128{1, 2}, 2) + cc, _ := core.NewSparseCOO(idx, cv, []int{2, 2}) + if _, _, err := SpEigenGeneral(cc, 1, nil); err == nil { + t.Error("SpEigenGeneral accepted complex values") + } + rv, _ := core.FromFloats([]float64{1, 2}, 2) + rc, _ := core.NewSparseCOO(idx, rv, []int{2, 2}) + if _, _, err := SpEigenGeneralComplex(rc, 1, nil); err == nil { + t.Error("SpEigenGeneralComplex accepted real values") + } +} diff --git a/linalg/sparseigen.go b/linalg/sparseigen.go new file mode 100644 index 0000000..dd73dd9 --- /dev/null +++ b/linalg/sparseigen.go @@ -0,0 +1,432 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "slices" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Sparse symmetric eigensolver. The dense `Eigen` answers a full +// spectrum by tridiagonalising the whole matrix, which costs O(n³) +// and materialises n² entries; both are wasteful when the matrix is +// sparse and only a few extreme eigenpairs are wanted. SpEigen runs +// the Lanczos process instead, whose every step costs one +// sparse-matrix-vector product, then reads the Ritz pairs off the +// small tridiagonal matrix it builds. +// +// Lanczos with the plain three-term recurrence loses orthogonality +// between the basis vectors as it converges, which shows up as +// duplicated ("ghost") eigenvalues. The classical cure is full +// reorthogonalisation, applied here twice per step: one pass removes +// the loss, the second removes what the first reintroduces at rounding +// level. That costs O(m·n) per step but keeps the Ritz values honest, +// and m stays small because only a few eigenpairs are wanted. + +// spEigenSeed is the fixed seed used when SpEigen gets no generator, +// so an unseeded call is reproducible across runs. It is arbitrary but +// fixed; any non-zero seed avoids the all-ones start vector, which is +// orthogonal to every zero-sum eigenvector and would silently miss +// them. +const spEigenSeed = 20260912 + +// spEigenBlock is how many Lanczos steps to run past the number of +// eigenpairs asked for. The extra steps buy convergence of the wanted +// extremes: Ritz values settle well before Ritz vectors do, so the +// budget is set by the vectors. A matrix of n ≤ k+spEigenBlock is +// decomposed fully, which makes the answer exact rather than an +// approximation; the partial path is what large n takes. +const spEigenBlock = 40 + +// sparseCSR is a compressed-sparse-row view of a square real matrix, +// built once per solve so the repeated matrix-vector products of the +// Lanczos loop stream through contiguous value slices. +type sparseCSR struct { + rowStart []int + colIdx []int + vals []float64 + n int +} + +// cooToCSR converts a core.SparseCOO into CSR form, summing duplicate +// coordinates the way COO semantics require. name names the calling +// entry point in validation errors. +func cooToCSR(s *core.SparseCOO, name string) (*sparseCSR, error) { + ndim := len(s.Shape) + nnz := s.Indices.Shape()[0] + type entry struct { + row, col int + val float64 + } + entries := make([]entry, nnz) + ints := s.Indices.RawInts() + for i := range nnz { + row, col := 0, 0 + off := i * ndim + for d := range ndim { + idx := int(ints[off+d]) + if idx < 0 || idx >= s.Shape[d] { + return nil, base.Errf("%s: index [%d]=%d out of range for dim %d of size %d", + name, i, idx, d, s.Shape[d]) + } + if d == 0 { + row = idx + } + if d == 1 { + col = idx + } + } + entries[i] = entry{row: row, col: col, val: s.Values.FloatAt(i)} + } + // Sort by (row, col) so duplicates merge and each row is contiguous. + slices.SortFunc(entries, func(a, b entry) int { + if a.row != b.row { + return a.row - b.row + } + return a.col - b.col + }) + c := &sparseCSR{ + rowStart: make([]int, s.Shape[0]+1), + colIdx: make([]int, 0, nnz), + vals: make([]float64, 0, nnz), + n: s.Shape[0], + } + // The duplicates are adjacent after the sort, so each run merges into + // one accumulated value on the way into the compressed form. + for p := 0; p < len(entries); { + e := entries[p] + v := e.val + p++ + for p < len(entries) && entries[p].row == e.row && entries[p].col == e.col { + v += entries[p].val + p++ + } + if v == 0 { + continue + } + c.colIdx = append(c.colIdx, e.col) + c.vals = append(c.vals, v) + c.rowStart[e.row+1]++ + } + for i := range c.n { + c.rowStart[i+1] += c.rowStart[i] + } + return c, nil +} + +// symmetricCSR converts a core.SparseCOO into CSR form and verifies the +// matrix is symmetric. Asymmetry beyond a scale-relative 1e-12 is +// refused, matching the tolerance the dense `Eigen` applies. +func symmetricCSR(s *core.SparseCOO, name string) (*sparseCSR, error) { + c, err := cooToCSR(s, name) + if err != nil { + return nil, err + } + if err := c.checkSymmetric(name); err != nil { + return nil, err + } + return c, nil +} + +// checkSymmetric verifies that every stored entry has a matching +// transposed entry of the same value, the property the Lanczos and +// conjugate-gradient recurrences rely on. The tolerance is purely +// relative to the largest magnitude present, deliberately without an +// absolute floor: a floor would approve a matrix of small scale whose +// asymmetry is a large fraction of it, and the recurrence would then +// answer for a matrix that is neither A nor Aᵀ, whatever the wording of +// the guard claims. Relative-only keeps the approval and the arithmetic +// in agreement. A non-finite entry is refused in its own right: NaN +// compares unequal to everything, so the mirror test alone would wave +// it through as symmetric. The wording of the refusal is unchanged. +func (c *sparseCSR) checkSymmetric(name string) error { + scale := 0.0 + for _, v := range c.vals { + if a := math.Abs(v); a > scale { + scale = a + } + } + tol := 1e-12 * scale + for i := range c.n { + for p := c.rowStart[i]; p < c.rowStart[i+1]; p++ { + j := c.colIdx[p] + v := c.vals[p] + if math.IsNaN(v) || math.IsInf(v, 0) { + return base.Errf("%s: entry [%d,%d] is not finite", name, i, j) + } + mirror, ok := c.at(j, i) + if !ok || math.Abs(mirror-v) > tol { + return base.Errf("%s: matrix is not symmetric within 1e-12 tolerance", name) + } + } + } + return nil +} + +// at reads the entry (row, col), reporting whether it is stored. The +// row's column indices are sorted, so the lookup is a binary search: +// the symmetry check calls it once per stored entry, and a linear scan +// made that check quadratic in the row length. +func (c *sparseCSR) at(row, col int) (float64, bool) { + lo, hi := c.rowStart[row], c.rowStart[row+1] + for lo < hi { + mid := int(uint(lo+hi) >> 1) + if c.colIdx[mid] < col { + lo = mid + 1 + } else { + hi = mid + } + } + if lo < c.rowStart[row+1] && c.colIdx[lo] == col { + return c.vals[lo], true + } + return 0, false +} + +// matVec computes y = A·x. Output rows are independent, so the row +// range splits across workers once a worker's share of the stored +// entries pays for the split, and runs on the calling goroutine below +// that. +func (c *sparseCSR) matVec(x, y []float64) { + // The three structure slices are taken once: every worker's rows walk + // them per stored entry, and the row loop holds no other state. + rowStart := c.rowStart + colIdx := c.colIdx + vals := c.vals + if sparseMatVecSplit(c.n, len(vals)) { + engine.ParallelMin(c.n, 1, func(s, e int) { + csrMatVecRange(s, e, rowStart, colIdx, vals, x, y) + }) + return + } + csrMatVecRange(0, c.n, rowStart, colIdx, vals, x, y) +} + +// SpEigen returns the k eigenvalues of largest magnitude of a real +// symmetric sparse matrix, each with its unit eigenvector. Values are +// ordered by descending magnitude and vectors holds the matching +// eigenvectors as columns, so column j goes with values[j]. +// +// The matrix must be square, real and symmetric; complex sparse +// matrices are refused, as they are everywhere else on the sparse +// surface. Requesting k = n recovers the complete spectrum, though at +// that point the dense `Eigen` is the cheaper path. +// +// The starting vector is drawn from gen; passing nil uses a fixed +// seed, so the result is reproducible. Because Lanczos converges to +// the extremes of the spectrum first, the eigenvalues returned are +// approximations whose accuracy improves with the iteration budget, +// not the exact answers `Eigen` computes. +func SpEigen(s *core.SparseCOO, k int, gen *core.Generator) (values, vectors *core.Array, err error) { + if s.Values.Dtype() == core.Complex { + return nil, nil, base.Errf("SpEigen: complex sparse matrices are not supported") + } + if len(s.Shape) != 2 || s.Shape[0] != s.Shape[1] { + return nil, nil, base.Errf("SpEigen: needs a square 2-D sparse matrix, got shape %v", s.Shape) + } + n := s.Shape[0] + if n == 0 { + return nil, nil, base.Errf("SpEigen: zero-sized matrix, got shape %v", s.Shape) + } + if k < 1 || k > n { + return nil, nil, base.Errf("SpEigen: k must be in [1, %d], got %d", n, k) + } + c, err := symmetricCSR(s, "SpEigen") + if err != nil { + return nil, nil, err + } + // The projected tridiagonal carries the matrix's magnitude, and the + // symmetric sweep's deflation floor is taken against the sum of its + // squared entries: above about 1.3e154 that sum overflows to +Inf, so + // every subdiagonal reads as negligible, the sweep deflates the whole + // block at once and the raw Rayleigh quotients come back as the + // spectrum. The matrix is moved into the window for the recurrence and + // the Ritz values, which carry the scale, are moved back before they + // are returned; the basis vectors are unit-length and therefore + // scale-free. + ws := windowScale(maxMagF64(c.vals)) + if ws != 1 { + scaleFloats(c.vals, ws) + } + if gen == nil { + gen = core.NewGenerator(spEigenSeed) + } + alphas, betas, basis := c.lanczos(k, gen) + r := len(alphas) + + // The small tridiagonal matrix whose eigenpairs are the Ritz + // approximations. Its eigenvectors accumulate into the identity, + // so the columns of tVec are the tridiagonal eigenvectors. + tMat := make([]float64, r*r) + for i := range r { + tMat[i*r+i] = alphas[i] + if i+1 < r { + tMat[i*r+i+1] = betas[i] + tMat[(i+1)*r+i] = betas[i] + } + } + tVec := eye(r) + if err := symmetricQr(tMat, tVec, r); err != nil { + return nil, nil, base.Errf("SpEigen: %w", err) + } + + // Rank the Ritz pairs by descending magnitude and keep k. + idx := make([]int, r) + for i := range r { + idx[i] = i + } + slices.SortFunc(idx, func(a, b int) int { + da, db := math.Abs(tMat[a*r+a]), math.Abs(tMat[b*r+b]) + if da != db { + if da > db { + return -1 + } + return 1 + } + return a - b + }) + sel := idx[:k] + + outVals := make([]float64, k) + outVecs := make([]float64, n*k) + tmp := make([]float64, n) + for j, si := range sel { + outVals[j] = tMat[si*r+si] + // Ritz vector: lift the tridiagonal eigenvector through the + // Lanczos basis, v = Q·y. + clear(tmp) + for p := range r { + y := tVec[p*r+si] + if y == 0 { + continue + } + row := basis[p*n : (p+1)*n] + for i := range n { + tmp[i] += y * row[i] + } + } + normaliseF64(tmp) + for i := range n { + outVecs[i*k+j] = tmp[i] + } + } + // Only the Ritz values carry the matrix's scale; the vectors do not. + if ws != 1 { + unscaleFloats(outVals, ws) + } + return floatsToArray(outVals, []int{k}), floatsToArray(outVecs, []int{n, k}), nil +} + +// lanczos runs the process, returning the diagonal alpha, the +// off-diagonal beta and the basis vectors flattened as n-wide rows. +// When the recurrence collapses, meaning the Krylov space is +// exhausted, the block is closed off with a zero coupling and a fresh +// direction restarts it: the tridiagonal matrix becomes +// block-diagonal, and each block still yields valid Ritz pairs. +func (c *sparseCSR) lanczos(k int, gen *core.Generator) (alphas, betas []float64, basis []float64) { + n := c.n + steps := min(n, max(2*k, k+spEigenBlock)) + w := make([]float64, n) + q := make([]float64, n) + // The budget bounds the recurrence: one basis row and at most one + // coefficient per step, so all three are sized once rather than grown. + alphas = make([]float64, 0, steps) + betas = make([]float64, 0, steps) + basis = make([]float64, 0, steps*n) + scale := 0.0 + addScale := func(v float64) { + if a := math.Abs(v); a > scale { + scale = a + } + } + start := func() { + for i := range n { + q[i] = gen.NormalUnit() + } + // Orthogonalise against every direction already taken so a + // restarted block cannot re-enter an earlier one. + for p := range len(alphas) { + row := basis[p*n : (p+1)*n] + d := dotF64(q, row) + for i := range n { + q[i] -= d * row[i] + } + } + normaliseF64(q) + } + start() + for range steps { + basis = append(basis, q...) + c.matVec(q, w) + alpha := dotF64(q, w) + alphas = append(alphas, alpha) + addScale(alpha) + // Strip the two previous directions out of w, the three-term + // recurrence itself. + for i := range n { + w[i] -= alpha * q[i] + } + if len(alphas) > 1 { + beta := betas[len(betas)-1] + prev := basis[(len(alphas)-2)*n : (len(alphas)-1)*n] + for i := range n { + w[i] -= beta * prev[i] + } + } + // Full reorthogonalisation, twice: one pass removes the + // accumulated loss, the second what the first reintroduces. + for range 2 { + for p := range len(alphas) { + row := basis[p*n : (p+1)*n] + d := dotF64(w, row) + for i := range n { + w[i] -= d * row[i] + } + } + } + beta := norm2F64(w) + // The exhaustion threshold is purely relative to the running + // spectral scale: an absolute floor would treat a legitimate + // tiny-scale matrix as one big deflated block and hand back + // random-start Rayleigh quotients instead of the spectrum. + if beta <= float64(n)*base.EpsF*scale { + // The Krylov space is exhausted. Close the block with a + // zero coupling and restart in a fresh direction, unless + // the basis already spans the whole space. + if len(alphas) >= n { + break + } + betas = append(betas, 0) + start() + continue + } + betas = append(betas, beta) + addScale(beta) + for i := range n { + q[i] = w[i] / beta + } + } + // The last beta couples to a vector that was never formed, so it + // does not belong on the tridiagonal diagonal-band. + if len(betas) >= len(alphas) { + betas = betas[:len(alphas)-1] + } + return alphas, betas, basis +} + +// normaliseF64 scales a vector to unit length in place, leaving a +// zero vector untouched rather than dividing by zero. +func normaliseF64(a []float64) { + n := norm2F64(a) + if n == 0 { + return + } + for i := range a { + a[i] /= n + } +} diff --git a/linalg/sparseigen_test.go b/linalg/sparseigen_test.go new file mode 100644 index 0000000..a52b055 --- /dev/null +++ b/linalg/sparseigen_test.go @@ -0,0 +1,503 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// sparseFromDense builds a symmetric SparseCOO from a dense matrix, +// failing the test if the input is not symmetric enough to qualify. +func sparseFromDense(t *testing.T, vals []float64, n int) *core.SparseCOO { + t.Helper() + dense, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.SparseFrom(dense) + if err != nil { + t.Fatalf("SparseFrom: %v", err) + } + return sp +} + +// residual norms ‖A·v − λ·v‖ for one eigenpair, the direct check that +// a returned pair is actually an eigenpair of the original matrix. +// vec holds eigenvectors as columns of an (n, k) array. +func residual(t *testing.T, sp *core.SparseCOO, vec *core.Array, lam float64, n, k, j int) float64 { + t.Helper() + x := make([]float64, n) + for i := range n { + x[i] = vec.FloatAt(i*k + j) + } + dense, err := sp.Dense() + if err != nil { + t.Fatalf("Dense: %v", err) + } + r := 0.0 + for i := range n { + sum := 0.0 + for p := range n { + sum += dense.FloatAt(i*n+p) * x[p] + } + r += (sum - lam*x[i]) * (sum - lam*x[i]) + } + return math.Sqrt(r) +} + +// TestSpEigenMatchesDense compares the sparse solver against the dense +// `Eigen` on the same matrix: the k largest-magnitude eigenvalues must +// agree, which is the property that matters for a partial solver. +func TestSpEigenMatchesDense(t *testing.T) { + cases := []struct { + name string + vals []float64 + n int + k int + }{ + { + name: "diagonal", + vals: []float64{ + 5, 0, 0, 0, + 0, -3, 0, 0, + 0, 0, 2, 0, + 0, 0, 0, -7, + }, + n: 4, k: 2, + }, + { + name: "tridiagonal", + vals: []float64{ + 4, 1, 0, 0, + 1, 4, 1, 0, + 0, 1, 4, 1, + 0, 0, 1, 4, + }, + n: 4, k: 2, + }, + { + name: "full_symmetric", + vals: []float64{ + 2, -1, 0.5, 0.3, + -1, 3, 0.2, -0.4, + 0.5, 0.2, 1, 0.6, + 0.3, -0.4, 0.6, 2.5, + }, + n: 4, k: 3, + }, + { + name: "repeated_spectrum", + vals: []float64{ + 1, 0, 0, 0, + 0, 1, 0, 0, + 0, 0, -2, 0, + 0, 0, 0, -2, + }, + n: 4, k: 2, + }, + } + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + sp := sparseFromDense(t, tt.vals, tt.n) + gotVals, gotVecs, err := SpEigen(sp, tt.k, nil) + if err != nil { + t.Fatalf("SpEigen: %v", err) + } + if gotVals.Len() != tt.k { + t.Fatalf("values len %d, want %d", gotVals.Len(), tt.k) + } + if gotVecs.Shape()[0] != tt.n || gotVecs.Shape()[1] != tt.k { + t.Fatalf("vectors shape %s, want [%d %d]", base.ShapeText(gotVecs.Shape()), tt.n, tt.k) + } + + dense, err := core.FromFloats(tt.vals, tt.n, tt.n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + refVals, _, err := Eigen(dense) + if err != nil { + t.Fatalf("Eigen: %v", err) + } + // Reference extremes by descending magnitude. + ref := make([]float64, tt.n) + for i := range tt.n { + ref[i] = refVals.FloatAt(i) + } + for i := range tt.k { + for j := i + 1; j < tt.n; j++ { + if math.Abs(ref[j]) > math.Abs(ref[i]) { + ref[i], ref[j] = ref[j], ref[i] + } + } + } + for i := range tt.k { + g, w := gotVals.FloatAt(i), ref[i] + if math.Abs(g-w) > 1e-8*(1+math.Abs(w)) { + t.Fatalf("value[%d] = %.12g, want %.12g", i, g, w) + } + } + + // Each returned pair must be an eigenpair of A itself. + for j := range tt.k { + r := residual(t, sp, gotVecs, gotVals.FloatAt(j), tt.n, tt.k, j) + if r > 1e-8 { + t.Fatalf("eigenpair %d: residual ‖Av-λv‖ = %.3g, want <= 1e-8", j, r) + } + } + }) + } +} + +// TestSpEigenFullSpectrum checks that k = n recovers every eigenvalue, +// matching the dense solver's complete output. +func TestSpEigenFullSpectrum(t *testing.T) { + const n = 5 + vals := []float64{ + 4, 1, 0, 0, 0, + 1, 3, 2, 0, 0, + 0, 2, 1, -1, 0, + 0, 0, -1, 2, 0.5, + 0, 0, 0, 0.5, 5, + } + sp := sparseFromDense(t, vals, n) + gotVals, _, err := SpEigen(sp, n, nil) + if err != nil { + t.Fatalf("SpEigen: %v", err) + } + dense, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + refVals, _, err := Eigen(dense) + if err != nil { + t.Fatalf("Eigen: %v", err) + } + got := make([]float64, n) + for i := range n { + got[i] = gotVals.FloatAt(i) + } + for i := range n { + for j := i + 1; j < n; j++ { + if math.Abs(got[j]) > math.Abs(got[i]) { + got[i], got[j] = got[j], got[i] + } + } + } + ref := make([]float64, n) + for i := range n { + ref[i] = refVals.FloatAt(i) + } + for i := range n { + for j := i + 1; j < n; j++ { + if math.Abs(ref[j]) > math.Abs(ref[i]) { + ref[i], ref[j] = ref[j], ref[i] + } + } + } + for i := range n { + if math.Abs(got[i]-ref[i]) > 1e-8*(1+math.Abs(ref[i])) { + t.Fatalf("value[%d] = %.12g, want %.12g", i, got[i], ref[i]) + } + } +} + +// TestSpEigenLargerMatrix exercises the solver where it earns its +// keep: a sparse matrix far larger than the eigenpair count, so the +// answer comes from a partial decomposition rather than a full one. +// The spectrum is deliberately spread out with well-separated +// extremes, which is the regime Lanczos converges in; a clustered +// spectrum needs a far larger iteration budget than a partial solve +// normally gets. +func TestSpEigenLargerMatrix(t *testing.T) { + const n = 24 + // A diagonally dominant tridiagonal matrix: diagonal spaced from + // +n down to +1, off-diagonals 1. The diagonal is kept positive so + // the spectrum is not symmetric about zero: a spectrum with ±pairs + // of equal magnitude would leave the descending-magnitude ordering + // with ties, and the test could not compare position by position. + // The extremes sit well clear of the rest, so the top Ritz values + // are accurate. + vals := make([]float64, n*n) + for i := range n { + vals[i*n+i] = float64(n - i) + if i+1 < n { + vals[i*n+i+1] = 1 + vals[(i+1)*n+i] = 1 + } + } + sp := sparseFromDense(t, vals, n) + const k = 3 + gotVals, gotVecs, err := SpEigen(sp, k, nil) + if err != nil { + t.Fatalf("SpEigen: %v", err) + } + dense, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + refVals, _, err := Eigen(dense) + if err != nil { + t.Fatalf("Eigen: %v", err) + } + ref := make([]float64, n) + for i := range n { + ref[i] = refVals.FloatAt(i) + } + for i := range k { + for j := i + 1; j < n; j++ { + if math.Abs(ref[j]) > math.Abs(ref[i]) { + ref[i], ref[j] = ref[j], ref[i] + } + } + } + for i := range k { + g, w := gotVals.FloatAt(i), ref[i] + if math.Abs(g-w) > 1e-8*(1+math.Abs(w)) { + t.Fatalf("value[%d] = %.12g, want %.12g", i, g, w) + } + } + for j := range k { + // A partial solve delivers ‖Av-λv‖ ≈ β_m·|y_m|, the last + // coupling times the tail of the tridiagonal eigenvector, + // rather than machine epsilon. The eigenvalues above are + // accurate to 1e-8 because the extremes are well separated; + // the vectors carry the looser bound. + r := residual(t, sp, gotVecs, gotVals.FloatAt(j), n, k, j) + if r > 1e-5 { + t.Fatalf("eigenpair %d: residual ‖Av-λv‖ = %.3g, want <= 1e-5", j, r) + } + } +} + +// TestSpEigenPartialDecomposition exercises the path the sparse +// eigensolver exists for: n far larger than the eigenpair count, so +// the iteration budget stops well short of the dimension and the +// answer is a truncated decomposition rather than an exact one. The +// smaller tests all take n ≤ k+40, which decomposes fully and never +// reaches this branch. +func TestSpEigenPartialDecomposition(t *testing.T) { + const n = 120 + // A diagonally dominant tridiagonal matrix whose diagonal falls off + // like 1000/(i+1), so the top few eigenvalues are widely separated + // relative to their magnitude. That separation is what lets a + // truncated Lanczos converge on the wanted extremes inside the + // budget; a tightly clustered spectrum (a near-constant diagonal, + // say) converges far slower and would make the assertion about the + // iteration budget rather than about the solver. The off-diagonal + // of 1 is negligible against the leading diagonal entries, so the + // spectrum stays close to the diagonal. + vals := make([]float64, n*n) + for i := range n { + vals[i*n+i] = 1000 / float64(i+1) + if i+1 < n { + vals[i*n+i+1] = 1 + vals[(i+1)*n+i] = 1 + } + } + sp := sparseFromDense(t, vals, n) + const k = 4 + gotVals, gotVecs, err := SpEigen(sp, k, nil) + if err != nil { + t.Fatalf("SpEigen: %v", err) + } + dense, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + refVals, _, err := Eigen(dense) + if err != nil { + t.Fatalf("Eigen: %v", err) + } + ref := make([]float64, n) + for i := range n { + ref[i] = refVals.FloatAt(i) + } + for i := range k { + for j := i + 1; j < n; j++ { + if math.Abs(ref[j]) > math.Abs(ref[i]) { + ref[i], ref[j] = ref[j], ref[i] + } + } + } + // The budget stops at k+40 = 44 of n = 120 steps, and the diagonal + // is tightly clustered near the top, so the extremes are converged + // to a relative few times 1e-7 rather than exactly. A tolerance + // tighter than that would be asserting an exactness a truncated + // decomposition does not claim. + for i := range k { + g, w := gotVals.FloatAt(i), ref[i] + if math.Abs(g-w) > 1e-5*(1+math.Abs(w)) { + t.Fatalf("value[%d] = %.12g, want %.12g", i, g, w) + } + } + for j := range k { + r := residual(t, sp, gotVecs, gotVals.FloatAt(j), n, k, j) + if r > 1e-4 { + t.Fatalf("eigenpair %d: residual ‖Av-λv‖ = %.3g, want <= 1e-4", j, r) + } + } +} + +// TestSpEigenSumsDuplicates pins the COO contract: a coordinate listed +// twice contributes the sum of both values, not whichever comes last. +func TestSpEigenSumsDuplicates(t *testing.T) { + // Build [[3,0],[0,1]] as (0,0)=1 plus (0,0)=2, plus the lower entry. + idx, err := core.FromInts([]int64{0, 0, 0, 0, 1, 1}, 3, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + vals, err := core.FromFloats([]float64{1, 2, 1}, 3) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.NewSparseCOO(idx, vals, []int{2, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + gotVals, _, err := SpEigen(sp, 2, nil) + if err != nil { + t.Fatalf("SpEigen: %v", err) + } + if math.Abs(gotVals.FloatAt(0)-3) > 1e-12 { + t.Fatalf("largest value = %.12g, want 3 (duplicates must sum)", gotVals.FloatAt(0)) + } + if math.Abs(gotVals.FloatAt(1)-1) > 1e-12 { + t.Fatalf("second value = %.12g, want 1", gotVals.FloatAt(1)) + } +} + +// TestSpEigenDeterminism checks that the unseeded call is reproducible +// and that an explicit seed is honoured, the contract callers rely on +// when they pin a start vector. +func TestSpEigenDeterminism(t *testing.T) { + const n = 6 + vals := make([]float64, n*n) + for i := range n { + vals[i*n+i] = float64(i) - 3 + if i+1 < n { + vals[i*n+i+1] = 0.5 + vals[(i+1)*n+i] = 0.5 + } + } + sp := sparseFromDense(t, vals, n) + v1, u1, err := SpEigen(sp, 3, nil) + if err != nil { + t.Fatalf("SpEigen #1: %v", err) + } + v2, u2, err := SpEigen(sp, 3, nil) + if err != nil { + t.Fatalf("SpEigen #2: %v", err) + } + for i := range 3 { + if v1.FloatAt(i) != v2.FloatAt(i) { + t.Fatalf("unseeded run %d: value %.12g vs %.12g", i, v1.FloatAt(i), v2.FloatAt(i)) + } + } + for i := range u1.Len() { + if u1.FloatAt(i) != u2.FloatAt(i) { + t.Fatalf("unseeded run: vector element %d differs", i) + } + } + // A different seed may pick a different start vector, but the + // eigenvalues of a symmetric matrix do not depend on it. + v3, _, err := SpEigen(sp, 3, core.NewGenerator(99)) + if err != nil { + t.Fatalf("SpEigen seeded: %v", err) + } + for i := range 3 { + if math.Abs(v1.FloatAt(i)-v3.FloatAt(i)) > 1e-8 { + t.Fatalf("seeded run %d: value %.12g vs %.12g", i, v1.FloatAt(i), v3.FloatAt(i)) + } + } +} + +// TestSpEigenRejectsInvalid pins the error contract for every input +// the solver cannot honestly answer. +func TestSpEigenRejectsInvalid(t *testing.T) { + t.Run("complex_sparse", func(t *testing.T) { + idx, err := core.FromInts([]int64{0, 0}, 1, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + vals, err := core.FromComplexes([]complex128{1}, 1) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + sp, err := core.NewSparseCOO(idx, vals, []int{1, 1}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, _, err := SpEigen(sp, 1, nil); err == nil { + t.Fatal("expected an error for a complex sparse matrix") + } + }) + t.Run("not_square", func(t *testing.T) { + idx, err := core.FromInts([]int64{0, 0}, 1, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + vals, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.NewSparseCOO(idx, vals, []int{1, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, _, err := SpEigen(sp, 1, nil); err == nil { + t.Fatal("expected an error for a non-square matrix") + } + }) + t.Run("zero_sized", func(t *testing.T) { + sp := &core.SparseCOO{Shape: []int{0, 0}} + idx, err := core.FromInts(nil, 0, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + vals, err := core.FromFloats(nil, 0) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp.Indices, sp.Values = idx, vals + if _, _, err := SpEigen(sp, 1, nil); err == nil { + t.Fatal("expected an error for a zero-sized matrix") + } + }) + t.Run("k_out_of_range", func(t *testing.T) { + sp := sparseFromDense(t, []float64{1, 0, 0, 1}, 2) + if _, _, err := SpEigen(sp, 0, nil); err == nil { + t.Fatal("expected an error for k = 0") + } + if _, _, err := SpEigen(sp, 3, nil); err == nil { + t.Fatal("expected an error for k > n") + } + }) + t.Run("asymmetric", func(t *testing.T) { + // A is not symmetric: A[0,1] = 1 but A[1,0] = 2. + sp := sparseFromDense(t, []float64{1, 1, 2, 1}, 2) + if _, _, err := SpEigen(sp, 1, nil); err == nil { + t.Fatal("expected an error for an asymmetric matrix") + } + }) + t.Run("index_out_of_range", func(t *testing.T) { + idx, err := core.FromInts([]int64{0, 5}, 1, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + vals, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.NewSparseCOO(idx, vals, []int{2, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, _, err := SpEigen(sp, 1, nil); err == nil { + t.Fatal("expected an error for an out-of-range index") + } + }) +} diff --git a/linalg/sparseilu.go b/linalg/sparseilu.go new file mode 100644 index 0000000..d54b963 --- /dev/null +++ b/linalg/sparseilu.go @@ -0,0 +1,197 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// ILU(0): the incomplete LU factorisation whose L and U factors keep +// exactly the sparsity pattern of A, dropping every fill-in entry the +// complete factorisation would create. The result is not a solver but +// a preconditioner: applying it approximates A⁻¹ at the cost of one +// forward and one backward substitution over the stored pattern, and +// the Krylov solvers converge in far fewer steps than under plain +// Jacobi scaling because the approximation carries the matrix's local +// coupling. For a tridiagonal matrix the complete factorisation +// creates no fill anyway, so ILU(0) equals the complete LU there; the +// wider the stencil, the bigger the accuracy gap and the cheaper the +// factorisation stays. +// +// The factorisation runs the row-oriented IKJ elimination: row i's +// left part scales by the pivots above it, and each elimination +// updates only the positions the pattern already stores, intersecting +// row i with the pivot row by a two-pointer walk over the sorted +// column indices. Cost tracks the non-zero count, not n³. + +// SparseILU holds the ILU(0) factorisation of a square sparse matrix +// in its CSR pattern: entries left of the diagonal are L multipliers +// with the unit diagonal implied, entries at or right of it are U +// values including the pivots on the diagonal. +type SparseILU struct { + vals []float64 + colIdx []int + rowStart []int + n int +} + +// NewSparseILU factors a square sparse matrix with ILU(0). Column +// indices within each row must be sorted ascending, which the COO to +// CSR conversions guarantee. A missing diagonal entry or a zero pivot +// is refused: the factorisation divides by both, and an incomplete +// factorisation cannot repair a structurally singular matrix. +// Non-finite entries are refused, exactly as NewSparseLU and +// NewSparseCholesky refuse them. +func NewSparseILU(a *core.SparseCOO) (*SparseILU, error) { + if a.Values.Dtype() == core.Complex { + return nil, base.Errf("NewSparseILU: complex sparse matrices are not supported") + } + if len(a.Shape) != 2 || a.Shape[0] != a.Shape[1] { + return nil, base.Errf("NewSparseILU: needs a square 2-D sparse matrix, got shape %v", a.Shape) + } + c, err := cooToCSR(a, "NewSparseILU") + if err != nil { + return nil, err + } + for i := range c.n { + for p := c.rowStart[i]; p < c.rowStart[i+1]; p++ { + v := c.vals[p] + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("NewSparseILU: entry [%d,%d] is not finite", i, c.colIdx[p]) + } + } + } + n := c.n + // Every row needs a pivot: the forward substitution divides by the + // pivots above the diagonal and the backward substitution by every + // diagonal, so a row without one is structurally singular. The + // elimination loop below only inspects the diagonals a later row + // actually pivots on, which never includes the last row (and misses + // any other row nothing below it references), so the scan is what + // makes the documented refusal hold for the whole matrix. + for i := range n { + if c.atDiagonal(i) == 0 { + return nil, base.Errf("NewSparseILU: missing or zero diagonal at row %d", i) + } + } + for i := 1; i < n; i++ { + rs, re := c.rowStart[i], c.rowStart[i+1] + // Walk row i's left part; every pivot row k < i is final by + // the time row i is processed. + for p := rs; p < re && c.colIdx[p] < i; p++ { + k := c.colIdx[p] + pivot := c.atDiagonal(k) + if pivot == 0 { + return nil, base.Errf("NewSparseILU: zero pivot at row %d", k) + } + mult := c.vals[p] / pivot + c.vals[p] = mult + // Two-pointer intersection of row i right of column k + // with row k right of column k. The elimination can + // overflow on legal finite input (the siblings in LU and + // Cholesky refuse the same), and Apply has no error + // return, so a poisoned factor must never leave here. + if math.IsInf(mult, 0) || math.IsNaN(mult) { + return nil, base.Errf("NewSparseILU: the elimination overflowed at row %d, column %d (multiplier %g)", i, k, mult) + } + q := c.rowStart[k] + kre := c.rowStart[k+1] + for q < kre && c.colIdx[q] <= k { + q++ + } + t := p + 1 + for t < re && q < kre { + switch { + case c.colIdx[t] == c.colIdx[q]: + c.vals[t] -= mult * c.vals[q] + if math.IsInf(c.vals[t], 0) || math.IsNaN(c.vals[t]) { + return nil, base.Errf("NewSparseILU: the elimination overflowed at row %d, column %d", i, c.colIdx[t]) + } + t++ + q++ + case c.colIdx[t] < c.colIdx[q]: + t++ + default: + q++ + } + } + } + } + return &SparseILU{ + vals: c.vals, + colIdx: c.colIdx, + rowStart: c.rowStart, + n: n, + }, nil +} + +// atDiagonal returns the diagonal entry of row i, or 0 when absent. +func (c *sparseCSR) atDiagonal(i int) float64 { + for p := c.rowStart[i]; p < c.rowStart[i+1] && c.colIdx[p] <= i; p++ { + if c.colIdx[p] == i { + return c.vals[p] + } + } + return 0 +} + +// Apply solves (L·U)·x = r over the stored pattern: a forward +// substitution with unit diagonal, then a backward substitution with +// the pivots. The result approximates A⁻¹·r to the accuracy the +// dropped fill allows. +func (ilu *SparseILU) Apply(r []float64) []float64 { + dst := make([]float64, ilu.n) + ilu.applyTo(dst, r) + return dst +} + +// applyTo writes (L·U)⁻¹·r into dst, which must be a buffer of the +// factor's dimension distinct from r. Both substitutions run in place: +// the forward pass overwrites each entry only after reading the ones +// below it, and the backward pass reads dst[i] before overwriting it +// with y[i], so an application allocates nothing. +func (ilu *SparseILU) applyTo(dst, r []float64) { + n := ilu.n + copy(dst, r) + for i := range n { + s := dst[i] + for p := ilu.rowStart[i]; p < ilu.rowStart[i+1] && ilu.colIdx[p] < i; p++ { + s -= ilu.vals[p] * dst[ilu.colIdx[p]] + } + dst[i] = s + } + for i := n - 1; i >= 0; i-- { + s := dst[i] + var diag float64 + for p := ilu.rowStart[i]; p < ilu.rowStart[i+1]; p++ { + j := ilu.colIdx[p] + if j == i { + diag = ilu.vals[p] + } else if j > i { + s -= ilu.vals[p] * dst[j] + } + } + if diag == 0 { + // NewSparseILU refuses zero pivots, so this is defensive. + diag = 1 + } + dst[i] = s / diag + } +} + +// iluPrecondition writes the preconditioned vector of r into dst: the +// caller's ILU factor when one was given, otherwise the Jacobi scaling +// by the diagonal. +func iluPrecondition(dst, r []float64, ilu *SparseILU, diag []float64) { + if ilu != nil { + ilu.applyTo(dst, r) + return + } + for i := range r { + dst[i] = r[i] / diag[i] + } +} diff --git a/linalg/sparseilu_test.go b/linalg/sparseilu_test.go new file mode 100644 index 0000000..c56cf74 --- /dev/null +++ b/linalg/sparseilu_test.go @@ -0,0 +1,208 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// triDiagCOO builds a tridiagonal SparseCOO with the given diagonals. +func triDiagCOO(t *testing.T, n int, lo, diag, up float64) *core.SparseCOO { + t.Helper() + idx := make([]int64, 0, 3*n) + vals := make([]float64, 0, 3*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, diag) + if i+1 < n { + add(i, i+1, up) + add(i+1, i, lo) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +// TestSparseILUTridiagonalExact uses the fact that a tridiagonal +// matrix creates no fill under LU: the ILU(0) factors must reproduce +// the complete LU exactly, so applying them to a unit vector answers +// what the dense Solve answers for the same right-hand side. +func TestSparseILUTridiagonalExact(t *testing.T) { + const n = 12 + coo := triDiagCOO(t, n, -1, 4, -2) + ilu, err := NewSparseILU(coo) + if err != nil { + t.Fatalf("NewSparseILU: %v", err) + } + denseVals := make([]float64, n*n) + for i := range n { + denseVals[i*n+i] = 4 + if i+1 < n { + denseVals[i*n+i+1] = -2 + denseVals[(i+1)*n+i] = -1 + } + } + dense := mustFloats(t, denseVals, n, n) + for j := range n { + e := make([]float64, n) + e[j] = 1 + got := ilu.Apply(e) + ref, err := Solve(dense, mustFloats(t, e)) + if err != nil { + t.Fatalf("Solve: %v", err) + } + for i := range n { + if math.Abs(got[i]-ref.FloatAt(i)) > 1e-11 { + t.Fatalf("column %d, row %d: %.14g, want %.14g", j, i, got[i], ref.FloatAt(i)) + } + } + } +} + +// TestSparseILUWiderStencil checks the wider-than-tridiagonal case: +// a five-point stencil with dropped fill cannot reproduce A exactly, +// but the preconditioned residual of the factorisation must still be +// far smaller than the identity preconditioner's. +func TestSparseILUWiderStencil(t *testing.T) { + const n = 20 + idx := make([]int64, 0, 5*n) + vals := make([]float64, 0, 5*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 5) + if i+1 < n { + add(i, i+1, -2) + add(i+1, i, -1) + } + if i+2 < n { + add(i, i+2, 0.3) + add(i+2, i, -0.2) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + ilu, err := NewSparseILU(coo) + if err != nil { + t.Fatalf("NewSparseILU: %v", err) + } + rhs := make([]float64, n) + for i := range n { + rhs[i] = math.Cos(float64(i)) + } + x := ilu.Apply(rhs) + // The residual ‖r − A·x‖ must sit well below ‖r‖: the incomplete + // factorisation captures most of A. + res := 0.0 + rhsNorm := 0.0 + for i := range n { + s := rhs[i] + rowSum := 0.0 + for k := range n { + v := 0.0 + switch { + case k == i: + v = 5 + case k == i+1: + v = -2 + case k == i-1: + v = -1 + case k == i+2: + v = 0.3 + case k == i-2: + v = -0.2 + } + rowSum += v * x[k] + } + d := s - rowSum + res += d * d + rhsNorm += s * s + } + if math.Sqrt(res) > 0.5*math.Sqrt(rhsNorm) { + t.Fatalf("ILU residual %g too large against ‖r‖ %g", math.Sqrt(res), math.Sqrt(rhsNorm)) + } +} + +// TestSparseILUPreconditionedSteps checks the preconditioner's purpose: +// on a convective 100×100 system the ILU-preconditioned BiCGSTAB must +// reach the same tolerance as the Jacobi run, with a residual that +// meets the bound. +func TestSparseILUPreconditionedSteps(t *testing.T) { + const n = 100 + coo := triDiagCOO(t, n, -1+0.5, 4, -1) + bv := make([]float64, n) + for i := range n { + bv[i] = math.Sin(float64(i+1) / float64(n+1)) + } + b := mustFloats(t, bv) + tol := 1e-10 + xJ, err := SpSolveBiCGSTAB(coo, b, tol, 0) + if err != nil { + t.Fatalf("Jacobi run: %v", err) + } + ilu, err := NewSparseILU(coo) + if err != nil { + t.Fatalf("NewSparseILU: %v", err) + } + xI, err := SpSolveBiCGSTAB(coo, b, tol, 0, ilu) + if err != nil { + t.Fatalf("ILU run: %v", err) + } + for i := range n { + if math.Abs(xJ.FloatAt(i)-xI.FloatAt(i)) > 1e-8 { + t.Fatalf("solutions disagree at %d: %.14g vs %.14g", i, + xJ.FloatAt(i), xI.FloatAt(i)) + } + } + // The tridiagonal ILU is the exact LU: the ILU-preconditioned + // system must converge in very few steps. Cap the budget where + // plain Jacobi needs far more. + if _, err := SpSolveBiCGSTAB(coo, b, tol, 3, ilu); err != nil { + t.Fatalf("exact-LU preconditioner should converge in 3 steps: %v", err) + } +} + +func TestSparseILUErrors(t *testing.T) { + if _, err := NewSparseILU(triDiagCOO(t, 3, -1, 0, -1)); err == nil { + t.Fatal("zero pivot: want an error") + } + rect, err := core.NewSparseCOO(mustInts2(t, []int64{0, 0, 0, 1, 0, 2}, 3, 2), + floatsToArray([]float64{1, 2, 3}, []int{3}), []int{2, 3}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseILU(rect); err == nil { + t.Fatal("non-square: want an error") + } +} + +// mustInts2 builds a 2-D int array for sparse constructors. +func mustInts2(t *testing.T, vals []int64, rows, cols int) *core.Array { + t.Helper() + a, err := core.FromInts(vals, rows, cols) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + return a +} diff --git a/linalg/sparselsqr.go b/linalg/sparselsqr.go new file mode 100644 index 0000000..f6d72e1 --- /dev/null +++ b/linalg/sparselsqr.go @@ -0,0 +1,807 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Overdetermined sparse least squares: min ‖A·x − b‖₂ for a sparse A +// with at least as many rows as columns. The dense `LeastSquares` +// factorises the whole matrix, which a large sparse system cannot +// afford and need not: LSQR and LSMR walk the Golub-Kahan +// bidiagonalisation instead, whose every step costs one A·v product +// and one Aᵀ·u product, so the cost tracks the non-zero count. LSQR +// minimises ‖b − A·x‖ over the growing Krylov space; LSMR minimises +// ‖Aᵀ(b − A·x)‖, which keeps the normal-equations residual monotone. +// +// Both solvers carry the stopping-criterion trio the roadmap names for +// this recursion (Paige and Saunders): +// +// - "residual": the estimate of ‖b − A·x‖₂ has fallen against the +// scaled target tol·(‖b‖ + ‖A‖·‖x‖); +// - "normal": the estimate of ‖Aᵀ(b − A·x)‖₂, the residual of the +// normal equations, has fallen against tol·‖A‖·‖b − A·x‖; +// - "condition": the iteration's conditional estimate of cond(A) +// has passed the limit conlim, the augmented estimate the +// recursion maintains from the bidiagonal entries. +// +// The iteration stops when ANY criterion fires, the standard LSQR and +// LSMR semantics, and the returned info records which one fired. When +// several fire on the same step the residual criterion wins, the +// priority the reference implementations apply. A condition stop is a +// stop, not a convergence, and Converged reports the difference. +// +// An underdetermined system (fewer rows than columns) is refused: the +// problem changes shape there, the minimum-norm answer of an +// underdetermined solve is a different contract than the least-squares +// answer this file implements, and the dense surface refuses it the +// same way. +// +// Both entry points are deterministic: every product runs serially +// over the stored entries in index order and no goroutine joins the +// iteration. The per-step allocation is nothing beyond the info's own +// estimate trajectory: the working vectors are allocated once and +// reused, and no step copies a matrix. +// +// Two guards keep the recursion honest where the plain recurrences +// would lie. A bidiagonal entry at or below 64·eps of the recursion's +// own scale, the constant `SolveTruncated` applies to a vanished +// singular value, carries no usable direction: it is treated as +// exactly zero, the current step folds on the clamped values and the +// iteration stops, because a direction the recursion cannot resolve +// would otherwise be amplified into the answer. And a fired criterion +// is verified against the explicitly recomputed residual before the +// answer is returned, so an estimate the drift of the recursion has +// detached from the truth can end the iteration only as an honest +// failure, never as a converged answer. + +// spLeastSquaresTol is the default relative tolerance when the caller +// passes zero, matching the package's iterative solvers. +const spLeastSquaresTol = 1e-10 + +// spLeastSquaresConLim is the default condition limit when the caller +// passes zero: large enough that a healthy system never trips it, low +// enough that a hopeless one stops instead of burning its budget. +const spLeastSquaresConLim = 1e8 + +// spLeastSquaresBreakdown is the breakdown constant, in multiples of +// the machine epsilon relative to the recursion's own scale, below +// which a bidiagonal entry is treated as exactly zero. +const spLeastSquaresBreakdown = 64 + +// The stopping criteria a sparse least-squares iteration can report. +const ( + // LeastSquaresResidual names the residual-norm criterion: the + // estimate of ‖b − A·x‖₂ met the scaled target. + LeastSquaresResidual = "residual" + // LeastSquaresNormal names the normal-equations criterion: the + // estimate of ‖Aᵀ(b − A·x)‖₂ met the scaled target. + LeastSquaresNormal = "normal" + // LeastSquaresCondition names the conditional criterion: the + // estimate of cond(A) passed the limit. + LeastSquaresCondition = "condition" +) + +// LeastSquaresInfo reports what a sparse least-squares iteration +// achieved and which stopping test ended it. +type LeastSquaresInfo struct { + // Iterations is the number of Golub-Kahan steps folded into the + // answer. + Iterations int + // Criterion names the test that ended the iteration: + // LeastSquaresResidual, LeastSquaresNormal or + // LeastSquaresCondition. It is empty when the exact answer x = 0 + // was returned without a step. + Criterion string + // ResidualNorm is the achieved ‖b − A·x‖₂, recomputed explicitly + // from the returned x once the iteration has ended, so it is the + // truth and not the in-loop estimate. + ResidualNorm float64 + // NormalResidual is the achieved ‖Aᵀ(b − A·x)‖₂, recomputed the + // same way. + NormalResidual float64 + // MatrixNorm is the iteration's running estimate of ‖A‖. + MatrixNorm float64 + // Condition is the iteration's running estimate of cond(A). + Condition float64 + // Converged reports whether a residual or normal-equations test + // fired. A condition stop leaves it false: the answer is the + // estimate reached when the condition limit passed, the standard + // LSQR and LSMR semantics, and an honest caller treats it as a + // warning, not a solution. + Converged bool + + // residualEstimates and normalEstimates hold the per-step in-loop + // estimates of ‖b − A·x‖ and ‖Aᵀ(b − A·x)‖ in iteration order, so + // the convergence behaviour stays inspectable. + residualEstimates []float64 + normalEstimates []float64 +} + +// lsqrOperator holds A and its transpose in CSR form, built once per +// solve so every iteration streams two products through contiguous +// slices instead of touching the coordinate form. +type lsqrOperator struct { + m, n int + rowStart []int + colIdx []int + vals []float64 + tStart []int + tIdx []int + tVals []float64 +} + +// newLSQROperator converts the coordinate form through the canonical +// CSR conversion (duplicates sum, explicit zeros drop) and counts the +// transpose in one pass. +func newLSQROperator(a *core.SparseCOO) (*lsqrOperator, error) { + csr, err := CSRFromCOO(a) + if err != nil { + return nil, err + } + t := csr.Transpose() + return &lsqrOperator{ + m: csr.Rows, n: csr.Cols, + rowStart: csr.RowStart, colIdx: csr.ColIdx, vals: csr.Values, + tStart: t.RowStart, tIdx: t.ColIdx, tVals: t.Values, + }, nil +} + +// matVec writes A·x into out, one row at a time in ascending order. +func (op *lsqrOperator) matVec(x, out []float64) { + for i := range op.m { + sum := 0.0 + for p := op.rowStart[i]; p < op.rowStart[i+1]; p++ { + sum += op.vals[p] * x[op.colIdx[p]] + } + out[i] = sum + } +} + +// tMatVec writes Aᵀ·x into out, one column of A at a time in ascending +// order, which keeps the reduction order deterministic. +func (op *lsqrOperator) tMatVec(x, out []float64) { + for j := range op.n { + sum := 0.0 + for p := op.tStart[j]; p < op.tStart[j+1]; p++ { + sum += op.tVals[p] * x[op.tIdx[p]] + } + out[j] = sum + } +} + +// checkSparseLeastSquares validates the inputs both least-squares +// solvers share: a non-empty real 2-D sparse matrix with at least as +// many rows as columns and a real right-hand side of the row count. It +// returns the operator and b as a working vector. +func checkSparseLeastSquares(name string, a *core.SparseCOO, b *core.Array) (*lsqrOperator, []float64, error) { + if a.Values.Dtype() == core.Complex { + return nil, nil, base.Errf("%s: complex sparse matrices are not supported", name) + } + if len(a.Shape) != 2 || a.Shape[0] == 0 || a.Shape[1] == 0 { + return nil, nil, base.Errf("%s: needs a non-empty 2-D sparse matrix, got shape %v", name, a.Shape) + } + m, n := a.Shape[0], a.Shape[1] + if m < n { + return nil, nil, base.Errf("%s: needs an overdetermined system with m ≥ n, got %d×%d; the underdetermined minimum-norm problem is a different contract", name, m, n) + } + op, err := newLSQROperator(a) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + if b.Dtype() == core.Complex { + return nil, nil, base.Errf("%s: complex right-hand sides are not supported", name) + } + if b.NDim() != 1 || b.Len() != m { + return nil, nil, base.Errf("%s: right-hand side must be a vector of length %d, got shape %s", + name, m, base.ShapeText(b.Shape())) + } + return op, vectorF64(b, m), nil +} + +// symOrtho computes the Givens rotation of the pair (a, b) that zeroes +// the second coordinate, in the stable form the reference +// implementations use: the signs are carried instead of cancelled, so +// no intermediate reaches 1/eps. r is the resulting non-negative norm. +func symOrtho(a, b float64) (c, s, r float64) { + switch { + case b == 0: + return sign(a), 0, math.Abs(a) + case a == 0: + return 0, sign(b), math.Abs(b) + case math.Abs(b) > math.Abs(a): + tau := a / b + s = sign(b) / math.Sqrt(1+tau*tau) + c = s * tau + r = b / s + default: + tau := b / a + c = sign(a) / math.Sqrt(1+tau*tau) + s = c * tau + r = a / c + } + return c, s, r +} + +// lsqrExplicitNorms recomputes ‖b − A·x‖ and ‖Aᵀ(b − A·x)‖ directly +// from a candidate x: two products, run once at the end of the +// iteration, so the reported achieved quantities are the truth rather +// than the in-loop estimates. ax and ar are scratch vectors of the +// operator's row and column counts, taken over from the iteration that +// just ended. +func lsqrExplicitNorms(op *lsqrOperator, bv, x, ax, ar []float64) (rNorm, arNorm float64) { + op.matVec(x, ax) + for i := range op.m { + ax[i] = bv[i] - ax[i] + } + op.tMatVec(ax, ar) + return norm2F64(ax), norm2F64(ar) +} + +// verifyLeastSquares checks a fired criterion against the explicitly +// recomputed norms, with the round-off slack the achieved norms can +// never beat, and relabels to the sibling test when the fired one +// fails while the other passes. It reports whether the achieved +// answer satisfies any residual test at all; a condition stop makes no +// accuracy claim and passes unverified. +func verifyLeastSquares(criterion string, rNorm, arNorm, bNorm, matrixNorm, xNorm, tol float64) (string, bool) { + targetR := tol*(bNorm+matrixNorm*xNorm) + 64*base.EpsF*(bNorm+matrixNorm*xNorm) + targetA := tol*matrixNorm*rNorm + 64*base.EpsF*matrixNorm*(bNorm+rNorm+xNorm) + passR := rNorm <= targetR + passA := arNorm <= targetA + switch criterion { + case LeastSquaresResidual: + if passR { + return LeastSquaresResidual, true + } + if passA { + return LeastSquaresNormal, true + } + case LeastSquaresNormal: + if passA { + return LeastSquaresNormal, true + } + if passR { + return LeastSquaresResidual, true + } + default: + return criterion, true + } + return criterion, false +} + +// zeroLeastSquaresAnswer builds the answer both solvers give when the +// iteration needs no step: b is zero, or b is orthogonal to the column +// space of A, and x = 0 is then the exact least-squares solution. +func zeroLeastSquaresAnswer(n int, bNorm, alpha float64) (*core.Array, *LeastSquaresInfo, error) { + x, err := core.Zeros(core.Float, n) + if err != nil { + return nil, nil, err + } + info := &LeastSquaresInfo{ + Criterion: "", + Converged: true, + ResidualNorm: bNorm, + NormalResidual: 0, + MatrixNorm: alpha, + Condition: 0, + } + return x, info, nil +} + +// SpLSQR returns the vector x minimising ‖A·x − b‖₂ over a sparse +// overdetermined A, by the Golub-Kahan bidiagonalisation recursion of +// Paige and Saunders. The stopping criterion trio is the residual-norm +// estimate, the normal-equations residual and the conditional +// estimate; the iteration stops when ANY criterion fires and the +// returned info names it and carries the achieved quantities. +// +// tol ≤ 0 means 1e-10 and feeds both the residual and the +// normal-equations test; maxIter ≤ 0 means 2n steps, the reference +// default of twice the Krylov dimension; conlim ≤ 0 means 1e8, and +// cond(A) passing it stops the iteration without claiming +// convergence. Running out of steps with every tolerance unmet is an +// error naming the residual achieved, with no estimate returned, in +// the style of the package's other solvers. +func SpLSQR(a *core.SparseCOO, b *core.Array, tol float64, maxIter int, conlim float64) (*core.Array, *LeastSquaresInfo, error) { + const name = "SpLSQR" + op, bv, err := checkSparseLeastSquares(name, a, b) + if err != nil { + return nil, nil, err + } + if tol <= 0 { + tol = spLeastSquaresTol + } + if conlim <= 0 { + conlim = spLeastSquaresConLim + } + ctol := 1 / conlim + m, n := op.m, op.n + if maxIter <= 0 { + maxIter = 2 * n + } + + x := make([]float64, n) + bNorm := norm2F64(bv) + if bNorm == 0 { + return zeroLeastSquaresAnswer(n, 0, 0) + } + // β₁u₁ = b, α₁v₁ = Aᵀu₁. A vanishing α₁ means b is orthogonal to + // the column space and x = 0 is the exact answer. + u := append([]float64(nil), bv...) + scale := 1 / bNorm + for i := range m { + u[i] *= scale + } + v := make([]float64, n) + op.tMatVec(u, v) + alpha := norm2F64(v) + if alpha == 0 { + return zeroLeastSquaresAnswer(n, bNorm, 0) + } + scale = 1 / alpha + for j := range n { + v[j] *= scale + } + w := append([]float64(nil), v...) + + // The rotated right-hand side: φ̄ carries the part of b no step has + // consumed yet, ρ̄ the pending bidiagonal entry. The estimates are + // ‖b − A·x‖ = |φ̄| and the normal-equations estimate α·|τ|, τ being + // the rotated off-diagonal, exactly as the reference recursion + // maintains them. + rhobar, phibar := alpha, bNorm + anorm, ddnorm, xxnorm := 0.0, 0.0, 0.0 + xnorm, z, cs2, sn2 := 0.0, 0.0, -1.0, 0.0 + info := &LeastSquaresInfo{} + criterion := "" + exhausted := false + steps := 0 + uBuf := make([]float64, m) + vBuf := make([]float64, n) + // The per-step estimates grow into buffers sized for the iteration + // budget, capped at the default budget so a caller's huge maxIter + // cannot allocate steps the recursion will not take. + estCap := min(maxIter, 2*n) + rEsts := make([]float64, 0, estCap) + arEsts := make([]float64, 0, estCap) + // hScale tracks the recursion's own magnitude for the breakdown + // floor: a bidiagonal entry at round-off of the scale already seen + // carries no direction. + hScale := max(alpha, bNorm) + + // hitBudget tells the three exits apart: a loop that ended by its + // own iteration count fired no criterion and broke down nowhere, + // which is the budget's fault and not an estimate's drift. + hitBudget := true + for itn := 1; itn <= maxIter; itn++ { + // βu = A·v − αu, αv = Aᵀu − βv: one step of the bidiagonalisation. + op.matVec(v, uBuf) + for i := range m { + uBuf[i] -= alpha * u[i] + } + if !vecFinite(uBuf) { + return nil, nil, base.Errf("%s: non-finite residual direction at step %d", name, itn) + } + beta := norm2F64(uBuf) + hScale = max(hScale, beta) + floor := spLeastSquaresBreakdown * base.EpsF * hScale + if beta <= floor { + // The u-direction is gone: fold the pending step as the + // exact-arithmetic completion, with β = 0, and stop. + beta = 0 + exhausted = true + } else { + scale = 1 / beta + for i := range m { + u[i] = uBuf[i] * scale + } + anorm = math.Sqrt(anorm*anorm + alpha*alpha + beta*beta) + op.tMatVec(u, vBuf) + for j := range n { + vBuf[j] -= beta * v[j] + } + if !vecFinite(vBuf) { + return nil, nil, base.Errf("%s: non-finite normal direction at step %d", name, itn) + } + alpha = norm2F64(vBuf) + hScale = max(hScale, alpha) + if alpha <= floor { + // The v-direction is gone: α is the exact zero the + // recursion would have produced, v the zero vector, and + // the iteration stops after this fold. + alpha = 0 + exhausted = true + clear(v) + } else if alpha > 0 { + scale = 1 / alpha + for j := range n { + v[j] = vBuf[j] * scale + } + } + } + // The plane rotation that eliminates the subdiagonal β. With a + // clamped β the rotation reads zero off the subdiagonal, which + // is the exact-arithmetic completion of the recursion; a stale + // α only feeds ᾱ, which the stop leaves unused. + cs, sn, rho := symOrtho(rhobar, beta) + if rho == 0 { + exhausted = true + hitBudget = false + break + } + theta := sn * alpha + rhobar = -cs * alpha + phi := cs * phibar + phibar = sn * phibar + tau := sn * phi + t1 := phi / rho + t2 := -theta / rho + // ddnorm feeds the condition estimate: the accumulated squared + // lengths of the consumed directions, read before w moves. + nw := norm2F64(w) + ddnorm += nw * nw / (rho * rho) + for j := range n { + x[j] += t1 * w[j] + w[j] = v[j] + t2*w[j] + } + steps = itn + // ‖x‖ through a running rotation on the triangular solve, the + // reference recursion's incremental form. cs2 starts at −1 and + // every update carries a nonzero gambar, so the divisions never + // see zero. + delta := sn2 * rho + gambar := -cs2 * rho + rhs := phi - delta*z + zbar := rhs / gambar + xnorm = math.Sqrt(xxnorm + zbar*zbar) + gamma := math.Hypot(gambar, theta) + cs2 = gambar / gamma + sn2 = theta / gamma + z = rhs / gamma + xxnorm += z * z + + rEst := math.Abs(phibar) + arEst := alpha * math.Abs(tau) + rEsts = append(rEsts, rEst) + arEsts = append(arEsts, arEst) + acond := anorm * math.Sqrt(ddnorm) + test1 := rEst / bNorm + test2 := arEst / (anorm*rEst + base.EpsF) + test3 := 1 / (acond + base.EpsF) + target := tol * (1 + anorm*xnorm/bNorm) + switch { + case test1 <= target: + criterion = LeastSquaresResidual + case test2 <= tol: + criterion = LeastSquaresNormal + case test3 <= ctol: + criterion = LeastSquaresCondition + } + if criterion != "" || exhausted { + hitBudget = false + break + } + } + + rNorm, arNorm := lsqrExplicitNorms(op, bv, x, uBuf, vBuf) + criterion = settleLeastSquares(criterion, exhausted, rNorm, arNorm, x, bNorm, anorm, tol) + if criterion == "" { + if exhausted { + return nil, nil, base.Errf("%s: no convergence: the recursion folded at step %d, residual %.3g (tolerance %.3g)", + name, steps, rNorm, tol) + } + if hitBudget { + return nil, nil, base.Errf("%s: no convergence in %d steps, residual %.3g (tolerance %.3g)", + name, maxIter, rNorm, tol) + } + return nil, nil, base.Errf("%s: a stopping criterion fired at step %d but its estimate drifted past the recomputed norms (residual %.3g, tolerance %.3g)", + name, steps, rNorm, tol) + } + info.Iterations = steps + info.Criterion = criterion + info.Converged = criterion != LeastSquaresCondition + info.ResidualNorm = rNorm + info.NormalResidual = arNorm + info.MatrixNorm = anorm + info.Condition = anorm * math.Sqrt(ddnorm) + info.residualEstimates = rEsts + info.normalEstimates = arEsts + return floatsToArray(x, []int{n}), info, nil +} + +// SpLSMR returns the vector x minimising ‖Aᵀ(b − A·x)‖₂ over a sparse +// overdetermined A, by Fong and Saunders' LSMR, the Golub-Kahan +// bidiagonalisation with the recurrences reorganised so that both +// ‖b − A·x‖ and ‖Aᵀ(b − A·x)‖ carry exact-norm estimates. The stopping +// criterion trio and the reporting contract are SpLSQR's; the +// normal-equations residual, LSMR's own minimisation target, moves +// monotonically. LSMR's answer on a consistent rank-deficient system +// is the minimum-norm one, as LSQR's is. +// +// maxIter ≤ 0 means 2n steps, twice the Krylov dimension, the same +// default SpLSQR applies; the breakdown floor ends the recursion +// before a step can fold a direction the space cannot resolve. +func SpLSMR(a *core.SparseCOO, b *core.Array, tol float64, maxIter int, conlim float64) (*core.Array, *LeastSquaresInfo, error) { + const name = "SpLSMR" + op, bv, err := checkSparseLeastSquares(name, a, b) + if err != nil { + return nil, nil, err + } + if tol <= 0 { + tol = spLeastSquaresTol + } + if conlim <= 0 { + conlim = spLeastSquaresConLim + } + ctol := 1 / conlim + m, n := op.m, op.n + if maxIter <= 0 { + maxIter = 2 * n + } + + x := make([]float64, n) + bNorm := norm2F64(bv) + if bNorm == 0 { + return zeroLeastSquaresAnswer(n, 0, 0) + } + u := append([]float64(nil), bv...) + scale := 1 / bNorm + for i := range m { + u[i] *= scale + } + v := make([]float64, n) + op.tMatVec(u, v) + alpha := norm2F64(v) + if alpha == 0 { + return zeroLeastSquaresAnswer(n, bNorm, 0) + } + scale = 1 / alpha + for j := range n { + v[j] *= scale + } + + // The two rotation pairs of LSMR: one turns the bidiagonal matrix + // upper triangular (ρ, c, s and the ᾱ carry), the other turns its + // transpose back (c̄, s̄, ρ̄), which is what makes ζ̄ the exact + // normal-equations norm the solver minimises. + zetabar := alpha * bNorm + alphabar := alpha + rho, rhobar := 1.0, 1.0 + cbar, sbar := 1.0, 0.0 + h := append([]float64(nil), v...) + hbar := make([]float64, n) + + // The residual-norm bookkeeping: ζ, βd and β̆ walk the rotated + // right-hand side so ‖r‖ = √((βd − τd)² + β̆²) is exact in exact + // arithmetic. + betadd, betad := bNorm, 0.0 + rhodold, tautildeold, thetatilde, zeta := 1.0, 0.0, 0.0, 0.0 + + // The ‖A‖ and cond(A) estimates: the squared bidiagonal entries and + // the extreme rotated diagonals. + normA2 := alpha * alpha + maxrbar, minrbar := 0.0, 1e100 + + info := &LeastSquaresInfo{} + criterion := "" + exhausted := false + steps := 0 + uBuf := make([]float64, m) + vBuf := make([]float64, n) + // The estimate buffers are sized as SpLSQR's, for the same reason. + estCap := min(maxIter, 2*n) + rEsts := make([]float64, 0, estCap) + arEsts := make([]float64, 0, estCap) + normA := math.Sqrt(normA2) + condA := 1.0 + normr := bNorm + hScale := max(alpha, bNorm) + // The rotation's cosine and sine live across the loop body; ρ + // itself carries across iterations, which is what rhoold reads. + var c, s float64 + + // hitBudget tells the three exits apart, exactly as SpLSQR's: a + // loop that ended by its own iteration count fired no criterion + // and broke down nowhere, which is the budget's fault. + hitBudget := true + for itn := 1; itn <= maxIter; itn++ { + // βu = A·v − αu, αv = Aᵀu − βv. + op.matVec(v, uBuf) + for i := range m { + uBuf[i] -= alpha * u[i] + } + if !vecFinite(uBuf) { + return nil, nil, base.Errf("%s: non-finite residual direction at step %d", name, itn) + } + beta := norm2F64(uBuf) + hScale = max(hScale, beta) + floor := spLeastSquaresBreakdown * base.EpsF * hScale + if beta <= floor { + beta = 0 + exhausted = true + } else { + scale = 1 / beta + for i := range m { + u[i] = uBuf[i] * scale + } + op.tMatVec(u, vBuf) + for j := range n { + vBuf[j] -= beta * v[j] + } + if !vecFinite(vBuf) { + return nil, nil, base.Errf("%s: non-finite normal direction at step %d", name, itn) + } + alpha = norm2F64(vBuf) + hScale = max(hScale, alpha) + if alpha <= floor { + // The v-direction is gone: α is treated as the exact + // zero the recursion would have produced, v as the zero + // vector, and the iteration stops after this fold. + alpha = 0 + exhausted = true + clear(v) + } else if alpha > 0 { + scale = 1 / alpha + for j := range n { + v[j] = vBuf[j] * scale + } + } + } + // First rotation pair: the damping fold first, as the reference + // applies it: (ᾱ, 0) gives the sign chat and the magnitude α̂, + // then (α̂, β) turns to ρ. c and s live in the loop, but ρ is + // declared outside it and carries: rhoold needs the previous + // step's value, and a `:=` here would shadow it. + chat := sign(alphabar) + rhoold := rho + c, s, rho = symOrtho(math.Abs(alphabar), beta) + if rho == 0 { + exhausted = true + hitBudget = false + break + } + thetanew := s * alpha + alphabar = c * alpha + // Second rotation pair: (c̄ρ, θ) to ρ̄, which moves the + // normal-equations estimate ζ̄. + rhobarold := rhobar + zetaold := zeta + thetabar := sbar * rho + rhotemp := cbar * rho + cbar, sbar, rhobar = symOrtho(cbar*rho, thetanew) + zeta = cbar * zetabar + zetabar = -sbar * zetabar + + // The direction recurrences and the answer update. + coef := thetabar * rho / (rhoold * rhobarold) + for j := range n { + hbar[j] = h[j] - coef*hbar[j] + } + xk := zeta / (rho * rhobar) + for j := range n { + x[j] += xk * hbar[j] + } + ht := -thetanew / rho + for j := range n { + h[j] = v[j] + ht*h[j] + } + steps = itn + + // The exact-form residual estimate: the pending right-hand side + // entries ride the rotations, chat carrying the sign the + // damping fold produced. + betaacute := chat * betadd + betahat := c * betaacute + betadd = -s * betaacute + thetatildeold := thetatilde + ctildeold, stildeold, rhotildeold := symOrtho(rhodold, thetabar) + thetatilde = stildeold * rhobar + rhodold = ctildeold * rhobar + betad = -stildeold*betad + ctildeold*betahat + tautildeold = (zetaold - thetatildeold*tautildeold) / rhotildeold + taud := (zeta - thetatilde*tautildeold) / rhodold + normr = math.Sqrt((betad-taud)*(betad-taud) + betadd*betadd) + + // The ‖A‖ and cond(A) estimates. + normA2 += beta * beta + normA = math.Sqrt(normA2) + normA2 += alpha * alpha + maxrbar = max(maxrbar, rhobarold) + if itn > 1 { + minrbar = min(minrbar, rhobarold) + } + condA = math.Inf(1) + if d := min(minrbar, rhotemp); d > 0 { + condA = max(maxrbar, rhotemp) / d + } + + normar := math.Abs(zetabar) + normx := norm2F64(x) + rEsts = append(rEsts, normr) + arEsts = append(arEsts, normar) + test1 := normr / bNorm + test2 := math.Inf(1) + if p := normA * normr; p != 0 { + test2 = normar / p + } + test3 := 1 / (condA + base.EpsF) + target := tol * (1 + normA*normx/bNorm) + switch { + case test1 <= target: + criterion = LeastSquaresResidual + case test2 <= tol: + criterion = LeastSquaresNormal + case test3 <= ctol: + criterion = LeastSquaresCondition + } + if criterion != "" || exhausted { + hitBudget = false + break + } + } + + rNorm, arNorm := lsqrExplicitNorms(op, bv, x, uBuf, vBuf) + criterion = settleLeastSquares(criterion, exhausted, rNorm, arNorm, x, bNorm, normA, tol) + if criterion == "" { + if exhausted { + return nil, nil, base.Errf("%s: no convergence: the recursion folded at step %d, residual %.3g (tolerance %.3g)", + name, steps, rNorm, tol) + } + if hitBudget { + return nil, nil, base.Errf("%s: no convergence in %d steps, residual %.3g (tolerance %.3g)", + name, maxIter, rNorm, tol) + } + return nil, nil, base.Errf("%s: a stopping criterion fired at step %d but its estimate drifted past the recomputed norms (residual %.3g, tolerance %.3g)", + name, steps, rNorm, tol) + } + info.Iterations = steps + info.Criterion = criterion + info.Converged = criterion != LeastSquaresCondition + info.ResidualNorm = rNorm + info.NormalResidual = arNorm + info.MatrixNorm = normA + info.Condition = condA + info.residualEstimates = rEsts + info.normalEstimates = arEsts + return floatsToArray(x, []int{n}), info, nil +} + +// finishLeastSquares recomputes the explicit norms of a candidate x +// and hands the verdict to settleLeastSquares, the entry for a caller +// that holds only the operator and the vectors. The solvers call +// settleLeastSquares directly with the norms they already hold, so a +// solve never runs those two products twice. +func finishLeastSquares(name string, criterion string, exhausted bool, steps int, op *lsqrOperator, bv, x []float64, bNorm, matrixNorm, tol float64) string { + ax := make([]float64, op.m) + ar := make([]float64, op.n) + rNorm, arNorm := lsqrExplicitNorms(op, bv, x, ax, ar) + return settleLeastSquares(criterion, exhausted, rNorm, arNorm, x, bNorm, matrixNorm, tol) +} + +// settleLeastSquares settles the criterion against the truth. A fired +// criterion is verified against the explicitly recomputed residual, +// with the round-off slack, so a drifted estimate cannot dress a +// broken answer up as converged; a collapse that fired no test is +// judged by its explicit norms the same way. It returns the criterion +// to report, or the empty string when the honest answer is an error. +func settleLeastSquares(criterion string, exhausted bool, rNorm, arNorm float64, x []float64, bNorm, matrixNorm, tol float64) string { + if criterion == "" && exhausted { + // The recursion clamped to a stop without a test firing: judge + // the folded answer by its explicit norms. + criterion = LeastSquaresResidual + } + if criterion == "" { + return "" + } + final, ok := verifyLeastSquares(criterion, rNorm, arNorm, bNorm, matrixNorm, norm2F64(x), tol) + if !ok { + return "" + } + return final +} diff --git a/linalg/sparselsqr_test.go b/linalg/sparselsqr_test.go new file mode 100644 index 0000000..9329f9b --- /dev/null +++ b/linalg/sparselsqr_test.go @@ -0,0 +1,468 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// overdeterminedLSFixture builds a deterministic sparse overdetermined +// system: thirty rows, twelve columns, three stored entries per row, +// a right-hand side taken from a known x through the matrix plus a +// small inconsistent part, and its dense twin for the reference +// solvers. +func overdeterminedLSFixture(t *testing.T) (coo *core.SparseCOO, b *core.Array, dense *core.Array) { + t.Helper() + const m, n = 30, 12 + g := core.NewGenerator(7) + idx := make([]int64, 0, 3*m+2) + vals := make([]float64, 0, 3*m+2) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range m { + add(i, i%n, 1+float64(g.Next()%100)/200) + add(i, (i*7+3)%n, -1+float64(g.Next()%100)/100) + add(i, (i*13+5)%n, float64(g.Next()%100)/100) + } + add(0, 0, 3) + add(1, 2, 2) + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err = core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{m, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + xTrue := make([]float64, n) + for i := range n { + xTrue[i] = math.Sin(0.7*float64(i)) + float64(i%5)*0.3 - 0.6 + } + b, err = csr.MatVec(floatsToArray(xTrue, []int{n})) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + for i := range m { + b.RawFloats()[i] += 0.01 * math.Cos(float64(i)) + } + denseVals := make([]float64, m*n) + for i := range len(vals) { + denseVals[int(idx[2*i])*n+int(idx[2*i+1])] += vals[i] + } + return coo, b, floatsToArray(denseVals, []int{m, n}) +} + +// lsAchievedNorms recomputes ‖b − A·x‖ and ‖Aᵀ(b − A·x)‖ through the +// public sparse surface, independently of the solver's own operator. +func lsAchievedNorms(t *testing.T, coo *core.SparseCOO, b, x *core.Array) (rNorm, arNorm float64) { + t.Helper() + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + ax, err := csr.MatVec(x) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + r := core.New(core.Float, b.Len()) + for i := range b.Len() { + r.RawFloats()[i] = b.FloatAt(i) - ax.FloatAt(i) + } + rNorm = norm2F64(vectorF64(r, b.Len())) + at, err := csr.Transpose().MatVec(r) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + arNorm = norm2F64(vectorF64(at, at.Len())) + return rNorm, arNorm +} + +func TestSpLSQRMatchesDenseLeastSquares(t *testing.T) { + coo, b, dense := overdeterminedLSFixture(t) + ref, err := LeastSquares(dense, b) + if err != nil { + t.Fatalf("LeastSquares: %v", err) + } + x, info, err := SpLSQR(coo, b, 1e-12, 0, 0) + if err != nil { + t.Fatalf("SpLSQR: %v", err) + } + worst, refMax := 0.0, 0.0 + for i := range ref.Len() { + if d := math.Abs(x.FloatAt(i) - ref.FloatAt(i)); d > worst { + worst = d + } + if v := math.Abs(ref.FloatAt(i)); v > refMax { + refMax = v + } + } + if worst > 1e-9*refMax { + t.Fatalf("LSQR disagrees with the dense solve by %.3g (scale %.3g)", worst, refMax) + } + if !info.Converged { + t.Fatalf("LSQR stopped without convergence: %+v", info) + } + if info.Criterion != LeastSquaresResidual && info.Criterion != LeastSquaresNormal { + t.Fatalf("criterion %q is not a residual test", info.Criterion) + } + if info.Iterations > 2*12 { + t.Fatalf("LSQR took %d steps for a 12-column Krylov space", info.Iterations) + } + // The achieved quantities must be the explicit truth, not the + // in-loop estimates. + rNorm, arNorm := lsAchievedNorms(t, coo, b, x) + if math.Abs(rNorm-info.ResidualNorm) > 1e-9*info.ResidualNorm { + t.Fatalf("reported residual %.3g does not match the explicit %.3g", info.ResidualNorm, rNorm) + } + if math.Abs(arNorm-info.NormalResidual) > 1e-9*info.NormalResidual { + t.Fatalf("reported normal residual %.3g does not match the explicit %.3g", info.NormalResidual, arNorm) + } +} + +func TestSpLSMRMatchesDenseLeastSquares(t *testing.T) { + coo, b, dense := overdeterminedLSFixture(t) + ref, err := LeastSquares(dense, b) + if err != nil { + t.Fatalf("LeastSquares: %v", err) + } + x, info, err := SpLSMR(coo, b, 1e-12, 0, 0) + if err != nil { + t.Fatalf("SpLSMR: %v", err) + } + worst, refMax := 0.0, 0.0 + for i := range ref.Len() { + if d := math.Abs(x.FloatAt(i) - ref.FloatAt(i)); d > worst { + worst = d + } + if v := math.Abs(ref.FloatAt(i)); v > refMax { + refMax = v + } + } + if worst > 1e-9*refMax { + t.Fatalf("LSMR disagrees with the dense solve by %.3g (scale %.3g)", worst, refMax) + } + if !info.Converged || info.Criterion == "" { + t.Fatalf("LSMR stopped without convergence: %+v", info) + } + // LSMR's own minimisation target moves monotonically. + for i := 1; i < len(info.normalEstimates); i++ { + if info.normalEstimates[i] > info.normalEstimates[i-1] { + t.Fatalf("LSMR normal-equations estimate rose at step %d: %.6g after %.6g", + i, info.normalEstimates[i], info.normalEstimates[i-1]) + } + } +} + +func TestSpLSQREstimateTrajectoriesMonotone(t *testing.T) { + coo, b, _ := overdeterminedLSFixture(t) + // LSQR's residual estimate is |φ̄|, shrank every step by |sn| ≤ 1; + // LSMR's normal estimate is |ζ̄|, shrank by |s̄| ≤ 1. Both are + // monotone by construction and must stay so in float. + _, lsInfo, err := SpLSQR(coo, b, 1e-12, 0, 0) + if err != nil { + t.Fatalf("SpLSQR: %v", err) + } + for i := 1; i < len(lsInfo.residualEstimates); i++ { + if lsInfo.residualEstimates[i] > lsInfo.residualEstimates[i-1] { + t.Fatalf("LSQR residual estimate rose at step %d: %.6g after %.6g", + i, lsInfo.residualEstimates[i], lsInfo.residualEstimates[i-1]) + } + } + _, lsmrInfo, err := SpLSMR(coo, b, 1e-12, 0, 0) + if err != nil { + t.Fatalf("SpLSMR: %v", err) + } + for i := 1; i < len(lsmrInfo.normalEstimates); i++ { + if lsmrInfo.normalEstimates[i] > lsmrInfo.normalEstimates[i-1] { + t.Fatalf("LSMR normal estimate rose at step %d: %.6g after %.6g", + i, lsmrInfo.normalEstimates[i], lsmrInfo.normalEstimates[i-1]) + } + } +} + +// TestSpLeastSquaresMinimumNorm pins the hand-checkable rank-deficient +// consistent system A = [[1,0,1],[0,1,1],[1,1,2]], b = (1,2,3): the +// solution family is (1−t, 2−t, t) and the minimum-norm member is +// (0, 1, 1), which both solvers must answer from a zero start. +func TestSpLeastSquaresMinimumNorm(t *testing.T) { + indices, err := core.FromInts([]int64{0, 0, 0, 2, 1, 1, 1, 2, 2, 0, 2, 1, 2, 2}, 7, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray([]float64{1, 1, 1, 1, 1, 1, 2}, []int{7}), []int{3, 3}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + b, err := csr.MatVec(floatsToArray([]float64{1, 2, 0}, []int{3})) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + want := []float64{0, 1, 1} + for name, run := range map[string]func(*core.SparseCOO, *core.Array) (*core.Array, *LeastSquaresInfo, error){ + "LSQR": func(a *core.SparseCOO, bb *core.Array) (*core.Array, *LeastSquaresInfo, error) { + return SpLSQR(a, bb, 1e-13, 0, 0) + }, + "LSMR": func(a *core.SparseCOO, bb *core.Array) (*core.Array, *LeastSquaresInfo, error) { + return SpLSMR(a, bb, 1e-13, 0, 0) + }, + } { + x, info, err := run(coo, b) + if err != nil { + t.Fatalf("%s: %v", name, err) + } + for i := range want { + if math.Abs(x.FloatAt(i)-want[i]) > 1e-10 { + t.Fatalf("%s: x[%d] = %.12g, want the minimum-norm %.12g", name, i, x.FloatAt(i), want[i]) + } + } + if !info.Converged { + t.Fatalf("%s: not converged: %+v", name, info) + } + } + // The same answer must agree with the SVD's minimum-norm solution. + dense := floatsToArray([]float64{1, 0, 1, 0, 1, 1, 1, 1, 2}, []int{3, 3}) + pinv, err := Pinverse(dense, 0) + if err != nil { + t.Fatalf("Pinverse: %v", err) + } + ref := core.New(core.Float, 3) + for i := range 3 { + s := 0.0 + for j := range 3 { + s += pinv.FloatAt(i*3+j) * b.FloatAt(j) + } + ref.RawFloats()[i] = s + } + for i := range 3 { + if math.Abs(ref.FloatAt(i)-want[i]) > 1e-12 { + t.Fatalf("Pinverse reference %.12g disagrees with the hand solution %.12g", ref.FloatAt(i), want[i]) + } + } +} + +// TestSpLeastSquaresIllConditioned pins convergence with the criterion +// recorded on a system whose diagonal decays by four orders: the +// conditional estimate stays under the limit and a residual test +// ends the iteration. +func TestSpLeastSquaresIllConditioned(t *testing.T) { + const n = 10 + idx := make([]int64, 0, 3*n) + vals := make([]float64, 0, 3*n) + for i := range n { + idx = append(idx, int64(i), int64(i)) + vals = append(vals, math.Pow(10, -0.45*float64(i))) + if i+1 < n { + idx = append(idx, int64(i), int64(i+1), int64(i+1), int64(i)) + vals = append(vals, 1e-7, 1e-7) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + b, err := csr.MatVec(floatsToArray([]float64{1, 1, 1, 1, 1, 1, 1, 1, 1, 1}, []int{n})) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + for name, run := range map[string]func(*core.SparseCOO, *core.Array) (*core.Array, *LeastSquaresInfo, error){ + "LSQR": func(a *core.SparseCOO, bb *core.Array) (*core.Array, *LeastSquaresInfo, error) { + return SpLSQR(a, bb, 1e-9, 0, 0) + }, + "LSMR": func(a *core.SparseCOO, bb *core.Array) (*core.Array, *LeastSquaresInfo, error) { + return SpLSMR(a, bb, 1e-9, 0, 0) + }, + } { + x, info, err := run(coo, b) + if err != nil { + t.Fatalf("%s: %v", name, err) + } + if !info.Converged || info.Criterion == "" { + t.Fatalf("%s: ill-conditioned system ended without a recorded criterion: %+v", name, info) + } + if info.Condition <= 0 || info.MatrixNorm <= 0 { + t.Fatalf("%s: estimates not reported: %+v", name, info) + } + rNorm, _ := lsAchievedNorms(t, coo, b, x) + if rNorm > 1e-8 { + t.Fatalf("%s: achieved residual %.3g is too coarse", name, rNorm) + } + } +} + +func TestSpLeastSquaresErrors(t *testing.T) { + coo, b, _ := overdeterminedLSFixture(t) + // Underdetermined systems are refused, in the dense surface's own + // words. + small, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0, 0, 1}, 2, 2), + floatsToArray([]float64{1, 1}, []int{2}), + []int{2, 3}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + for name, run := range map[string]func(*core.SparseCOO, *core.Array) (*core.Array, *LeastSquaresInfo, error){ + "LSQR": func(a *core.SparseCOO, bb *core.Array) (*core.Array, *LeastSquaresInfo, error) { + return SpLSQR(a, bb, 0, 0, 0) + }, + "LSMR": func(a *core.SparseCOO, bb *core.Array) (*core.Array, *LeastSquaresInfo, error) { + return SpLSMR(a, bb, 0, 0, 0) + }, + } { + if _, _, err := run(small, mustFloats(t, []float64{1, 2}, 2)); err == nil || !strings.Contains(err.Error(), "overdetermined") { + t.Fatalf("%s: underdetermined system accepted: %v", name, err) + } + } + // Complex inputs. + complexValues, err := core.FromComplexes([]complex128{1}, 1) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + complexCOO, err := core.NewSparseCOO(mustInts(t, []int64{0, 0}, 1, 2), complexValues, []int{1, 1}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, _, err := SpLSQR(complexCOO, mustFloats(t, []float64{1}, 1), 0, 0, 0); err == nil { + t.Fatal("a complex matrix was accepted") + } + if _, _, err := SpLSQR(coo, mustFloats(t, []float64{1, 2}, 2), 0, 0, 0); err == nil { + t.Fatal("a short right-hand side was accepted") + } + if _, _, err := SpLSQR(coo, core.New(core.Float, 3, 3), 0, 0, 0); err == nil { + t.Fatal("a rank-2 right-hand side was accepted") + } + // A budget that runs out with every tolerance unmet is the budget's + // own fault: the error names the steps, not a drift no estimate + // committed. A tolerance of 1e-14 with two steps fires no criterion + // (the residual is still 1.59), and neither does 1e-300 in one. + if _, _, err := SpLSQR(coo, b, 1e-14, 2, 0); err == nil || + !strings.Contains(err.Error(), "no convergence in 2 steps") || + strings.Contains(err.Error(), "drifted") { + t.Fatalf("exhausted budget misreported: %v", err) + } + if _, _, err := SpLSQR(coo, b, 1e-300, 1, 0); err == nil || + !strings.Contains(err.Error(), "no convergence in 1 steps") || + strings.Contains(err.Error(), "drifted") { + t.Fatalf("exhausted budget misreported: %v", err) + } + if _, _, err := SpLSMR(coo, b, 1e-300, 1, 0); err == nil || + !strings.Contains(err.Error(), "no convergence in 1 steps") || + strings.Contains(err.Error(), "drifted") { + t.Fatalf("exhausted budget misreported: %v", err) + } + // A zero right-hand side answers the exact zero without a step. + zero := floatsToArray(make([]float64, b.Len()), []int{b.Len()}) + x, info, err := SpLSQR(coo, zero, 0, 0, 0) + if err != nil { + t.Fatalf("zero right-hand side: %v", err) + } + for i := range x.Len() { + if x.FloatAt(i) != 0 { + t.Fatalf("zero right-hand side answered %.3g", x.FloatAt(i)) + } + } + if info.Criterion != "" || !info.Converged || info.Iterations != 0 { + t.Fatalf("zero right-hand side info: %+v", info) + } + // A non-empty matrix with no column-space component of b answers + // the exact zero as well: both columns lie along (1,0) and + // b = (0,1) is orthogonal to the column space. + null, err := core.NewSparseCOO(mustInts(t, []int64{0, 0, 0, 1}, 2, 2), floatsToArray([]float64{1, 2}, []int{2}), []int{2, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + ortho := mustFloats(t, []float64{0, 1}, 2) + x, info, err = SpLSQR(null, ortho, 0, 0, 0) + if err != nil { + t.Fatalf("orthogonal right-hand side: %v", err) + } + for i := range x.Len() { + if x.FloatAt(i) != 0 { + t.Fatalf("orthogonal right-hand side answered %.3g", x.FloatAt(i)) + } + } + if info.ResidualNorm <= 0 { + t.Fatalf("the achieved residual of an orthogonal system vanished: %+v", info) + } +} + +// TestFinishLeastSquaresDriftRefusal pins the drift guard at its own +// gate: a fired criterion whose recomputed norms do not carry it is +// refused with the empty string, while the same answer at a tolerance +// it genuinely meets is carried. The end-to-end budget pins live in +// TestSpLeastSquaresErrors; the drift itself needs the estimate and +// the truth to disagree, which is settled here without a recursion. +func TestFinishLeastSquaresDriftRefusal(t *testing.T) { + coo, err := core.NewSparseCOO(mustInts(t, []int64{0, 0}, 1, 2), mustFloats(t, []float64{1}, 1), []int{1, 1}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + op, err := newLSQROperator(coo) + if err != nil { + t.Fatalf("newLSQROperator: %v", err) + } + // x = 0.5 against A = [1], b = [1] leaves both the residual and + // the normal residual at 0.5 with ‖b‖ = 1. + if got := finishLeastSquares("SpLSQR", LeastSquaresResidual, false, 3, op, []float64{1}, []float64{0.5}, 1, 1, 1e-14); got != "" { + t.Fatalf("a drifted criterion carried as %q", got) + } + if got := finishLeastSquares("SpLSQR", LeastSquaresResidual, false, 3, op, []float64{1}, []float64{0.5}, 1, 1, 0.4); got != LeastSquaresResidual { + t.Fatalf("an honestly met criterion refused: %q", got) + } +} + +// TestSpLSQREstimateSeriesIdentities pins what each estimate series +// carries, not merely that it shrinks. LSQR's residual estimate tracks +// ‖b − A·x‖ and its normal-equations estimate tracks ‖Aᵀ(b − A·x)‖, so +// on a converged fit the first ends beside the explicitly recomputed +// info.ResidualNorm while the second sits orders of magnitude below it. +// Filling the residual series with the normal estimate, or the reverse, +// keeps both series non-increasing and fails here. +func TestSpLSQREstimateSeriesIdentities(t *testing.T) { + coo, b, _ := overdeterminedLSFixture(t) + _, info, err := SpLSQR(coo, b, 1e-12, 0, 0) + if err != nil { + t.Fatalf("SpLSQR: %v", err) + } + re, ne := info.residualEstimates, info.normalEstimates + if len(re) != info.Iterations || len(ne) != info.Iterations { + t.Fatalf("estimate series of %d and %d against %d recorded steps", len(re), len(ne), info.Iterations) + } + if len(re) == 0 { + t.Fatal("no estimates recorded") + } + lastRe, lastNe := re[len(re)-1], ne[len(ne)-1] + if dev := math.Abs(lastRe - info.ResidualNorm); dev > 0.5*info.ResidualNorm { + t.Fatalf("the residual estimate ends at %.6g against the explicit residual %.6g, want the two on the same scale", + lastRe, info.ResidualNorm) + } + if lastNe > 1e-3*info.ResidualNorm { + t.Fatalf("the normal-equations estimate ends at %.6g against a residual of %.6g, want the normal residual's own, much smaller scale", + lastNe, info.ResidualNorm) + } + if lastRe == lastNe { + t.Fatalf("both series end at %.6g, want the residual and the normal-equations estimates", lastRe) + } +} diff --git a/linalg/sparselu.go b/linalg/sparselu.go new file mode 100644 index 0000000..2761b67 --- /dev/null +++ b/linalg/sparselu.go @@ -0,0 +1,352 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "slices" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The sparse LU factorisation with partial pivoting: the left-looking +// column elimination of Gilbert and Peierls. Column k is gathered +// into a dense working vector, the columns that update it are found +// by a depth-first search over the factor built so far, the pivot row +// is the largest working entry at or below the diagonal, and a pivot +// swap relabels the stored factor's rows in place. L is unit lower +// triangular stored by columns, U is upper triangular stored by rows; +// the asymmetry is what lets a row swap move whole U rows while L's +// row labels travel with their values. + +// SparseLU carries the factorisation of a square matrix: P·A = L·U +// with a row permutation only, the columns eliminated in stored +// order. One factorisation solves any number of right hand sides. +type SparseLU struct { + // piv[k] is the original index of the row eliminated k-th. + piv []int + n int + // L by columns, unit diagonal: colRows[j] holds the rows i > j of + // column j with colVals[j] beside it, sorted ascending after the + // elimination. + colRows [][]int + colVals [][]float64 + // U by rows: rowCols[k] holds the columns j > k of row k with + // rowVals[k] beside it, diag[k] the pivot itself. + rowCols [][]int + rowVals [][]float64 + diag []float64 + nnz int +} + +// NewSparseLU factors a square matrix with partial pivoting: the +// pivot row of column k is the largest working entry at or below the +// diagonal, ties resolved by the order the working positions were +// gathered in, which the input alone fixes, so the factorisation is a +// pure function of the input. Complex, rectangular +// and non-finite inputs are refused; a zero pivot column means the +// matrix is singular and is reported. +func NewSparseLU(a *core.SparseCOO) (*SparseLU, error) { + const name = "NewSparseLU" + if a.Values.Dtype() == core.Complex { + return nil, base.Errf("%s: complex sparse matrices are not supported", name) + } + if len(a.Shape) != 2 || a.Shape[0] != a.Shape[1] { + return nil, base.Errf("%s: needs a square 2-D sparse matrix, got shape %v", name, a.Shape) + } + c, err := CSCFromCOO(a) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + for j := range c.Cols { + for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ { + v := c.Values[p] + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: entry [%d,%d] is not finite", name, c.RowIdx[p], j) + } + } + } + f := &SparseLU{n: c.Cols} + f.piv = make([]int, c.Cols) + for i := range f.piv { + f.piv[i] = i + } + if err := f.factor(c); err != nil { + return nil, base.Errf("%s: %w", name, err) + } + return f, nil +} + +// factor runs the left-looking column elimination. +func (f *SparseLU) factor(c *SparseCSC) error { + const name = "NewSparseLU" + n := c.Cols + f.colRows = make([][]int, n) + f.colVals = make([][]float64, n) + f.rowCols = make([][]int, n) + f.rowVals = make([][]float64, n) + f.diag = make([]float64, n) + x := make([]float64, n) + mark := make([]bool, n) + visited := make([]bool, n) + pattern := make([]int, 0, n) + reach := make([]int, 0, n) + // inversePiv maps an original row to the position it sits at now, and + // piv is its inverse. A pivot swap exchanges two entries of each. + // + // L's stored rows carry the ORIGINAL index of the row, not its + // position, and every read translates through inversePiv. A swap + // then costs two entries of inversePiv instead of a walk over every + // stored multiplier, which for a system that pivots often is the + // difference between a pass over the factor per swap and a pass over + // the factor per factorisation. The final labels are exactly the + // ones the position-space storage would hold: the translation is + // applied once, after the elimination, before the columns are + // sorted. + inversePiv := make([]int, n) + for i := range inversePiv { + inversePiv[i] = i + } + // The depth-first search stack: a column and how far into its + // stored entries the walk has come. + type frame struct { + col int + deg int + } + stack := make([]frame, 0, n) + for k := range n { + // x holds column k of the working matrix: the stored entries + // minus the updates the search produces below. mark flags the + // working positions; visited flags the columns the search has + // walked, which the scatter must not suppress. + pattern = pattern[:0] + for p := c.ColStart[k]; p < c.ColStart[k+1]; p++ { + i := inversePiv[c.RowIdx[p]] + if !mark[i] { + mark[i] = true + pattern = append(pattern, i) + } + x[i] = c.Values[p] + } + reach = reach[:0] + // Depth-first search over the factor's columns from the + // stored entries above the diagonal: it visits every column + // the updates flow through, marks every position they can + // reach, and records the visit post-order, whose reverse is + // the order the updates must run in. + for p := c.ColStart[k]; p < c.ColStart[k+1]; p++ { + // The seed enters in the factor's current row space: the + // stored row may have been swapped since it was written, + // and the search walks the labels as they stand now. + seed := inversePiv[c.RowIdx[p]] + if seed >= k || visited[seed] { + continue + } + visited[seed] = true + stack = stack[:0] + stack = append(stack, frame{seed, 0}) + for len(stack) > 0 { + // The frame is addressed by index, never by pointer: + // an append can reallocate the stack and a held + // pointer would keep writing into the stale array. + top := len(stack) - 1 + entries := f.colRows[stack[top].col] + advanced := false + for stack[top].deg < len(entries) { + // The stored row is an original index; the search + // walks positions, so it translates as it reads. + i := inversePiv[entries[stack[top].deg]] + stack[top].deg++ + if !visited[i] { + visited[i] = true + if !mark[i] { + mark[i] = true + pattern = append(pattern, i) + } + if i < k { + stack = append(stack, frame{i, 0}) + advanced = true + break + } + } + } + if !advanced { + reach = append(reach, stack[top].col) + stack = stack[:top] + } + } + } + for _, j := range reach { + visited[j] = false + } + // The updates run over the post-order list backwards: a + // column's own inputs are all applied before the column + // pushes its U entry down its rows. + for _, j := range slices.Backward(reach) { + + t := x[j] + if t == 0 { + continue + } + // First touch of a U row capacity-sizes it once, so the + // entries that arrive across eliminations append without a + // doubling chain per row. + if cap(f.rowCols[j]) == 0 { + f.rowCols[j] = make([]int, 0, 4) + f.rowVals[j] = make([]float64, 0, 4) + } + f.rowCols[j] = append(f.rowCols[j], k) + f.rowVals[j] = append(f.rowVals[j], t) + for p, i := range f.colRows[j] { + x[inversePiv[i]] -= f.colVals[j][p] * t + } + } + // Partial pivoting: the largest working entry at or below the + // diagonal; ties stay with the position the deterministic + // traversal gathered first. + pivot := -1 + best := 0.0 + for _, i := range pattern { + if i < k { + continue + } + if pivot == -1 || math.Abs(x[i]) > best { + pivot = i + best = math.Abs(x[i]) + } + } + if pivot == -1 || best == 0 { + return base.Errf("the column %d has no pivot; the matrix is singular", k) + } + if pivot != k { + a, b := f.piv[k], f.piv[pivot] + f.piv[k], f.piv[pivot] = b, a + inversePiv[a], inversePiv[b] = inversePiv[b], inversePiv[a] + f.rowCols[k], f.rowCols[pivot] = f.rowCols[pivot], f.rowCols[k] + f.rowVals[k], f.rowVals[pivot] = f.rowVals[pivot], f.rowVals[k] + x[k], x[pivot] = x[pivot], x[k] + } + if math.IsInf(x[k], 0) { + return base.Errf("the pivot at [%d,%d] overflowed to %g; the elimination has left the float64 range", k, k, x[k]) + } + if math.IsNaN(x[k]) || x[k] == 0 { + return base.Errf("the pivot at [%d,%d] is %g; the matrix is singular", k, k, x[k]) + } + f.diag[k] = x[k] + // Column k of L: the working entries below the pivot row, stored + // under the original index of the row they belong to. The column + // is capacity-sized from the working pattern once, so the fill + // appends without a growth chain. + f.colRows[k] = make([]int, 0, max(len(pattern)-k-1, 0)) + f.colVals[k] = make([]float64, 0, max(len(pattern)-k-1, 0)) + for _, i := range pattern { + if i > k && x[i] != 0 { + f.colRows[k] = append(f.colRows[k], f.piv[i]) + f.colVals[k] = append(f.colVals[k], x[i]/f.diag[k]) + } + } + // Restore the working vector and both sets of flags; every + // visited node is in the pattern, so the reset is one pass. + for _, i := range pattern { + x[i] = 0 + mark[i] = false + visited[i] = false + } + x[k] = 0 + } + // Translate the stored original rows into the elimination positions + // they ended at, then sort the columns into canonical ascending row + // order. Both passes are per factorisation, and their result is what + // position-space storage would have held entry for entry. One sort + // scratch serves every column. + var sortScratch []rowValuePair + for j := range n { + rows := f.colRows[j] + for q, i := range rows { + rows[q] = inversePiv[i] + } + sortScratch = f.sortColumn(j, sortScratch) + } + // U's strict triangle is stored by rows, so it escapes the column + // walk's count: add it here, keeping NNZ at L strict + U strict + + // the diagonal. + for j := range n { + f.nnz += len(f.rowCols[j]) + } + return nil +} + +// sortColumn orders one L column by ascending row, values travelling +// with their rows, and records the factor's non-zero count. The caller +// supplies the scratch across columns; the grown buffer comes back for +// the next column. Stored rows are unique within a column, so the +// sorted placement is unique and the scratch cannot move a value. +func (f *SparseLU) sortColumn(j int, scratch []rowValuePair) []rowValuePair { + rows := f.colRows[j] + if !slices.IsSorted(rows) { + pairs := scratch[:0] + for p, r := range rows { + pairs = append(pairs, rowValuePair{r, f.colVals[j][p]}) + } + slices.SortFunc(pairs, func(x, y rowValuePair) int { return x.row - y.row }) + for p, e := range pairs { + f.colRows[j][p] = e.row + f.colVals[j][p] = e.val + } + scratch = pairs + } + f.nnz += len(rows) + return scratch +} + +// Solve computes x = A⁻¹·b for a dense vector b: permute, forward +// substitution with the unit lower L, backward substitution with U, +// undo the permutation. +func (f *SparseLU) Solve(b *core.Array) (*core.Array, error) { + const name = "Solve" + if b.NDim() != 1 { + return nil, base.Errf("%s: the right hand side must be rank 1", name) + } + if b.Dtype() == core.Complex { + return nil, base.Errf("%s: complex right hand sides are not supported", name) + } + if b.Len() != f.n { + return nil, base.Errf("%s: right hand side length %d does not match %d rows", name, b.Len(), f.n) + } + x := core.New(core.Float, []int{f.n}...) + xf := x.RawFloats() + for k := range f.n { + xf[k] = b.FloatAt(f.piv[k]) + } + // L is unit lower triangular: the divide at the diagonal is the + // identity and each column push reads its own x[j] first. + for j := range f.n { + xj := xf[j] + for p, i := range f.colRows[j] { + xf[i] -= f.colVals[j][p] * xj + } + } + for j := f.n - 1; j >= 0; j-- { + sum := xf[j] + for p, c := range f.rowCols[j] { + sum -= f.rowVals[j][p] * xf[c] + } + xf[j] = sum / f.diag[j] + } + // A = Pᵀ·L·U, so U⁻¹·L⁻¹·P·b is the answer itself: with the + // single-sided permutation nothing is undone at the end. + return x, nil +} + +// Permutation returns the row elimination order: position k holds the +// original index of the row factored k-th, with P·A = L·U. The slice +// is a copy, so the caller cannot move the factor's state. +func (f *SparseLU) Permutation() []int { + return slices.Clone(f.piv) +} + +// NNZ returns the count of stored non-zeros in the factor: L's strict +// columns, U's strict rows and the diagonal of U, L's unit diagonal +// excluded. +func (f *SparseLU) NNZ() int { return f.nnz + f.n } diff --git a/linalg/sparselu_test.go b/linalg/sparselu_test.go new file mode 100644 index 0000000..73b6294 --- /dev/null +++ b/linalg/sparselu_test.go @@ -0,0 +1,307 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// nonsymmetricCOO builds a sparse nonsymmetric matrix from the seeded +// generator: a diagonally dominant backbone with extra couplings that +// force fill into the factor, large enough that pivoting has real +// choices. +func nonsymmetricCOO(t *testing.T, seed int64, n int) *core.SparseCOO { + t.Helper() + g := core.NewGenerator(seed) + idx := make([]int64, 0, 6*n) + vals := make([]float64, 0, 6*n) + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 6+2*g.Unit()) + if i+1 < n { + add(i, i+1, -1-g.Unit()) + add(i+1, i, -1-g.Unit()) + } + if i+2 < n { + add(i, i+2, -0.5*g.Unit()) + } + if i+3 < n { + add(i+3, i, -0.3*g.Unit()) + } + if i+5 < n { + add(i, i+5, -0.2*g.Unit()) + } + } + indices, err := core.FromInts(idx, len(vals), 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + return coo +} + +func TestSparseLUSolvesNonsymmetric(t *testing.T) { + const n = 20 + for _, seed := range []int64{1, 2, 3} { + coo := nonsymmetricCOO(t, seed, n) + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + g := core.NewGenerator(seed + 100) + xTrue := core.New(core.Float, n) + for i := range n { + xTrue.RawFloats()[i] = g.NormalUnit() + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + f, err := NewSparseLU(coo) + if err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + x, err := f.Solve(b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + worst := 0.0 + for i := range n { + if d := math.Abs(x.FloatAt(i) - xTrue.FloatAt(i)); d > worst { + worst = d + } + } + if worst > 1e-9 { + t.Fatalf("seed %d: worst solution error %.3g, want under 1e-9", seed, worst) + } + } +} + +// TestSparseLUMatchesDense requires the sparse factor to answer what +// the dense LU with partial pivoting answers for the same system. +func TestSparseLUMatchesDense(t *testing.T) { + const n = 16 + coo := nonsymmetricCOO(t, 7, n) + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + g := core.NewGenerator(8) + xTrue := core.New(core.Float, n) + for i := range n { + xTrue.RawFloats()[i] = g.NormalUnit() + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatalf("MatVec: %v", err) + } + f, err := NewSparseLU(coo) + if err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + sparse, err := f.Solve(b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + denseVals := make([]float64, n*n) + nnz := coo.Indices.Shape()[0] + for i := range nnz { + r := int(coo.Indices.RawInts()[i*2]) + c := int(coo.Indices.RawInts()[i*2+1]) + denseVals[r*n+c] = coo.Values.FloatAt(i) + } + dense, err := Solve(floatsToArray(denseVals, []int{n, n}), b) + if err != nil { + t.Fatalf("dense Solve: %v", err) + } + scale := 0.0 + for i := range n { + if d := math.Abs(dense.FloatAt(i)); d > scale { + scale = d + } + } + for i := range n { + if math.Abs(sparse.FloatAt(i)-dense.FloatAt(i)) > 1e-9*scale { + t.Fatalf("entry %d: sparse %.12g vs dense %.12g", i, sparse.FloatAt(i), dense.FloatAt(i)) + } + } +} + +// TestSparseLUPivotsRowSwap pins the reason partial pivoting exists: +// the anti-diagonal matrix has a zero at (0,0) and only a row swap +// can start the elimination, and the swap must be visible in the +// reported permutation. +func TestSparseLUPivotsRowSwap(t *testing.T) { + anti, err := core.NewSparseCOO( + mustInts(t, []int64{0, 1, 1, 0, 1, 1}, 3, 2), + floatsToArray([]float64{2, 1, 3}, []int{3}), []int{2, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + // Entries: (0,1)=2, (1,0)=1, (1,1)=3: the matrix [[0,2],[1,3]]. + f, err := NewSparseLU(anti) + if err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + if got := f.Permutation(); !slices.Equal(got, []int{1, 0}) { + t.Fatalf("permutation %v, want the rows swapped to [1 0]", got) + } + b := floatsToArray([]float64{4, 5}, []int{2}) + x, err := f.Solve(b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + // [[0,2],[1,3]]·x = 4x2=4... solve by hand: x2 = 4/2 = 2, + // x1 + 3·2 = 5 → x1 = −1. + if math.Abs(x.FloatAt(0)-(-1)) > 1e-12 || math.Abs(x.FloatAt(1)-2) > 1e-12 { + t.Fatalf("solve = [%.6g %.6g], want [-1 2]", x.FloatAt(0), x.FloatAt(1)) + } +} + +func TestSparseLURefusals(t *testing.T) { + good := nonsymmetricCOO(t, 9, 6) + if _, err := NewSparseLU(good); err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + // Complex input. + cplx, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0}, 1, 2), + mustFromComplexes(t, []complex128{1 + 1i}, 1), + []int{1, 1}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseLU(cplx); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("complex input: %v", err) + } + // Rectangular input. + rect, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0}, 1, 2), + floatsToArray([]float64{1}, []int{1}), + []int{1, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseLU(rect); err == nil { + t.Fatal("a rectangular matrix was accepted") + } + // Non-finite stored value. + bad, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0, 1, 1}, 2, 2), + floatsToArray([]float64{4, math.NaN()}, []int{2}), + []int{2, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseLU(bad); err == nil || !strings.Contains(err.Error(), "finite") { + t.Fatalf("NaN entry: %v", err) + } + // Structurally singular: the last column stores nothing, and no + // permutation can pivot an empty column. + sing, err := core.NewSparseCOO( + mustInts(t, []int64{0, 0, 1, 0, 1, 1, 2, 0}, 4, 2), + floatsToArray([]float64{4, 1, 5, 6}, []int{4}), + []int{3, 3}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := NewSparseLU(sing); err == nil || !strings.Contains(err.Error(), "singular") { + t.Fatalf("an empty column: %v", err) + } + // Solve-side refusals. + f, err := NewSparseLU(good) + if err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + if _, err := f.Solve(core.New(core.Float, 3, 3)); err == nil { + t.Fatal("a rank 2 right hand side was accepted") + } + if _, err := f.Solve(floatsToArray([]float64{1, 2, 3}, []int{3})); err == nil { + t.Fatal("a short right hand side was accepted") + } +} + +// TestSparseLUIsDeterministic factors the same matrix twice and +// requires the factor's stored values to come out bit for bit equal, +// the contract every Tensor entry point carries. +func TestSparseLUIsDeterministic(t *testing.T) { + coo := nonsymmetricCOO(t, 11, 18) + f1, err := NewSparseLU(coo) + if err != nil { + t.Fatalf("first factorisation: %v", err) + } + f2, err := NewSparseLU(coo) + if err != nil { + t.Fatalf("second factorisation: %v", err) + } + if f1.NNZ() != f2.NNZ() { + t.Fatalf("factor sizes differ: %d vs %d", f1.NNZ(), f2.NNZ()) + } + for j := range f1.n { + if slices.Compare(f1.colRows[j], f2.colRows[j]) != 0 { + t.Fatalf("column %d patterns differ", j) + } + for p := range f1.colRows[j] { + if f1.colVals[j][p] != f2.colVals[j][p] { + t.Fatalf("L value %d of column %d differs: %.17g vs %.17g", + p, j, f1.colVals[j][p], f2.colVals[j][p]) + } + } + if slices.Compare(f1.rowCols[j], f2.rowCols[j]) != 0 { + t.Fatalf("row %d patterns differ", j) + } + for p := range f1.rowCols[j] { + if f1.rowVals[j][p] != f2.rowVals[j][p] { + t.Fatalf("U value %d of row %d differs: %.17g vs %.17g", + p, j, f1.rowVals[j][p], f2.rowVals[j][p]) + } + } + } + if slices.Compare(f1.piv, f2.piv) != 0 { + t.Fatal("permutations differ") + } +} + +// TestSparseLUFillStaysSparse pins the point of the sparse structure: +// on a banded matrix the factor must stay close to the bandwidth, not +// drift toward a dense triangle. The bound carries slack; the +// measured number is recorded in the log line. +func TestSparseLUFillStaysSparse(t *testing.T) { + const n = 60 + coo := nonsymmetricCOO(t, 13, n) + f, err := NewSparseLU(coo) + if err != nil { + t.Fatalf("NewSparseLU: %v", err) + } + t.Logf("n=%d: LU nnz %d (bandwidth-5 pattern)", n, f.NNZ()) + // NNZ counts L strict + U strict + the U diagonal (L's unit diagonal + // excluded): a bandwidth-5 pattern holds about 2·5n strict entries, + // so the bound sits at 10n, still far from the dense n²/2 triangle. + if f.NNZ() > 10*n { + t.Fatalf("LU nnz %d drifted past 10n on a banded pattern", f.NNZ()) + } + perm := f.Permutation() + sorted := slices.Clone(perm) + slices.Sort(sorted) + for i := range sorted { + if sorted[i] != i { + t.Fatalf("permutation entry %d holds %d; not a permutation", i, sorted[i]) + } + } + perm[0] = -1 + if f.Permutation()[0] == -1 { + t.Fatal("Permutation exposed the factor's internal slice") + } +} diff --git a/linalg/sparseops.go b/linalg/sparseops.go new file mode 100644 index 0000000..55dae05 --- /dev/null +++ b/linalg/sparseops.go @@ -0,0 +1,150 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "slices" + +// Operations on the CSR sparse matrix beyond the matrix-vector +// product: dense and sparse products, and the transpose. The sparse +// product keeps its result sparse by construction, row by row, so +// chained sparse expressions never materialise a dense intermediate +// unless the caller asks for one. + +// MatMulDense returns A·x for a dense matrix x of shape (Cols, k): +// every output column is one matrix-vector product against the same +// sparse structure, so the non-zeros are streamed once for the whole +// product. +func (c *SparseCSR) MatMulDense(x *core.Array) (*core.Array, error) { + if x.Dtype() == core.Complex { + return nil, base.Errf("MatMulDense: complex operands are not supported") + } + if x.NDim() != 2 || x.Shape()[0] != c.Cols { + return nil, base.Errf("MatMulDense: operand shape %s does not match (Rows, Cols) = (%d, %d)", + base.ShapeText(x.Shape()), c.Rows, c.Cols) + } + k := x.Shape()[1] + out := core.New(core.Float, []int{c.Rows, k}...) + xc := contiguousF64(x) + // Row-by-row with all k columns in the inner loop keeps the + // dense operand's rows hot; the output payload is taken once so the + // update is a plain store. + yc := out.RawFloats() + for i := range c.Rows { + rs, re := c.RowStart[i], c.RowStart[i+1] + yi := yc[i*k : i*k+k : i*k+k] + for p := rs; p < re; p++ { + v := c.Values[p] + j := c.ColIdx[p] + xj := xc[j*k : j*k+k : j*k+k] + for m := range yi { + yi[m] += v * xj[m] + } + } + } + return out, nil +} + +// MatMulSparse returns the sparse product A·B, valid when A.Cols +// equals B.Rows. Each output row gathers the products of the row's +// non-zeros with the corresponding rows of B into an accumulator +// indexed by column, so the cost is the number of stored multiply-adds +// plus one sort of the columns the row touched, and the result stores +// only what survives. +func (c *SparseCSR) MatMulSparse(o *SparseCSR) (*SparseCSR, error) { + if c.Cols != o.Rows { + return nil, base.Errf("MatMulSparse: inner dimensions disagree, %d vs %d", c.Cols, o.Rows) + } + // An output row gathers one multiply-add per stored entry of A it + // reaches, but it can store no more entries than it has columns, so + // the row's bound is the smaller of the two and their sum bounds the + // result's non-zeros, which sizes the output in one pass. The bound + // is tight for a dense product and never exceeds it. + work := 0 + for i := range c.Rows { + row := 0 + for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ { + r := c.ColIdx[p] + row += o.RowStart[r+1] - o.RowStart[r] + } + work += min(row, o.Cols) + } + out := &SparseCSR{Rows: c.Rows, Cols: o.Cols} + out.RowStart = make([]int, c.Rows+1) + out.ColIdx = make([]int, 0, work) + out.Values = make([]float64, 0, work) + // The accumulator is dense over the columns and addressed by the + // column index itself instead of by a hash of it. touched holds the + // columns the current row has contributed to, each once, so the sort + // below orders exactly the columns that were written; inAcc marks the + // ones the list already holds, which keeps a second contribution to + // the same column from listing it twice. + acc := make([]float64, o.Cols) + inAcc := make([]bool, o.Cols) + touched := make([]int, 0, 16) + for i := range c.Rows { + touched = touched[:0] + for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ { + v := c.Values[p] + r := c.ColIdx[p] + rs, re := o.RowStart[r], o.RowStart[r+1] + oCols := o.ColIdx[rs:re:re] + oVals := o.Values[rs:re:re] + for q := range oCols { + j := oCols[q] + if !inAcc[j] { + inAcc[j] = true + touched = append(touched, j) + } + acc[j] += v * oVals[q] + } + } + slices.Sort(touched) + for _, j := range touched { + if v := acc[j]; v != 0 { + out.ColIdx = append(out.ColIdx, j) + out.Values = append(out.Values, v) + out.RowStart[i+1]++ + } + acc[j] = 0 + inAcc[j] = false + } + } + for i := range c.Rows { + out.RowStart[i+1] += out.RowStart[i] + } + return out, nil +} + +// Transpose returns the transpose in CSR form, built by the standard +// counting pass: one sweep counts the entries per output row, the +// second places them. +func (c *SparseCSR) Transpose() *SparseCSR { + out := &SparseCSR{Rows: c.Cols, Cols: c.Rows} + out.RowStart = make([]int, c.Cols+1) + for _, j := range c.ColIdx { + out.RowStart[j+1]++ + } + for i := range c.Cols { + out.RowStart[i+1] += out.RowStart[i] + } + next := make([]int, c.Cols) + copy(next, out.RowStart[:c.Cols]) + out.ColIdx = make([]int, len(c.ColIdx)) + out.Values = make([]float64, len(c.Values)) + for i := range c.Rows { + for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ { + j := c.ColIdx[p] + q := next[j] + out.ColIdx[q] = i + out.Values[q] = c.Values[p] + next[j]++ + } + } + return out +} diff --git a/linalg/sparseops_test.go b/linalg/sparseops_test.go new file mode 100644 index 0000000..d5aae09 --- /dev/null +++ b/linalg/sparseops_test.go @@ -0,0 +1,226 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// sparseCSRFromEntries builds a SparseCSR straight from (row, col, +// value) triples, going through the public COO path. +func sparseCSRFromEntries(t *testing.T, rows, cols int, entries [][3]float64) *SparseCSR { + t.Helper() + idx := make([]int64, 0, len(entries)*2) + vals := make([]float64, 0, len(entries)) + for _, e := range entries { + idx = append(idx, int64(e[0]), int64(e[1])) + vals = append(vals, e[2]) + } + indices, ierr := core.FromInts(idx, len(entries), 2) + if ierr != nil { + t.Fatalf("FromInts: %v", ierr) + } + coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(entries)}), []int{rows, cols}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatalf("CSRFromCOO: %v", err) + } + return csr +} + +// TestSparseMatMulDense checks the sparse-dense product against an +// explicit dense computation on a ragged 4×5 matrix. +func TestSparseMatMulDense(t *testing.T) { + csr := sparseCSRFromEntries(t, 4, 5, [][3]float64{ + {0, 0, 2}, {0, 3, -1}, + {1, 1, 3}, {1, 4, 0.5}, + {2, 0, -1}, {2, 2, 4}, {2, 4, 2}, + {3, 3, 1}, + }) + x := mustFloats(t, []float64{ + 1, -1, 2, 0, 3, + 0.5, 1, 1, 1, -2, + }, 5, 2) + got, err := csr.MatMulDense(x) + if err != nil { + t.Fatalf("MatMulDense: %v", err) + } + if got.Shape()[0] != 4 || got.Shape()[1] != 2 { + t.Fatalf("result shape %v, want (4, 2)", got.Shape()) + } + for i := range 4 { + for m := range 2 { + s := 0.0 + for j := range 5 { + s += csrValuesAt(csr, i, j) * x.FloatAt(j*2+m) + } + if math.Abs(got.FloatAt(i*2+m)-s) > 1e-14 { + t.Fatalf("out(%d, %d) = %g, want %g", i, m, got.FloatAt(i*2+m), s) + } + } + } + if _, err := csr.MatMulDense(mustFloats(t, []float64{1, 2, 3, 4, 5})); err == nil { + t.Fatal("vector operand: want an error") + } +} + +// csrValuesAt reads a CSR entry that may be stored or implicit zero. +func csrValuesAt(c *SparseCSR, i, j int) float64 { + for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ { + if c.ColIdx[p] == j { + return c.Values[p] + } + } + return 0 +} + +// TestSparseMatMulSparse checks the sparse product against its dense +// equivalent, including cancellation producing an implicit zero and +// the transpose route for the asymmetric case. +func TestSparseMatMulSparse(t *testing.T) { + a := sparseCSRFromEntries(t, 3, 4, [][3]float64{ + {0, 0, 1}, {0, 2, 2}, + {1, 1, -1}, {1, 3, 4}, + {2, 0, 0.5}, {2, 2, 1}, + }) + b := sparseCSRFromEntries(t, 4, 3, [][3]float64{ + {0, 1, 3}, + {1, 0, 2}, {1, 2, -1}, + {2, 1, 1.5}, {2, 2, 2}, + {3, 0, -0.5}, + }) + ab, err := a.MatMulSparse(b) + if err != nil { + t.Fatalf("MatMulSparse: %v", err) + } + if ab.Rows != 3 || ab.Cols != 3 { + t.Fatalf("product shape %d×%d, want 3×3", ab.Rows, ab.Cols) + } + for i := range 3 { + for j := range 3 { + s := 0.0 + for k := range 4 { + s += csrValuesAt(a, i, k) * csrValuesAt(b, k, j) + } + if math.Abs(csrValuesAt(ab, i, j)-s) > 1e-14 { + t.Fatalf("(%d, %d) = %g, want %g", i, j, csrValuesAt(ab, i, j), s) + } + } + } + if _, err := a.MatMulSparse(sparseCSRFromEntries(t, 3, 2, [][3]float64{{0, 0, 1}})); err == nil { + t.Fatal("inner dimension mismatch: want an error") + } +} + +// TestSparseTranspose checks the transpose against dense arithmetic +// and against double transposition restoring the original. +func TestSparseTranspose(t *testing.T) { + csr := sparseCSRFromEntries(t, 3, 5, [][3]float64{ + {0, 1, 2}, {0, 4, -3}, + {1, 0, 1.5}, + {2, 1, 4}, {2, 2, -1}, {2, 4, 0.25}, + }) + at := csr.Transpose() + if at.Rows != 5 || at.Cols != 3 { + t.Fatalf("transpose shape %d×%d, want 5×3", at.Rows, at.Cols) + } + for i := range 5 { + for j := range 3 { + if csrValuesAt(at, i, j) != csrValuesAt(csr, j, i) { + t.Fatalf("transpose(%d, %d) = %g, want %g", i, j, + csrValuesAt(at, i, j), csrValuesAt(csr, j, i)) + } + } + } + back := at.Transpose() + for i := range 3 { + for j := range 5 { + if csrValuesAt(back, i, j) != csrValuesAt(csr, i, j) { + t.Fatalf("double transpose(%d, %d) = %g, want %g", i, j, + csrValuesAt(back, i, j), csrValuesAt(csr, i, j)) + } + } + } + // The transpose product Aᵀ·A via MatMulSparse must agree with the + // dense Gram matrix. + g, err := at.MatMulSparse(csr) + if err != nil { + t.Fatalf("Aᵀ·A: %v", err) + } + for i := range 5 { + for j := range 5 { + s := 0.0 + for k := range 3 { + s += csrValuesAt(csr, k, i) * csrValuesAt(csr, k, j) + } + if math.Abs(csrValuesAt(g, i, j)-s) > 1e-14 { + t.Fatalf("Gram(%d, %d) = %g, want %g", i, j, csrValuesAt(g, i, j), s) + } + } + } +} + +// TestSparseOpsLargeCheck runs the three operations against dense +// arithmetic on a banded 40×40 structure, the shape PDE stencils +// produce. +func TestSparseOpsLargeCheck(t *testing.T) { + const n = 40 + vals := make([][3]float64, 0, 3*n) + for i := range n { + vals = append(vals, [3]float64{float64(i), float64(i), 4}) + if i+1 < n { + vals = append(vals, [3]float64{float64(i), float64(i + 1), -1}) + vals = append(vals, [3]float64{float64(i + 1), float64(i), -1}) + } + } + a := sparseCSRFromEntries(t, n, n, vals) + x := make([]float64, n*n) + for i := range n { + for j := range n { + x[i*n+j] = math.Sin(float64(i + j + 1)) + } + } + xd := floatsToArray(x, []int{n, n}) + got, err := a.MatMulDense(xd) + if err != nil { + t.Fatalf("MatMulDense: %v", err) + } + for i := range n { + for j := range n { + s := 0.0 + for k := range n { + if v := csrValuesAt(a, i, k); v != 0 { + s += v * x[k*n+j] + } + } + if math.Abs(got.FloatAt(i*n+j)-s) > 1e-12 { + t.Fatalf("(%d, %d): %g vs %g", i, j, got.FloatAt(i*n+j), s) + } + } + } + sq, err := a.MatMulSparse(a) + if err != nil { + t.Fatalf("A·A: %v", err) + } + for i := range n { + for j := range n { + s := 0.0 + for k := range n { + if v := csrValuesAt(a, i, k); v != 0 { + if w := csrValuesAt(a, k, j); w != 0 { + s += v * w + } + } + } + if math.Abs(csrValuesAt(sq, i, j)-s) > 1e-12 { + t.Fatalf("A²(%d, %d): %g vs %g", i, j, csrValuesAt(sq, i, j), s) + } + } + } +} diff --git a/linalg/sparsesolve.go b/linalg/sparsesolve.go new file mode 100644 index 0000000..d034f4b --- /dev/null +++ b/linalg/sparsesolve.go @@ -0,0 +1,207 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Sparse symmetric positive-definite solve. The dense `Solve` factors +// A once and back-substitutes, which costs O(n³) time and O(n²) memory +// and is the right answer for a modest, dense system. When A is large +// and sparse that factorisation is both unaffordable and wasteful: the +// fill-in alone can exceed the memory the non-zeros needed. SpSolve +// runs the preconditioned conjugate gradient instead, whose every step +// is one sparse-matrix-vector product plus a handful of vector +// operations, so the cost tracks the non-zero count rather than the +// dimension. +// +// Conjugate gradient applies only to a symmetric positive-definite +// matrix, and it is not a drop-in for `Solve`: the answer is an +// approximation within a stated residual tolerance, not the exact +// solution a factorisation returns. Symmetry is verified the same way +// `SpEigen` verifies it. Positive-definiteness cannot be checked +// cheaply up front, so it is caught in flight: a non-positive search +// curvature is the matrix declaring itself indefinite. + +// spSolveTol is the default relative residual when the caller passes +// zero. The iteration default is the dimension rather than a constant: +// an exact arithmetic conjugate gradient terminates in at most n +// steps, so a larger budget cannot buy convergence, only round-off +// work. +const spSolveTol = 1e-10 + +// checkSparseSquare validates the inputs every real sparse solver +// shares: a is a square 2-D matrix of non-complex dtype and, when b is +// not nil, b is a real right-hand side vector of the matrix dimension. +// It returns the dimension n. +func checkSparseSquare(name string, a *core.SparseCOO, b *core.Array) (int, error) { + if a.Values.Dtype() == core.Complex { + return 0, base.Errf("%s: complex sparse matrices are not supported", name) + } + if len(a.Shape) != 2 || a.Shape[0] != a.Shape[1] { + return 0, base.Errf("%s: needs a square 2-D sparse matrix, got shape %v", name, a.Shape) + } + n := a.Shape[0] + if n == 0 { + return 0, base.Errf("%s: zero-sized matrix, got shape %v", name, a.Shape) + } + if b != nil { + if b.NDim() != 1 || b.Shape()[0] != n { + return 0, base.Errf("%s: right-hand side must be a vector of length %d, got shape %s", + name, n, base.ShapeText(b.Shape())) + } + if b.Dtype() == core.Complex { + return 0, base.Errf("%s: complex right-hand side is not supported", name) + } + } + return n, nil +} + +// pickPreconditioner resolves the optional ILU argument of the Krylov +// solvers, falling back to the Jacobi diagonal of c when none was +// given. An ILU built for another dimension is refused: applying it +// would index past its vectors, or precondition with the wrong +// factorisation without a word. +func pickPreconditioner(name string, c *sparseCSR, precond []*SparseILU) (*SparseILU, []float64, error) { + if len(precond) > 0 { + if precond[0] == nil { + return nil, nil, base.Errf("%s: the preconditioner is nil", name) + } + if precond[0].n != c.n { + return nil, nil, base.Errf("%s: the preconditioner was built for dimension %d, the system is %d", + name, precond[0].n, c.n) + } + return precond[0], nil, nil + } + diag, err := c.diagonal(name) + if err != nil { + return nil, nil, err + } + return nil, diag, nil +} + +// vectorF64 flattens a rank-1 array into a plain float64 working +// vector of the expected length n. +func vectorF64(b *core.Array, n int) []float64 { + r := make([]float64, n) + for i := range n { + r[i] = b.FloatAt(i) + } + return r +} + +// SpSolve returns the vector x solving A·x = b for a real symmetric +// positive-definite sparse A, by preconditioned conjugate gradient +// with a Jacobi (diagonal) preconditioner. The right-hand side b is a +// rank-1 vector of length n. +// +// The stopping rule is the relative residual ‖b − A·x‖₂ ≤ tol·‖b‖₂. +// A tol ≤ 0 uses 1e-10, and maxIter ≤ 0 uses n steps. x starts at +// zero, so the initial residual is b itself and the loop needs no +// separate first matvec. The default preconditioner is Jacobi +// scaling; an ILU(0) factorisation from NewSparseILU may be passed to +// replace it. +// +// An unconverged solve is an error naming the residual achieved, with +// no estimate returned, rather than a silent approximation. +// The matrix must be symmetric, and a zero diagonal entry is refused: +// the Jacobi preconditioner divides by it, and for a symmetric +// positive-definite matrix the diagonal is necessarily positive. +func SpSolve(a *core.SparseCOO, b *core.Array, tol float64, maxIter int, precond ...*SparseILU) (*core.Array, error) { + const name = "SpSolve" + n, err := checkSparseSquare(name, a, b) + if err != nil { + return nil, err + } + c, err := symmetricCSR(a, name) + if err != nil { + return nil, err + } + ilu, diag, err := pickPreconditioner(name, c, precond) + if err != nil { + return nil, err + } + + if tol <= 0 { + tol = spSolveTol + } + if maxIter <= 0 { + maxIter = n + } + + // x starts at zero, so r = b − A·x is just b and the loop needs no + // opening matvec. + r := vectorF64(b, n) + bNorm := norm2F64(r) + if bNorm == 0 { + // The exact solution of A·0 = 0 is the zero vector; a relative + // stopping rule would otherwise divide by zero. + return core.Zeros(core.Float, n) + } + + z := make([]float64, n) + iluPrecondition(z, r, ilu, diag) + p := append([]float64(nil), z...) + rz := dotF64(r, z) + x := make([]float64, n) + ap := make([]float64, n) + + for iter := range maxIter { + c.matVec(p, ap) + pAp := dotF64(p, ap) + if !finiteF64(pAp) || pAp <= 0 { + // A non-positive or non-finite curvature means the matrix + // is not positive-definite, or has overflowed the + // representable range; either way conjugate gradient + // cannot proceed. A non-finite value must be an error, + // not a NaN that slips through the <= 0 test. + return nil, base.Errf("%s: matrix is not positive-definite (curvature %.3g at step %d)", + name, pAp, iter+1) + } + alpha := rz / pAp + for i := range n { + x[i] += alpha * p[i] + r[i] -= alpha * ap[i] + } + // norm2F64 skips NaN entries, so an all-NaN residual reads as a + // zero norm; a non-finite state is a breakdown before the + // convergence test can mistake it for an exact solve. + if !vecFinite(r) { + return nil, base.Errf("%s: non-finite residual at step %d", name, iter+1) + } + if norm2F64(r) <= tol*bNorm { + return floatsToArray(x, []int{n}), nil + } + iluPrecondition(z, r, ilu, diag) + rzNext := dotF64(r, z) + if rz == 0 { + return nil, base.Errf("%s: breakdown at step %d (preconditioned residual vanished)", name, iter+1) + } + beta := rzNext / rz + for i := range n { + p[i] = z[i] + beta*p[i] + } + rz = rzNext + } + return nil, base.Errf("%s: no convergence in %d steps, residual %.3g (tolerance %.3g)", + name, maxIter, norm2F64(r), tol*bNorm) +} + +// diagonal returns the main diagonal, which the Jacobi preconditioner +// divides by. A missing or zero entry makes the preconditioner +// undefined and, for a symmetric positive-definite matrix, signals +// that the input is not one. +func (c *sparseCSR) diagonal(name string) ([]float64, error) { + d := make([]float64, c.n) + for i := range c.n { + v, ok := c.at(i, i) + if !ok || v == 0 { + return nil, base.Errf("%s: zero or missing diagonal entry at %d", name, i) + } + d[i] = v + } + return d, nil +} diff --git a/linalg/sparsesolve_test.go b/linalg/sparsesolve_test.go new file mode 100644 index 0000000..bf021ff --- /dev/null +++ b/linalg/sparsesolve_test.go @@ -0,0 +1,375 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// spdTridiagonal builds a symmetric diagonally dominant tridiagonal +// matrix, which is positive-definite by the Gershgorin bound, so the +// conjugate gradient is guaranteed to apply. The diagonal is kept +// above the off-diagonal sum in absolute value. +func spdTridiagonal(n int, diag, off float64) []float64 { + vals := make([]float64, n*n) + for i := range n { + vals[i*n+i] = diag + if i+1 < n { + vals[i*n+i+1] = off + vals[(i+1)*n+i] = off + } + } + return vals +} + +// denseSolve solves A·x = b through the dense LU path, the reference +// the sparse solver is compared against. +func denseSolve(t *testing.T, vals []float64, n int, rhs []float64) []float64 { + t.Helper() + a, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + b, err := core.FromFloats(rhs, n) + if err != nil { + t.Fatalf("FromFloats rhs: %v", err) + } + x, err := Solve(a, b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + out := make([]float64, n) + for i := range n { + out[i] = x.FloatAt(i) + } + return out +} + +// residualNorm norms ‖b − A·x‖₂, the direct check that a returned x +// solves the system rather than merely being returned. +func residualNorm(t *testing.T, vals []float64, x, rhs []float64, n int) float64 { + t.Helper() + sum := 0.0 + for i := range n { + ax := 0.0 + for j := range n { + ax += vals[i*n+j] * x[j] + } + d := rhs[i] - ax + sum += d * d + } + return math.Sqrt(sum) +} + +// TestSpSolveMatchesDense compares the sparse solver against the dense +// `Solve` on symmetric positive-definite systems, which is the +// contract the two share. +func TestSpSolveMatchesDense(t *testing.T) { + cases := []struct { + name string + n int + diag float64 + off float64 + }{ + {name: "identity_like", n: 3, diag: 4, off: 0}, + {name: "weakly_coupled", n: 5, diag: 4, off: -1}, + {name: "strongly_coupled", n: 8, diag: 10, off: 3}, + {name: "larger_system", n: 40, diag: 5, off: 1.5}, + } + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + vals := spdTridiagonal(tt.n, tt.diag, tt.off) + sp := sparseFromDense(t, vals, tt.n) + rhs := make([]float64, tt.n) + for i := range tt.n { + rhs[i] = float64(i+1) / float64(tt.n) + } + b, err := core.FromFloats(rhs, tt.n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := SpSolve(sp, b, 0, 0) + if err != nil { + t.Fatalf("SpSolve: %v", err) + } + want := denseSolve(t, vals, tt.n, rhs) + x := make([]float64, tt.n) + for i := range tt.n { + x[i] = got.FloatAt(i) + } + for i := range tt.n { + if math.Abs(x[i]-want[i]) > 1e-8*(1+math.Abs(want[i])) { + t.Fatalf("x[%d] = %.12g, want %.12g", i, x[i], want[i]) + } + } + if r := residualNorm(t, vals, x, rhs, tt.n); r > 1e-8 { + t.Fatalf("residual ‖b-Ax‖ = %.3g, want <= 1e-8", r) + } + }) + } +} + +// TestSpSolveNonDiagonalMatrix exercises a symmetric +// positive-definite matrix that is not tridiagonal, so the sparse +// structure carries a genuinely two-dimensional sparsity pattern. +func TestSpSolveNonDiagonalMatrix(t *testing.T) { + const n = 6 + // A = L·Lᵀ for a lower-triangular L with a positive diagonal, + // which is symmetric positive-definite by construction. + l := []float64{ + 2, 0, 0, 0, 0, 0, + 1, 3, 0, 0, 0, 0, + 0, 1, 2, 0, 0, 0, + 1, 0, 1, 4, 0, 0, + 0, 0, 0, 1, 2, 0, + 0, 1, 0, 0, 1, 3, + } + vals := make([]float64, n*n) + for i := range n { + for j := range n { + s := 0.0 + for k := range n { + if k <= i && k <= j { + s += l[i*n+k] * l[j*n+k] + } + } + vals[i*n+j] = s + } + } + sp := sparseFromDense(t, vals, n) + rhs := []float64{1, 2, 3, 4, 5, 6} + b, err := core.FromFloats(rhs, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := SpSolve(sp, b, 0, 0) + if err != nil { + t.Fatalf("SpSolve: %v", err) + } + x := make([]float64, n) + for i := range n { + x[i] = got.FloatAt(i) + } + want := denseSolve(t, vals, n, rhs) + for i := range n { + if math.Abs(x[i]-want[i]) > 1e-8*(1+math.Abs(want[i])) { + t.Fatalf("x[%d] = %.12g, want %.12g", i, x[i], want[i]) + } + } +} + +// TestSpSolveZeroRightHandSide pins the degenerate contract: the +// solution of A·x = 0 is the zero vector, returned without dividing +// by a zero residual norm. +func TestSpSolveZeroRightHandSide(t *testing.T) { + const n = 4 + vals := spdTridiagonal(n, 4, -1) + sp := sparseFromDense(t, vals, n) + b, err := core.Zeros(core.Float, n) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + got, err := SpSolve(sp, b, 0, 0) + if err != nil { + t.Fatalf("SpSolve: %v", err) + } + for i := range n { + if got.FloatAt(i) != 0 { + t.Fatalf("x[%d] = %.12g, want 0 for a zero right-hand side", i, got.FloatAt(i)) + } + } +} + +// TestSpSolveDeterminism checks that the solve is reproducible: the +// iteration starts from a fixed zero guess and draws nothing random. +func TestSpSolveDeterminism(t *testing.T) { + const n = 12 + vals := spdTridiagonal(n, 6, -1) + sp := sparseFromDense(t, vals, n) + rhs := make([]float64, n) + for i := range n { + rhs[i] = math.Sin(float64(i + 1)) + } + b, err := core.FromFloats(rhs, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + x1, err := SpSolve(sp, b, 0, 0) + if err != nil { + t.Fatalf("SpSolve #1: %v", err) + } + x2, err := SpSolve(sp, b, 0, 0) + if err != nil { + t.Fatalf("SpSolve #2: %v", err) + } + for i := range n { + if x1.FloatAt(i) != x2.FloatAt(i) { + t.Fatalf("x[%d] = %.12g vs %.12g across runs", i, x1.FloatAt(i), x2.FloatAt(i)) + } + } +} + +// TestSpSolveTolerance drives the stopping rule explicitly: a loose +// tolerance stops early with a larger residual, a tight one converges +// further, and both are reported honestly. +func TestSpSolveTolerance(t *testing.T) { + const n = 60 + vals := spdTridiagonal(n, 4, -1) + sp := sparseFromDense(t, vals, n) + rhs := make([]float64, n) + for i := range n { + rhs[i] = 1 + } + b, err := core.FromFloats(rhs, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + loose, err := SpSolve(sp, b, 1e-2, 0) + if err != nil { + t.Fatalf("SpSolve loose: %v", err) + } + tight, err := SpSolve(sp, b, 1e-12, 0) + if err != nil { + t.Fatalf("SpSolve tight: %v", err) + } + xl := make([]float64, n) + xt := make([]float64, n) + for i := range n { + xl[i] = loose.FloatAt(i) + xt[i] = tight.FloatAt(i) + } + rl := residualNorm(t, vals, xl, rhs, n) + rt := residualNorm(t, vals, xt, rhs, n) + if rl < rt { + t.Fatalf("loose tolerance gave residual %.3g, tighter than the tight run's %.3g", rl, rt) + } + if rt > 1e-9 { + t.Fatalf("tight tolerance gave residual %.3g, want <= 1e-9", rt) + } +} + +// TestSpSolveNoConvergence pins the unconverged contract: the budget +// runs out and the solve reports the residual instead of returning a +// silent approximation. +func TestSpSolveNoConvergence(t *testing.T) { + const n = 200 + vals := spdTridiagonal(n, 2, -0.999) + sp := sparseFromDense(t, vals, n) + rhs := make([]float64, n) + for i := range n { + rhs[i] = 1 + } + b, err := core.FromFloats(rhs, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + // A nearly singular system needs far more than two steps. + if _, err := SpSolve(sp, b, 1e-14, 2); err == nil { + t.Fatal("expected an error when the iteration budget is exhausted") + } +} + +// TestSpSolveRejectsInvalid pins the error contract for every input +// the solver cannot honestly answer. +func TestSpSolveRejectsInvalid(t *testing.T) { + t.Run("complex_sparse", func(t *testing.T) { + idx, err := core.FromInts([]int64{0, 0}, 1, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + vals, err := core.FromComplexes([]complex128{1}, 1) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + sp, err := core.NewSparseCOO(idx, vals, []int{1, 1}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + b, _ := core.FromFloats([]float64{1}, 1) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for a complex sparse matrix") + } + }) + t.Run("not_square", func(t *testing.T) { + idx, err := core.FromInts([]int64{0, 0}, 1, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + vals, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + sp, err := core.NewSparseCOO(idx, vals, []int{1, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + b, _ := core.FromFloats([]float64{1}, 1) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for a non-square matrix") + } + }) + t.Run("zero_sized", func(t *testing.T) { + idx, _ := core.FromInts(nil, 0, 2) + vals, _ := core.FromFloats(nil, 0) + sp := &core.SparseCOO{Indices: idx, Values: vals, Shape: []int{0, 0}} + b, _ := core.FromFloats(nil, 0) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for a zero-sized matrix") + } + }) + t.Run("rhs_wrong_length", func(t *testing.T) { + vals := spdTridiagonal(3, 4, -1) + sp := sparseFromDense(t, vals, 3) + b, _ := core.FromFloats([]float64{1, 2}, 2) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for a right-hand side of the wrong length") + } + }) + t.Run("rhs_not_vector", func(t *testing.T) { + vals := spdTridiagonal(2, 4, -1) + sp := sparseFromDense(t, vals, 2) + b, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for a rank-2 right-hand side") + } + }) + t.Run("complex_rhs", func(t *testing.T) { + vals := spdTridiagonal(2, 4, -1) + sp := sparseFromDense(t, vals, 2) + b, _ := core.FromComplexes([]complex128{1, 1}, 2) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for a complex right-hand side") + } + }) + t.Run("zero_diagonal", func(t *testing.T) { + // Symmetric but singular, and the Jacobi preconditioner + // cannot divide by a zero diagonal. + sp := sparseFromDense(t, []float64{0, 1, 1, 0}, 2) + b, _ := core.FromFloats([]float64{1, 1}, 2) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for a zero diagonal entry") + } + }) + t.Run("asymmetric", func(t *testing.T) { + sp := sparseFromDense(t, []float64{4, 1, 2, 4}, 2) + b, _ := core.FromFloats([]float64{1, 1}, 2) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for an asymmetric matrix") + } + }) + t.Run("indefinite", func(t *testing.T) { + // Symmetric with eigenvalues 3 and -1, so not + // positive-definite. The right-hand side must excite the + // negative eigendirection [1,-1]: with b = [1,1] the system + // still has the exact solution [1/3,1/3] and the curvature + // stays positive, so the matrix would not be caught. + sp := sparseFromDense(t, []float64{1, 2, 2, 1}, 2) + b, _ := core.FromFloats([]float64{1, -1}, 2) + if _, err := SpSolve(sp, b, 0, 0); err == nil { + t.Fatal("expected an error for an indefinite matrix") + } + }) +} diff --git a/linalg/sparsexp.go b/linalg/sparsexp.go new file mode 100644 index 0000000..d717ffe --- /dev/null +++ b/linalg/sparsexp.go @@ -0,0 +1,197 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// Sparse matrix exponential applied to a vector. The dense +// `MatrixExp` forms the whole n×n exponential, which is O(n³) time and +// O(n²) memory and answers the question "what is exp(A)". Often the +// question is narrower: "what is exp(A)·v", for one vector, as arises +// in diffusion, network propagation and the action of a matrix +// function on an initial state. Forming the whole exponential to +// multiply one vector wastes both. +// +// SpExpApply computes exp(A)·v through a Krylov projection. The +// Krylov subspace K_m(A, v) = span{v, A·v, A²·v, …} is where the +// action of exp(A) on v actually lives, and on that subspace A +// restricts to a small tridiagonal matrix T. The answer is then +// Q·exp(T)·(‖v‖·e₁): one small dense exponential, which reuses the +// very `MatrixExp` already written, plus a lift back through the +// Lanczos basis. +// +// The subspace must be exactly K_m(A, v), so this variant of Lanczos +// stops when the recurrence collapses rather than restarting in a +// fresh direction the way the eigensolver does: a restarted direction +// would leave the Krylov space of v and the projection would no longer +// approximate exp(A)·v. A collapse simply means the space is smaller +// than the budget, which is exact rather than a failure. + +// spExpSteps is the Krylov dimension used when the caller passes zero: +// 40, matching the block the sparse eigensolver budgets past its +// eigenpair count. A matrix of n ≤ 40 gets its whole space and an +// exact answer; above that the result is an approximation whose +// accuracy grows with the step count. +const spExpSteps = 40 + +// SpExpApply returns exp(A)·v for a real symmetric sparse A and a +// rank-1 vector v of length n, by Krylov projection onto the subspace +// generated by A and v. +// +// steps is the Krylov dimension: a value ≤ 0 uses min(n, 40), and a +// value ≥ n decomposes the whole space, so the answer is exact rather +// than approximate. The matrix must be square, real and symmetric, +// verified with the same 1e-12 tolerance the other sparse entry points +// apply. A zero v gives the zero vector exactly, since exp(A)·0 = 0. +func SpExpApply(a *core.SparseCOO, v *core.Array, steps int) (*core.Array, error) { + if a.Values.Dtype() == core.Complex { + return nil, base.Errf("SpExpApply: complex sparse matrices are not supported") + } + if len(a.Shape) != 2 || a.Shape[0] != a.Shape[1] { + return nil, base.Errf("SpExpApply: needs a square 2-D sparse matrix, got shape %v", a.Shape) + } + n := a.Shape[0] + if n == 0 { + return nil, base.Errf("SpExpApply: zero-sized matrix, got shape %v", a.Shape) + } + if v.NDim() != 1 || v.Shape()[0] != n { + return nil, base.Errf("SpExpApply: vector must have length %d, got shape %s", n, base.ShapeText(v.Shape())) + } + if v.Dtype() == core.Complex { + return nil, base.Errf("SpExpApply: complex vectors are not supported") + } + c, err := symmetricCSR(a, "SpExpApply") + if err != nil { + return nil, err + } + if steps <= 0 { + steps = min(n, spExpSteps) + } + if steps > n { + steps = n + } + + start := make([]float64, n) + copy(start, contiguousF64(v)) + nrmv := norm2F64(start) + if nrmv == 0 { + return core.Zeros(core.Float, n) + } + for i := range n { + start[i] /= nrmv + } + + alphas, betas, basis := c.krylov(start, steps) + + // The small tridiagonal matrix that A restricts to on the Krylov + // subspace; its exponential is the action in those coordinates. + m := len(alphas) + tMat := make([]float64, m*m) + for i := range m { + tMat[i*m+i] = alphas[i] + if i+1 < m { + tMat[i*m+i+1] = betas[i] + tMat[(i+1)*m+i] = betas[i] + } + } + tExp, err := MatrixExp(floatsToArray(tMat, []int{m, m})) + if err != nil { + return nil, err + } + + // exp(T)·(‖v‖·e₁) is ‖v‖ times the first column of exp(T); the + // lift back through the basis then needs no separate product. + out := make([]float64, n) + tc := contiguousF64(tExp) + for p := range m { + y := nrmv * tc[p*m] + if y == 0 { + continue + } + row := basis[p*n : (p+1)*n] + for i := range n { + out[i] += y * row[i] + } + } + return floatsToArray(out, []int{n}), nil +} + +// krylov runs the Lanczos recurrence from an explicit start vector, +// returning the diagonal alpha, the off-diagonal beta and the basis +// vectors flattened as n-wide rows. Unlike the eigensolver's variant +// it stops at a collapse rather than restarting, because the subspace +// has to stay K_m(A, v) for the projection to answer exp(A)·v; a +// smaller space reached before the budget is exact, not a failure. +func (c *sparseCSR) krylov(start []float64, steps int) (alphas, betas []float64, basis []float64) { + n := c.n + q := append([]float64(nil), start...) + w := make([]float64, n) + // The budget bounds the recurrence: one basis row and at most one + // coefficient per step, so all three are sized once rather than grown. + alphas = make([]float64, 0, steps) + betas = make([]float64, 0, steps) + basis = make([]float64, 0, steps*n) + scale := 0.0 + addScale := func(v float64) { + if a := math.Abs(v); a > scale { + scale = a + } + } + for range steps { + basis = append(basis, q...) + c.matVec(q, w) + alpha := dotF64(q, w) + alphas = append(alphas, alpha) + addScale(alpha) + // Strip the two previous directions out of w, the three-term + // recurrence itself. + for i := range n { + w[i] -= alpha * q[i] + } + if len(alphas) > 1 { + beta := betas[len(betas)-1] + prev := basis[(len(alphas)-2)*n : (len(alphas)-1)*n] + for i := range n { + w[i] -= beta * prev[i] + } + } + // Full reorthogonalisation, twice: one pass removes the + // accumulated loss, the second what the first reintroduces. + for range 2 { + for p := range len(alphas) { + row := basis[p*n : (p+1)*n] + d := dotF64(w, row) + for i := range n { + w[i] -= d * row[i] + } + } + } + beta := norm2F64(w) + // Purely scale-relative exhaustion threshold: an absolute floor + // would truncate the Krylov space of a tiny-norm matrix at + // dimension one and replace the first-order correction with its + // projection, a relative error of order one. + if beta <= float64(n)*base.EpsF*scale { + // The Krylov space is exhausted; the vectors already + // taken span the whole action of A on v. + break + } + betas = append(betas, beta) + addScale(beta) + for i := range n { + q[i] = w[i] / beta + } + } + // The last beta couples to a vector that was never formed, so it + // does not belong on the tridiagonal band. + if len(betas) >= len(alphas) { + betas = betas[:len(alphas)-1] + } + return alphas, betas, basis +} diff --git a/linalg/sparsexp_test.go b/linalg/sparsexp_test.go new file mode 100644 index 0000000..bf8d7aa --- /dev/null +++ b/linalg/sparsexp_test.go @@ -0,0 +1,251 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// denseExpApply forms the whole exponential densely and multiplies it +// by v, the reference the Krylov projection is compared against. +func denseExpApply(t *testing.T, vals, v []float64, n int) []float64 { + t.Helper() + a, err := core.FromFloats(vals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + ea, err := MatrixExp(a) + if err != nil { + t.Fatalf("MatrixExp: %v", err) + } + out := make([]float64, n) + for i := range n { + s := 0.0 + for j := range n { + s += ea.FloatAt(i*n+j) * v[j] + } + out[i] = s + } + return out +} + +// TestSpExpApplyMatchesDense compares the Krylov projection against +// the dense exponential on symmetric matrices, at both a full and a +// truncated Krylov dimension. +func TestSpExpApplyMatchesDense(t *testing.T) { + cases := []struct { + name string + n int + diag float64 + off float64 + steps int + }{ + // n ≤ 40 gets the whole Krylov space, so the answer is exact. + {name: "exact_small", n: 6, diag: 3, off: -1, steps: 0}, + {name: "exact_at_budget", n: 40, diag: 4, off: -1, steps: 0}, + // A truncated dimension on a larger matrix is approximate. + {name: "truncated", n: 100, diag: 5, off: -1, steps: 40}, + // A near-diagonal matrix, where the action is close to + // elementwise and the projection converges immediately. + {name: "weakly_coupled", n: 30, diag: 2, off: -0.01, steps: 0}, + } + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + vals := spdTridiagonal(tt.n, tt.diag, tt.off) + sp := sparseFromDense(t, vals, tt.n) + v := make([]float64, tt.n) + for i := range tt.n { + v[i] = math.Sin(float64(i+1)) * 0.5 + } + vArr, err := core.FromFloats(v, tt.n) + if err != nil { + t.Fatalf("FromFloats v: %v", err) + } + got, err := SpExpApply(sp, vArr, tt.steps) + if err != nil { + t.Fatalf("SpExpApply: %v", err) + } + want := denseExpApply(t, vals, v, tt.n) + for i := range tt.n { + g, w := got.FloatAt(i), want[i] + if math.Abs(g-w) > 1e-8*(1+math.Abs(w)) { + t.Fatalf("element %d = %.12g, want %.12g", i, g, w) + } + } + }) + } +} + +// TestSpExpApplyDiagonal pins the closed form: for a diagonal matrix, +// exp(A)·v is the elementwise exponential of the diagonal times v. +func TestSpExpApplyDiagonal(t *testing.T) { + const n = 5 + diag := []float64{1, 2, -1, 0.5, 3} + vals := make([]float64, n*n) + for i := range n { + vals[i*n+i] = diag[i] + } + sp := sparseFromDense(t, vals, n) + v := []float64{1, 1, 1, 1, 1} + vArr, err := core.FromFloats(v, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := SpExpApply(sp, vArr, 0) + if err != nil { + t.Fatalf("SpExpApply: %v", err) + } + for i := range n { + want := math.Exp(diag[i]) + if math.Abs(got.FloatAt(i)-want) > 1e-12*(1+math.Abs(want)) { + t.Fatalf("element %d = %.12g, want %.12g", i, got.FloatAt(i), want) + } + } +} + +// TestSpExpApplyZeroVector pins the degenerate contract: exp(A)·0 is +// the zero vector, returned without dividing by a zero norm. +func TestSpExpApplyZeroVector(t *testing.T) { + const n = 4 + vals := spdTridiagonal(n, 3, -1) + sp := sparseFromDense(t, vals, n) + v, err := core.Zeros(core.Float, n) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + got, err := SpExpApply(sp, v, 0) + if err != nil { + t.Fatalf("SpExpApply: %v", err) + } + for i := range n { + if got.FloatAt(i) != 0 { + t.Fatalf("element %d = %.12g, want 0 for a zero vector", i, got.FloatAt(i)) + } + } +} + +// TestSpExpApplyDeterminism checks reproducibility: the projection +// starts from v itself and draws nothing random. +func TestSpExpApplyDeterminism(t *testing.T) { + const n = 50 + vals := spdTridiagonal(n, 4, -1) + sp := sparseFromDense(t, vals, n) + v := make([]float64, n) + for i := range n { + v[i] = float64(i+1) / float64(n) + } + vArr, err := core.FromFloats(v, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + g1, err := SpExpApply(sp, vArr, 0) + if err != nil { + t.Fatalf("SpExpApply #1: %v", err) + } + g2, err := SpExpApply(sp, vArr, 0) + if err != nil { + t.Fatalf("SpExpApply #2: %v", err) + } + for i := range n { + if g1.FloatAt(i) != g2.FloatAt(i) { + t.Fatalf("element %d = %.12g vs %.12g across runs", i, g1.FloatAt(i), g2.FloatAt(i)) + } + } +} + +// TestSpExpApplyStepsClamped checks that a step count past the +// dimension is clamped to it, rather than overrunning the space. +func TestSpExpApplyStepsClamped(t *testing.T) { + const n = 4 + vals := spdTridiagonal(n, 3, -1) + sp := sparseFromDense(t, vals, n) + v := []float64{1, 2, 3, 4} + vArr, err := core.FromFloats(v, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := SpExpApply(sp, vArr, 1000) + if err != nil { + t.Fatalf("SpExpApply: %v", err) + } + want := denseExpApply(t, vals, v, n) + for i := range n { + if math.Abs(got.FloatAt(i)-want[i]) > 1e-10*(1+math.Abs(want[i])) { + t.Fatalf("element %d = %.12g, want %.12g", i, got.FloatAt(i), want[i]) + } + } +} + +// TestSpExpApplyRejectsInvalid pins the error contract for every input +// the projection cannot honestly answer. +func TestSpExpApplyRejectsInvalid(t *testing.T) { + vector := func(n int) *core.Array { + v, err := core.FromFloats(make([]float64, n), n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return v + } + t.Run("complex_sparse", func(t *testing.T) { + idx, _ := core.FromInts([]int64{0, 0}, 1, 2) + vals, _ := core.FromComplexes([]complex128{1}, 1) + sp, err := core.NewSparseCOO(idx, vals, []int{1, 1}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := SpExpApply(sp, vector(1), 0); err == nil { + t.Fatal("expected an error for a complex sparse matrix") + } + }) + t.Run("not_square", func(t *testing.T) { + idx, _ := core.FromInts([]int64{0, 0}, 1, 2) + vals, _ := core.FromFloats([]float64{1}, 1) + sp, err := core.NewSparseCOO(idx, vals, []int{1, 2}) + if err != nil { + t.Fatalf("NewSparseCOO: %v", err) + } + if _, err := SpExpApply(sp, vector(1), 0); err == nil { + t.Fatal("expected an error for a non-square matrix") + } + }) + t.Run("zero_sized", func(t *testing.T) { + idx, _ := core.FromInts(nil, 0, 2) + vals, _ := core.FromFloats(nil, 0) + sp := &core.SparseCOO{Indices: idx, Values: vals, Shape: []int{0, 0}} + if _, err := SpExpApply(sp, vector(0), 0); err == nil { + t.Fatal("expected an error for a zero-sized matrix") + } + }) + t.Run("vector_wrong_length", func(t *testing.T) { + vals := spdTridiagonal(3, 3, -1) + sp := sparseFromDense(t, vals, 3) + if _, err := SpExpApply(sp, vector(2), 0); err == nil { + t.Fatal("expected an error for a vector of the wrong length") + } + }) + t.Run("vector_not_rank1", func(t *testing.T) { + vals := spdTridiagonal(2, 3, -1) + sp := sparseFromDense(t, vals, 2) + m, _ := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2) + if _, err := SpExpApply(sp, m, 0); err == nil { + t.Fatal("expected an error for a rank-2 vector") + } + }) + t.Run("complex_vector", func(t *testing.T) { + vals := spdTridiagonal(2, 3, -1) + sp := sparseFromDense(t, vals, 2) + v, _ := core.FromComplexes([]complex128{1, 1}, 2) + if _, err := SpExpApply(sp, v, 0); err == nil { + t.Fatal("expected an error for a complex vector") + } + }) + t.Run("asymmetric", func(t *testing.T) { + sp := sparseFromDense(t, []float64{4, 1, 2, 4}, 2) + if _, err := SpExpApply(sp, vector(2), 0); err == nil { + t.Fatal("expected an error for an asymmetric matrix") + } + }) +} diff --git a/linalg/spline.go b/linalg/spline.go new file mode 100644 index 0000000..20f34c6 --- /dev/null +++ b/linalg/spline.go @@ -0,0 +1,134 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// Natural cubic spline interpolation. Given n points (xᵢ, yᵢ) with +// strictly increasing x, the natural spline is the unique piecewise +// cubic that passes through every point, is C² across the interior +// knots, and has zero second derivative at both ends. The second +// derivatives come from an n−2 tridiagonal system solved by the +// library's `SolveTridiagonal`, so the setup reuses the sparse +// machinery rather than a second LU. + +// CubicSpline holds the natural cubic spline interpolant over sorted +// knots. At returns the interpolated value; Evaluate maps an array of +// query points element-wise. +type CubicSpline struct { + xs, ys, m []float64 // knots, values, second derivatives + n int +} + +// NewCubicSpline builds a natural cubic spline through the given +// points. The abscissae must be strictly increasing; at least 3 +// points are needed for a piecewise cubic to have interior knots. +func NewCubicSpline(xs, ys *core.Array) (*CubicSpline, error) { + n := xs.Len() + if err := requireReal("NewCubicSpline", xs, ys); err != nil { + return nil, err + } + if xs.NDim() != 1 || ys.NDim() != 1 { + return nil, base.Errf("NewCubicSpline: xs and ys must be 1-D, got shapes %s and %s", + base.ShapeText(xs.Shape()), base.ShapeText(ys.Shape())) + } + if n < 3 { + return nil, base.Errf("NewCubicSpline: at least 3 points are needed, got %d", n) + } + if ys.Len() != n { + return nil, base.Errf("NewCubicSpline: xs has %d points but ys has %d", n, ys.Len()) + } + knots := make([]float64, n) + for i := range n { + knots[i] = xs.FloatAt(i) + if i > 0 && knots[i] <= knots[i-1] { + return nil, base.Errf("NewCubicSpline: abscissae must be strictly increasing, found %v after %v", + knots[i], knots[i-1]) + } + } + vals := make([]float64, n) + for i := range n { + vals[i] = ys.FloatAt(i) + } + // Second derivatives via the tridiagonal system for the interior + // knots, with natural (zero) boundary conditions. + h := make([]float64, n-1) + for i := range n - 1 { + h[i] = knots[i+1] - knots[i] + } + m := n - 2 // interior knots: the tridiagonal system is m×m + sub := make([]float64, m-1) + main := make([]float64, m) + sup := make([]float64, m-1) + rhs := make([]float64, m) + for i := range m { + main[i] = 2 * (h[i] + h[i+1]) + if i > 0 { + sub[i-1] = h[i] + } + if i < m-1 { + sup[i] = h[i+1] + } + slope1 := (vals[i+1] - vals[i]) / h[i] + slope2 := (vals[i+2] - vals[i+1]) / h[i+1] + rhs[i] = 6 * (slope2 - slope1) + } + subArr := floatsToArray(sub, []int{m - 1}) + mainArr := floatsToArray(main, []int{m}) + supArr := floatsToArray(sup, []int{m - 1}) + rhsArr := floatsToArray(rhs, []int{m}) + mArr, err := SolveTridiagonal(subArr, mainArr, supArr, rhsArr) + if err != nil { + return nil, base.Errf("NewCubicSpline: %w", err) + } + sec := make([]float64, n) + sec[0], sec[n-1] = 0, 0 + for i := 1; i < n-1; i++ { + sec[i] = mArr.FloatAt(i - 1) + } + return &CubicSpline{xs: knots, ys: vals, m: sec, n: n}, nil +} + +// At evaluates the spline at one query point. Values outside the +// knot range return NaN, mirroring the convention that extrapolation +// without a boundary condition is undefined. +func (s *CubicSpline) At(x float64) float64 { + if x < s.xs[0] || x > s.xs[s.n-1] { + return math.NaN() + } + // Binary search for the interval [x_lo, x_lo+1]. + lo, hi := 0, s.n-1 + for hi-lo > 1 { + mid := (lo + hi) / 2 + if s.xs[mid] <= x { + lo = mid + } else { + hi = mid + } + } + h := s.xs[lo+1] - s.xs[lo] + a := (s.xs[lo+1] - x) / h + b := (x - s.xs[lo]) / h + return a*s.ys[lo] + b*s.ys[lo+1] + + ((a*a*a-a)*s.m[lo]+(b*b*b-b)*s.m[lo+1])*h*h/6 +} + +// Evaluate maps an array of query points through the spline. Like +// NewCubicSpline it reads its argument as real numbers, so a complex +// query array is refused. +func (s *CubicSpline) Evaluate(x *core.Array) (*core.Array, error) { + if err := requireReal("CubicSpline.Evaluate", x); err != nil { + return nil, err + } + out := core.New(core.Float, append([]int{}, x.Shape()...)...) + for i := range x.Len() { + out.RawFloats()[i] = s.At(x.FloatAt(i)) + } + return out, nil +} diff --git a/linalg/spline_test.go b/linalg/spline_test.go new file mode 100644 index 0000000..435e001 --- /dev/null +++ b/linalg/spline_test.go @@ -0,0 +1,109 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "testing" +) + +// TestCubicSplineExactness checks the defining property of the +// spline: it passes exactly through every knot. +func TestCubicSplineExactness(t *testing.T) { + xs := mustFloats(t, []float64{0, 1, 3, 5, 8}) + ys := mustFloats(t, []float64{1, 3, -2, 4, 0}) + s, err := NewCubicSpline(xs, ys) + if err != nil { + t.Fatalf("NewCubicSpline: %v", err) + } + for i := range xs.Len() { + v := s.At(xs.FloatAt(i)) + if math.Abs(v-ys.FloatAt(i)) > 1e-14 { + t.Fatalf("S(%v) = %v, want %v", xs.FloatAt(i), v, ys.FloatAt(i)) + } + } +} + +// TestCubicSplineAgainstSin checks the interpolated values against +// sin on a dense grid between knots: with a few knots the natural +// spline tracks the smooth function to the 1e-5 level. +func TestCubicSplineAgainstSin(t *testing.T) { + knots := []float64{0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.14159265} + xs := mustFloats(t, knots, len(knots)) + ys := make([]float64, len(knots)) + for i := range knots { + ys[i] = math.Sin(knots[i]) + } + ysArr := mustFloats(t, ys, len(knots)) + s, err := NewCubicSpline(xs, ysArr) + if err != nil { + t.Fatalf("NewCubicSpline: %v", err) + } + for _, q := range []float64{0.25, 0.75, 1.25, 1.75, 2.25, 3.0} { + got := s.At(q) + want := math.Sin(q) + if math.Abs(got-want) > 1e-3 { + t.Fatalf("S(%v) = %v, want ≈ %v", q, got, want) + } + } +} + +// TestCubicSplineNatural checks the natural boundary: the second +// derivative at both endpoints must vanish. +func TestCubicSplineNatural(t *testing.T) { + xs := mustFloats(t, []float64{0, 1, 2, 3}) + ys := mustFloats(t, []float64{0, 1, 0, 1}) + s, err := NewCubicSpline(xs, ys) + if err != nil { + t.Fatalf("NewCubicSpline: %v", err) + } + if s.m[0] != 0 || s.m[s.n-1] != 0 { + t.Fatalf("natural boundary: M₀ = %v, M₃ = %v, both want 0", s.m[0], s.m[s.n-1]) + } +} + +// TestCubicSplineEvaluate checks the array-valued entry point and the +// NaN contract outside the knot range. +func TestCubicSplineEvaluate(t *testing.T) { + xs := mustFloats(t, []float64{0, 1, 2, 3, 4, 5}) + ys := mustFloats(t, []float64{0, 1, 0, 1, 0, 1}) + s, err := NewCubicSpline(xs, ys) + if err != nil { + t.Fatalf("NewCubicSpline: %v", err) + } + // Inside: exact at knots. + queries := mustFloats(t, []float64{0, 2, 5}, 3) + got, err := s.Evaluate(queries) + if err != nil { + t.Fatalf("Evaluate: %v", err) + } + for i, xi := range []float64{0, 2, 5} { + if math.Abs(got.FloatAt(i)-ys.FloatAt(int(xi))) > 1e-14 { + t.Fatalf("S(%v) = %v, want %v", xi, got.FloatAt(i), ys.FloatAt(int(xi))) + } + } + // Outside: NaN. + out, _ := s.Evaluate(mustFloats(t, []float64{-1, 10}, 2)) + if !math.IsNaN(out.FloatAt(0)) || !math.IsNaN(out.FloatAt(1)) { + t.Fatalf("outside the knots expected NaN, got %v and %v", out.FloatAt(0), out.FloatAt(1)) + } +} + +// TestCubicSplineRejectsInvalid pins the input contract: non-increasing +// abscissae and a too-short point set are refused. +func TestCubicSplineRejectsInvalid(t *testing.T) { + small := mustFloats(t, []float64{0, 1}) + ys := mustFloats(t, []float64{0, 1}) + if _, err := NewCubicSpline(small, ys); err == nil { + t.Fatal("expected an error for fewer than 3 points") + } + dup := mustFloats(t, []float64{0, 1, 1, 3}) + if _, err := NewCubicSpline(dup, mustFloats(t, []float64{0, 1, 2, 3})); err == nil { + t.Fatal("expected an error for duplicate abscissae") + } + dec := mustFloats(t, []float64{0, 2, 1}) + if _, err := NewCubicSpline(dec, mustFloats(t, []float64{0, 1, 2})); err == nil { + t.Fatal("expected an error for decreasing abscissae") + } +} diff --git a/linalg/svdsolve.go b/linalg/svdsolve.go new file mode 100644 index 0000000..f1d06f7 --- /dev/null +++ b/linalg/svdsolve.go @@ -0,0 +1,138 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Rank-deficient and ill-posed systems through the singular value +// decomposition. Where the QR-based LeastSquares fails outright on a +// rank-deficient matrix, the SVD route scales each singular +// direction's contribution individually: truncation zeroes the +// directions below a chosen rank, Tikhonov damping shrinks every +// direction by σ/(σ²+λ). Both answer x = V·W·Uᵀb over the thin +// factorisation, the minimum-norm least-squares solution of the +// system they solve. + +// SolveTruncated solves A·x = b keeping only the rank largest +// singular values, the truncated pseudoinverse: directions beyond the +// rank, whatever noise they carry, contribute nothing. a is m×n and b +// is m or m×k; a singular value that vanishes to round-off below the +// requested rank is an error, since the system is more deficient than +// the truncation asked for. +func SolveTruncated(a, b *core.Array, rank int) (*core.Array, error) { + return svdSolve(a, b, "SolveTruncated", func(sigma []float64) ([]float64, error) { + if rank < 1 || rank > len(sigma) { + return nil, base.Errf("SolveTruncated: rank must be between 1 and %d, got %d", len(sigma), rank) + } + floor := 64 * base.EpsF * sigma[0] + w := make([]float64, len(sigma)) + for i, s := range sigma { + if i >= rank { + break + } + if s <= floor { + return nil, base.Errf("SolveTruncated: the system is rank-deficient below rank %d (σ%d = %g)", + rank, i, s) + } + w[i] = 1 / s + } + return w, nil + }) +} + +// SolveTikhonov solves the regularised least-squares problem +// min ‖A·x − b‖² + λ‖x‖² by Tikhonov damping with the identity +// prior: every singular direction is scaled by σ/(σ²+λ), which keeps +// ill-conditioned directions from amplifying whatever noise sits on +// b while biasing the answer only where the data carries no +// information. a is m×n and b is m or m×k; a non-positive λ is an +// error. +func SolveTikhonov(a, b *core.Array, lambda float64) (*core.Array, error) { + return svdSolve(a, b, "SolveTikhonov", func(sigma []float64) ([]float64, error) { + if lambda <= 0 { + return nil, base.Errf("SolveTikhonov: lambda must be positive, got %g", lambda) + } + w := make([]float64, len(sigma)) + for i, s := range sigma { + w[i] = s / (s*s + lambda) + } + return w, nil + }) +} + +// svdSolve applies x = V·W·Uᵀb over the thin SVD of a, with W the +// diagonal of per-direction weights the caller chooses. +func svdSolve(a, b *core.Array, name string, weight func(sigma []float64) ([]float64, error)) (*core.Array, error) { + if a.Dtype() == core.Complex { + return nil, base.Errf("%s: complex matrices are not supported", name) + } + // The float fill of b below is the same gap the a gate closes: a + // complex right-hand side has no float payload to read. + if b.Dtype() == core.Complex { + return nil, base.Errf("%s: complex right-hand sides are not supported", name) + } + if a.NDim() != 2 || a.Shape()[0] == 0 || a.Shape()[1] == 0 { + return nil, base.Errf("%s: a must be a non-empty 2-D matrix, got shape %s", name, base.ShapeText(a.Shape())) + } + if b.NDim() != 1 && b.NDim() != 2 { + return nil, base.Errf("%s: b must be 1-D or 2-D, got shape %s", name, base.ShapeText(b.Shape())) + } + m, n := a.Shape()[0], a.Shape()[1] + if b.Shape()[0] != m { + return nil, base.Errf("%s: b rows (%d) must match a rows (%d)", name, b.Shape()[0], m) + } + cols := 1 + if b.NDim() == 2 { + cols = b.Shape()[1] + } + r := min(n, m) + u, sigma, vt, err := SVD(a) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + sVals := make([]float64, r) + for i := range r { + sVals[i] = sigma.RawFloats()[i] // SVD returns a dense float array + } + w, err := weight(sVals) + if err != nil { + return nil, err + } + // c = Uᵀb (r×cols), damped row-wise, then x = Vᵀ·d. + bFlat := make([]float64, m*cols) + if b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == m*cols { + copy(bFlat, b.RawFloats()) + } else { + for i := range m { + for j := range cols { + bFlat[i*cols+j] = b.FloatAt(i*cols + j) + } + } + } + uFlat := denseFloats(u, m, r) + vtFlat := denseFloats(vt, r, n) + x := make([]float64, n*cols) + for i := range r { + if w[i] == 0 { + continue + } + for j := range cols { + c := 0.0 + for l := range m { + c += uFlat[l*r+i] * bFlat[l*cols+j] + } + c *= w[i] + for jj := range n { + x[jj*cols+j] += vtFlat[i*n+jj] * c + } + } + } + if b.NDim() == 1 { + return floatsToArray(x, []int{n}), nil + } + return floatsToArray(x, []int{n, cols}), nil +} diff --git a/linalg/svdsolve_test.go b/linalg/svdsolve_test.go new file mode 100644 index 0000000..63fcbf3 --- /dev/null +++ b/linalg/svdsolve_test.go @@ -0,0 +1,180 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// diagonalSpectrum returns the n×n diagonal matrix with the given +// singular values, the cleanest stage for checking what truncation +// and damping do per direction. +func diagonalSpectrum(t *testing.T, sigma []float64) *core.Array { + t.Helper() + vals := make([]float64, len(sigma)*len(sigma)) + for i, s := range sigma { + vals[i*len(sigma)+i] = s + } + a, err := core.FromFloats(vals, len(sigma), len(sigma)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestSolveTruncatedFullRank checks the untruncated case against the +// plain solve on a well-conditioned system. +func TestSolveTruncatedFullRank(t *testing.T) { + a := mustFloats(t, []float64{3, 0, 1, 1, 2, 1, 1, 1, 2, 1, 4, 0, 0, 1, 1, 5}, 4, 4) + b := mustFloats(t, []float64{1, 2, 3, 4}, 4) + x, err := SolveTruncated(a, b, 4) + if err != nil { + t.Fatalf("SolveTruncated: %v", err) + } + want, err := Solve(a, b) + if err != nil { + t.Fatalf("Solve: %v", err) + } + for i := range 4 { + if math.Abs(x.FloatAt(i)-want.FloatAt(i)) > 1e-8 { + t.Fatalf("x[%d] = %.12g, want %.12g", i, x.FloatAt(i), want.FloatAt(i)) + } + } +} + +// TestSolveTruncatedCutsNoise shows the point of truncation: with a +// spectrum of 10, 1, 1e-8, 1e-12 and b = 1 everywhere, the full solve +// amplifies b by 1e12 while rank-2 truncation answers (0.1, 1, 0, 0). +func TestSolveTruncatedCutsNoise(t *testing.T) { + a := diagonalSpectrum(t, []float64{10, 1, 1e-8, 1e-12}) + b := mustFloats(t, []float64{1, 1, 1, 1}, 4) + x, err := SolveTruncated(a, b, 2) + if err != nil { + t.Fatalf("SolveTruncated: %v", err) + } + want := []float64{0.1, 1, 0, 0} + for i := range 4 { + if math.Abs(x.FloatAt(i)-want[i]) > 1e-12 { + t.Fatalf("x[%d] = %.12g, want %.12g", i, x.FloatAt(i), want[i]) + } + } +} + +// TestSolveTikhonovDiagonal pins the per-direction damping factor +// σ/(σ²+λ) on the diagonal system: large directions pass nearly +// untouched, tiny directions are crushed. +func TestSolveTikhonovDiagonal(t *testing.T) { + const lambda = 1.0 + a := diagonalSpectrum(t, []float64{10, 1, 1e-8, 1e-12}) + b := mustFloats(t, []float64{1, 1, 1, 1}, 4) + x, err := SolveTikhonov(a, b, lambda) + if err != nil { + t.Fatalf("SolveTikhonov: %v", err) + } + for i, s := range []float64{10, 1, 1e-8, 1e-12} { + want := s / (s*s + lambda) + if math.Abs(x.FloatAt(i)-want) > 1e-12*math.Max(1, want) { + t.Fatalf("x[%d] = %.14g, want %.14g", i, x.FloatAt(i), want) + } + } +} + +// TestSolveTikhonovNormalEquations cross-checks a rectangular +// problem against the normal equations (AᵀA + λI)x = Aᵀb solved by +// the plain dense solver. +func TestSolveTikhonovNormalEquations(t *testing.T) { + const lambda = 0.7 + a := mustFloats(t, []float64{2, 0, 1, 1, 0, 3, 1, 1, 1, 4, 2, 0}, 4, 3) + b := mustFloats(t, []float64{1, 0, 2, 1}, 4) + x, err := SolveTikhonov(a, b, lambda) + if err != nil { + t.Fatalf("SolveTikhonov: %v", err) + } + // Normal equations, assembled row-major for Solve. + ata := make([]float64, 9) + atb := make([]float64, 3) + for i := range 3 { + for j := range 3 { + s := 0.0 + for l := range 4 { + s += a.FloatAt(l*3+i) * a.FloatAt(l*3+j) + } + ata[i*3+j] = s + if i == j { + ata[i*3+j] += lambda + } + } + s := 0.0 + for l := range 4 { + s += a.FloatAt(l*3+i) * b.FloatAt(l) + } + atb[i] = s + } + want, err := Solve(mustFloats(t, ata, 3, 3), mustFloats(t, atb, 3)) + if err != nil { + t.Fatalf("Solve: %v", err) + } + for i := range 3 { + if math.Abs(x.FloatAt(i)-want.FloatAt(i)) > 1e-8 { + t.Fatalf("x[%d] = %.12g, want %.12g", i, x.FloatAt(i), want.FloatAt(i)) + } + } +} + +// TestSolveSVDMultiColumn checks a two-column right-hand side against +// two single-column solves. +func TestSolveSVDMultiColumn(t *testing.T) { + a := mustFloats(t, []float64{3, 0, 1, 1, 2, 1, 1, 1, 2, 1, 4, 0}, 4, 3) + b2 := mustFloats(t, []float64{1, 2, 3, 4, 4, 3, 2, 1}, 4, 2) + x, err := SolveTikhonov(a, b2, 0.5) + if err != nil { + t.Fatalf("SolveTikhonov: %v", err) + } + if x.NDim() != 2 || x.Shape()[0] != 3 || x.Shape()[1] != 2 { + t.Fatalf("solution shape %v, want [3 2]", x.Shape()) + } + for j := range 2 { + col := make([]float64, 4) + for i := range 4 { + col[i] = b2.FloatAt(i*2 + j) + } + single, err := SolveTikhonov(a, mustFloats(t, col, 4), 0.5) + if err != nil { + t.Fatalf("SolveTikhonov column %d: %v", j, err) + } + for i := range 3 { + if math.Abs(x.FloatAt(i*2+j)-single.FloatAt(i)) > 1e-9 { + t.Fatalf("column %d, row %d: %.12g, want %.12g", j, i, x.FloatAt(i*2+j), single.FloatAt(i)) + } + } + } +} + +// TestSolveSVDErrors pins the validation contract of both solvers. +func TestSolveSVDErrors(t *testing.T) { + a := diagonalSpectrum(t, []float64{10, 1, 1e-8, 1e-12}) + b := mustFloats(t, []float64{1, 1, 1, 1}, 4) + if _, err := SolveTruncated(a, b, 0); err == nil { + t.Fatal("expected an error for rank 0") + } + if _, err := SolveTruncated(a, b, 5); err == nil { + t.Fatal("expected an error for rank beyond the spectrum") + } + if _, err := SolveTikhonov(a, b, 0); err == nil { + t.Fatal("expected an error for a non-positive lambda") + } + shortB := mustFloats(t, []float64{1, 1}, 2) + if _, err := SolveTruncated(a, shortB, 2); err == nil { + t.Fatal("expected an error for mismatched b rows") + } + if _, err := SolveTikhonov(shortB, b, 1); err == nil { + t.Fatal("expected an error for a rank-1 matrix") + } + singular := mustFloats(t, []float64{1, 1, 1, 1, 1, 1, 1, 1, 1}, 3, 3) + if _, err := SolveTruncated(singular, mustFloats(t, []float64{1, 1, 1}, 3), 3); err == nil { + t.Fatal("expected an error for a system rank-deficient below the requested rank") + } +} diff --git a/linalg/tridiag.go b/linalg/tridiag.go new file mode 100644 index 0000000..b171a33 --- /dev/null +++ b/linalg/tridiag.go @@ -0,0 +1,141 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Tridiagonal solvers. The Thomas algorithm and its periodic +// (Sherman-Morrison) variant solve the tridiagonal systems that +// discretised one-dimensional differential operators produce, in O(n) +// time and O(n) memory, versus the O(n³) dense `Solve`. All inputs +// are rank-1 vectors; the sub-, main- and superdiagonals are separate +// vectors of lengths n−1, n and n−1. + +// SolveTridiagonal returns the vector x solving the tridiagonal system +// with lower diagonal a (length n−1), main diagonal b (length n), +// upper diagonal c (length n−1) and right-hand side d (length n), +// by the Thomas algorithm. A zero pivot is refused. +func SolveTridiagonal(a, b, c, d *core.Array) (*core.Array, error) { + n := b.Len() + if err := requireReal("SolveTridiagonal", a, b, c, d); err != nil { + return nil, err + } + if n == 0 { + return nil, base.Errf("SolveTridiagonal: empty system") + } + if a.Len() != n-1 || c.Len() != n-1 || d.Len() != n { + return nil, base.Errf("SolveTridiagonal: diagonal lengths %d/%d/%d do not match rhs %d", + a.Len(), b.Len(), c.Len(), d.Len()) + } + // The elimination is internal/base's shared kernel, the same one + // the integrate package's PDE sweeps run against their scratch. A + // dense float64 operand hands the kernel its payload directly, the + // same values FloatAt read; every other real dtype reaches the + // kernel through the FloatAt promotion the element-wise walk + // applied. + av, bv, cv, dv := triDiagonals(a, b, c, d, n) + cp := make([]float64, n) + dp := make([]float64, n) + x := make([]float64, n) + if err := base.TriSolve(x, cp, dp, av, bv, cv, dv); err != nil { + return nil, err + } + return floatsToArray(x, []int{n}), nil +} + +// triDiagonals materialises the four dense float64 operands the shared +// kernel reads. A dense float64 array's payload slice is exactly its +// visible range, so it is passed through untouched; any other real +// dtype is promoted element by element, the conversion FloatAt +// performs. A strided operand takes the promotion path, where the +// accessor walks the logical order. +func triDiagonals(a, b, c, d *core.Array, n int) (av, bv, cv, dv []float64) { + extract := func(arr *core.Array, m int) []float64 { + if arr.Dtype() == core.Float && !arr.Strided() { + return arr.RawFloats()[:m] + } + out := make([]float64, m) + for i := range m { + out[i] = arr.FloatAt(i) + } + return out + } + return extract(a, n-1), extract(b, n), extract(c, n-1), extract(d, n) +} + +// SolveCyclicTridiagonal returns the vector x solving the periodic +// tridiagonal system whose corner entries couple the ends: a[0] is +// the wrap corner A[0][n−1], c[n−1] is the wrap corner A[n−1][0], and +// a[1..n−1], b[0..n−1], c[0..n−2] fill the interior bands. Built on +// the Sherman-Morrison rank-one update of the Thomas elimination, +// still O(n). +func SolveCyclicTridiagonal(a, b, c, d *core.Array) (*core.Array, error) { + n := b.Len() + if err := requireReal("SolveCyclicTridiagonal", a, b, c, d); err != nil { + return nil, err + } + if n < 3 { + return nil, base.Errf("SolveCyclicTridiagonal: needs at least 3 unknowns, got %d", n) + } + if a.Len() != n || c.Len() != n || d.Len() != n { + return nil, base.Errf("SolveCyclicTridiagonal: diagonal lengths %d/%d/%d do not match rhs %d", + a.Len(), b.Len(), c.Len(), d.Len()) + } + alpha := a.FloatAt(0) + beta := c.FloatAt(n - 1) + gamma := -b.FloatAt(0) + if gamma == 0 { + return nil, base.Errf("SolveCyclicTridiagonal: gamma = −b[0] vanishes, system is singular") + } + + // Strip the rank-one correction: the interior tridiagonal T has + // b′[0] = b[0] − γ, b′[n−1] = b[n−1] − αβ/γ, no wrap corners. + bb := make([]float64, n) + for i := range n { + bb[i] = b.FloatAt(i) + } + bb[0] -= gamma + bb[n-1] -= alpha * beta / gamma + aInt := make([]float64, n-1) + cInt := make([]float64, n-1) + for i := 1; i < n; i++ { + aInt[i-1] = a.FloatAt(i) + } + for i := range n - 1 { + cInt[i] = c.FloatAt(i) + } + aArr := floatsToArray(aInt, []int{n - 1}) + cArr := floatsToArray(cInt, []int{n - 1}) + bArr := floatsToArray(bb, []int{n}) + + // Two Thomas solves against the modified matrix: one for the + // right-hand side, one for the correction vector u. + ym, err := SolveTridiagonal(aArr, bArr, cArr, d) + if err != nil { + return nil, base.Errf("SolveCyclicTridiagonal: %w", err) + } + uVec := make([]float64, n) + uVec[0] = gamma + uVec[n-1] = beta + uArr := floatsToArray(uVec, []int{n}) + zm, err := SolveTridiagonal(aArr, bArr, cArr, uArr) + if err != nil { + return nil, base.Errf("SolveCyclicTridiagonal: %w", err) + } + // v = (1, 0, …, 0, α/γ): x = y − z·(vᵀy)/(1 + vᵀz). + vty := ym.FloatAt(0) + alpha/gamma*ym.FloatAt(n-1) + vtz := zm.FloatAt(0) + alpha/gamma*zm.FloatAt(n-1) + if 1+vtz == 0 { + return nil, base.Errf("SolveCyclicTridiagonal: singular Sherman-Morrison denominator") + } + fact := vty / (1 + vtz) + x := make([]float64, n) + for i := range n { + x[i] = ym.FloatAt(i) - fact*zm.FloatAt(i) + } + return floatsToArray(x, []int{n}), nil +} diff --git a/linalg/tridiag_test.go b/linalg/tridiag_test.go new file mode 100644 index 0000000..c5d6f19 --- /dev/null +++ b/linalg/tridiag_test.go @@ -0,0 +1,167 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package linalg + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestSolveTridiagonal matches the Thomas solve against the dense +// solver on a representative system. +func TestSolveTridiagonal(t *testing.T) { + // Tridiagonal with b = 4 on the main diagonal, ±1 off it. + n := 6 + b := make([]float64, n) + a := make([]float64, n-1) + c := make([]float64, n-1) + d := make([]float64, n) + for i := range n { + b[i] = 4 + d[i] = float64(i + 1) + if i > 0 { + a[i-1] = -1 + } + if i < n-1 { + c[i] = 1 + } + } + x, err := SolveTridiagonal(floatsToArray(a, []int{n - 1}), floatsToArray(b, []int{n}), + floatsToArray(c, []int{n - 1}), floatsToArray(d, []int{n})) + if err != nil { + t.Fatalf("SolveTridiagonal: %v", err) + } + full := make([]float64, n*n) + for i := range n { + full[i*n+i] = 4 + if i > 0 { + full[i*n+i-1] = -1 + } + if i < n-1 { + full[i*n+i+1] = 1 + } + } + ref, err := Solve(mustFloats(t, full, n, n), mustFloats(t, d, n)) + if err != nil { + t.Fatalf("Solve: %v", err) + } + for i := range n { + if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-12 { + t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i)) + } + } + if _, err := SolveTridiagonal(floatsToArray(a, []int{n - 1}), floatsToArray(make([]float64, n), []int{n}), + floatsToArray(c, []int{n - 1}), floatsToArray(d, []int{n})); err == nil { + t.Fatal("zero pivot: want an error") + } +} + +// TestSolveTridiagonalOneByOne pins the 1×1 system: the c and a +// diagonals are empty per the length contract, and the seed must not +// read them (it used to panic on c.FloatAt(0)). +func TestSolveTridiagonalOneByOne(t *testing.T) { + x, err := SolveTridiagonal(floatsToArray(nil, []int{0}), mustFloats(t, []float64{2}), + floatsToArray(nil, []int{0}), mustFloats(t, []float64{6})) + if err != nil { + t.Fatalf("SolveTridiagonal 1×1: %v", err) + } + if math.Abs(x.FloatAt(0)-3) > 1e-14 { + t.Fatalf("x[0] = %v, want 3", x.FloatAt(0)) + } + // A 1×1 system with a zero main-diagonal entry is refused, not + // divided through. + if _, err := SolveTridiagonal(floatsToArray(nil, []int{0}), mustFloats(t, []float64{0}), + floatsToArray(nil, []int{0}), mustFloats(t, []float64{6})); err == nil { + t.Fatal("1×1 zero pivot: want an error") + } +} + +// TestNewCubicSplineThreePoints regresses the same seed panic through +// its smallest caller: three knots build a 1×1 interior system. +func TestNewCubicSplineThreePoints(t *testing.T) { + xs := mustFloats(t, []float64{0, 1, 2}) + ys := mustFloats(t, []float64{0, 1, 0}) + s, err := NewCubicSpline(xs, ys) + if err != nil { + t.Fatalf("NewCubicSpline with 3 points: %v", err) + } + for i := range 3 { + if v := s.At(xs.FloatAt(i)); math.Abs(v-ys.FloatAt(i)) > 1e-14 { + t.Fatalf("S(%v) = %v, want %v", xs.FloatAt(i), v, ys.FloatAt(i)) + } + } +} + +// cyclicApply multiplies the cyclic tridiagonal system out: a[0] is +// the A[0][n−1] wrap corner, c[n−1] the A[n−1][0] corner. +func cyclicApply(a, b, c, x []float64) []float64 { + n := len(b) + out := make([]float64, n) + for i := range n { + out[i] = b[i] * x[i] + if i > 0 { + out[i] += a[i] * x[i-1] + } + if i < n-1 { + out[i] += c[i] * x[i+1] + } + } + out[0] += a[0] * x[n-1] + out[n-1] += c[n-1] * x[0] + return out +} + +// TestSolveCyclicTridiagonalAsymmetricCorners pins the Sherman-Morrison +// corner assignment: the correction vector u ends with beta (= c[n−1]), +// not alpha (= a[0]). With alpha ≠ beta the swapped corner used to +// return a silently wrong answer. +func TestSolveCyclicTridiagonalAsymmetricCorners(t *testing.T) { + a := []float64{7, 2, 1, 3} + b := []float64{10, 11, 12, 13} + c := []float64{1, 4, 2, 5} + d := []float64{1, 2, 3, 4} + x, err := SolveCyclicTridiagonal(mustFloats(t, a), mustFloats(t, b), mustFloats(t, c), mustFloats(t, d)) + if err != nil { + t.Fatalf("SolveCyclicTridiagonal: %v", err) + } + got := cyclicApply(a, b, c, []float64{x.FloatAt(0), x.FloatAt(1), x.FloatAt(2), x.FloatAt(3)}) + for i := range len(d) { + if math.Abs(got[i]-d[i]) > 1e-9 { + t.Fatalf("(A·x)[%d] = %.16g, want %.16g", i, got[i], d[i]) + } + } +} + +// TestSolveCyclicTridiagonalSymmetric covers the classic periodic +// second-difference system, whose wrap corners are equal. +func TestSolveCyclicTridiagonalSymmetric(t *testing.T) { + n := 5 + a := make([]float64, n) + b := make([]float64, n) + c := make([]float64, n) + d := make([]float64, n) + for i := range n { + a[i], b[i], c[i] = -1, 3, -1 + d[i] = float64(i+1) * float64(i) + } + x, err := SolveCyclicTridiagonal(mustFloats(t, a), mustFloats(t, b), mustFloats(t, c), mustFloats(t, d)) + if err != nil { + t.Fatalf("SolveCyclicTridiagonal: %v", err) + } + got := cyclicApply(a, b, c, mustSlice(x, n)) + for i := range n { + if math.Abs(got[i]-d[i]) > 1e-9 { + t.Fatalf("(A·x)[%d] = %.16g, want %.16g", i, got[i], d[i]) + } + } +} + +func mustSlice(x *core.Array, n int) []float64 { + out := make([]float64, n) + for i := range n { + out[i] = x.FloatAt(i) + } + return out +} diff --git a/optim/bench_leastsqfit_test.go b/optim/bench_leastsqfit_test.go new file mode 100644 index 0000000..0610105 --- /dev/null +++ b/optim/bench_leastsqfit_test.go @@ -0,0 +1,86 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the fit surface of LevenbergMarquardtFit: the weighted +// residual path and the covariance report, the two options the plain +// entry point and its benchmarks never touch. + +// benchExponentialFit builds the two-parameter decay fit the older LM +// benchmarks run, with the observations and their variances handed +// back so the weighted and covariance variants measure the same model. +func benchExponentialFit() (residual func(*core.Array) (*core.Array, error), jacobian func(*core.Array) (*core.Array, error), p0 *core.Array) { + const nObs = 40 + t := make([]float64, nObs) + y := make([]float64, nObs) + for i := range nObs { + t[i] = float64(i) / 4 + y[i] = 2.5*math.Exp(-0.7*t[i]) + 0.02*math.Sin(float64(i)) + } + residual = func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, nObs) + vals := out.RawFloats() + for i := range nObs { + vals[i] = p.FloatAt(0)*math.Exp(-p.FloatAt(1)*t[i]) - y[i] + } + return out, nil + } + jacobian = func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, nObs, 2) + vals := out.RawFloats() + for i := range nObs { + e := math.Exp(-p.FloatAt(1) * t[i]) + vals[i*2] = e + vals[i*2+1] = -p.FloatAt(0) * t[i] * e + } + return out, nil + } + p0, _ = core.FromFloats([]float64{1, 0.2}, 2) + return residual, jacobian, p0 +} + +// BenchmarkLevenbergMarquardtFitWeighted measures the fit with one +// variance per residual: the whitening of the residual and of the +// finite-difference Jacobian rides every evaluation the run makes. +func BenchmarkLevenbergMarquardtFitWeighted(b *testing.B) { + residual, _, p0 := benchExponentialFit() + const nObs = 40 + variance := make([]float64, nObs) + for i := range nObs { + variance[i] = 1 + float64(i)/8 + } + sigma, err := core.FromFloats(variance, nObs) + if err != nil { + b.Fatal(err) + } + opts := LMOptions{MaxIterations: 30, Sigma: sigma} + b.ReportAllocs() + for b.Loop() { + if _, err := LevenbergMarquardtFit(residual, p0, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkLevenbergMarquardtFitCovariance measures the covariance +// report: the fit converges and then rebuilds the Jacobian at the +// answer, forms the normal equations from it and solves the identity +// through them. +func BenchmarkLevenbergMarquardtFitCovariance(b *testing.B) { + residual, jacobian, p0 := benchExponentialFit() + opts := LMOptions{MaxIterations: 30, Jacobian: jacobian, RequestCovariance: true} + b.ReportAllocs() + for b.Loop() { + if _, err := LevenbergMarquardtFit(residual, p0, opts); err != nil { + b.Fatal(err) + } + } +} diff --git a/optim/bench_solvers_test.go b/optim/bench_solvers_test.go new file mode 100644 index 0000000..748d1fe --- /dev/null +++ b/optim/bench_solvers_test.go @@ -0,0 +1,370 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "math/rand/v2" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the solver iterations whose cost is dominated by +// repeated dense linear algebra or by a difference stencil: the revised +// simplex, the active-set QP, CMA-ES, L-BFGS and the finite-difference +// paths. Every input comes from a fixed seed, so each run walks one +// deterministic trajectory. + +// solversRNG returns a generator with a fixed stream: the inputs are +// identical on every machine and every run. +func solversRNG() *rand.Rand { return rand.New(rand.NewPCG(0x5eed, 0x1234)) } + +// solverArray builds a float array, panicking on a bad shape: every +// caller passes a literal shape. +func solverArray(vals []float64, shape ...int) *core.Array { + a, err := core.FromFloats(vals, shape...) + if err != nil { + panic("solverArray: " + err.Error()) + } + return a +} + +// lpProblem builds a standard-form LP with m rows and n = 2m columns: +// A = [I | R] with every entry of R at least 0.5, b = 1 and a random +// cost. The identity block makes x = (1, …, 1, 0, …, 0) feasible, and +// the feasible set is bounded: the slack block forces R·x_R ≤ 1, whose +// positive coefficients bound the 1-norm of x_R by 2, and x_I = 1 − +// R·x_R is bounded with it. The run therefore ends at a vertex rather +// than on an unbounded ray. +func lpProblem(m int, rng *rand.Rand) (c, a, b *core.Array) { + n := 2 * m + av := make([]float64, m*n) + for i := range m { + av[i*n+i] = 1 + for j := range m { + av[i*n+m+j] = 0.5 + rng.Float64() + } + } + cv := make([]float64, n) + for j := range n { + cv[j] = 2*rng.Float64() - 1 + } + bv := make([]float64, m) + for i := range m { + bv[i] = 1 + } + return solverArray(cv, n), solverArray(av, m, n), solverArray(bv, m) +} + +// qpProblem builds a strictly convex quadratic in n variables with r +// two-sided rows that all admit the origin, so the start point is +// feasible and the benchmark measures the active-set iteration itself. +func qpProblem(n, r int, rng *rand.Rand) (*core.Array, *core.Array, LinearConstraints) { + hv := make([]float64, n*n) + for i := range n { + hv[i*n+i] = 1 + rng.Float64() + for j := i + 1; j < n; j++ { + v := 0.25 * (2*rng.Float64() - 1) + hv[i*n+j], hv[j*n+i] = v, v + } + } + cv := make([]float64, n) + for j := range n { + cv[j] = 2*rng.Float64() - 1 + } + av := make([]float64, r*n) + lower := make([]float64, r) + upper := make([]float64, r) + for i := range r { + for j := range n { + av[i*n+j] = 2*rng.Float64() - 1 + } + lower[i], upper[i] = math.Inf(-1), 0.05+0.1*rng.Float64() + } + cons := LinearConstraints{A: solverArray(av, r, n), Lower: lower, Upper: upper} + return solverArray(hv, n, n), solverArray(cv, n), cons +} + +// solverBowl returns a coupled bowl around (0.5, …, 0.5) and a start +// point away from it. +func solverBowl(n int, rng *rand.Rand) (f func(*core.Array) (float64, error), x0 *core.Array) { + weights := make([]float64, n) + for i := range n { + weights[i] = 1 + 2*rng.Float64() + } + f = func(p *core.Array) (float64, error) { + total := 0.0 + for i := range n { + d := p.FloatAt(i) - 0.5 + total += weights[i] * d * d + if i+1 < n { + total += 0.3 * d * (p.FloatAt(i+1) - 0.5) + } + } + return total, nil + } + start := make([]float64, n) + for i := range start { + start[i] = 1.5 + 0.1*float64(i) + } + return f, solverArray(start, n) +} + +// BenchmarkSolversLinearProgram measures the two-phase revised simplex +// on a bounded LP with 60 rows and 120 columns, the shape the per-pivot +// refactorisation pays for. +func BenchmarkSolversLinearProgram(b *testing.B) { + rng := solversRNG() + c, a, rhs := lpProblem(60, rng) + opts := LinearProgramOptions{} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseLinear(c, a, rhs, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversLinearProgramRows measures an LP of the same family +// through the two-sided-row wrapper, whose conversion builds the +// standard form before the same simplex runs. The equality rows carry +// the LP itself and every variable is boxed, so the feasible set the +// wrapper sees is bounded. +func BenchmarkSolversLinearProgramRows(b *testing.B) { + rng := solversRNG() + const m = 8 + c, a, rhs := lpProblem(m, rng) + n := c.Len() + rows := m + n + av := make([]float64, rows*n) + for i := range m { + for j := range n { + av[i*n+j] = a.FloatAt(i*n + j) + } + } + lower := make([]float64, rows) + upper := make([]float64, rows) + for i := range m { + lower[i], upper[i] = rhs.FloatAt(i), rhs.FloatAt(i) + } + for i := range n { + av[(m+i)*n+i] = 1 + lower[m+i], upper[m+i] = -3, 3 + } + cons := LinearConstraints{A: solverArray(av, rows, n), Lower: lower, Upper: upper} + opts := LinearProgramOptions{} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseLinearRows(c, cons, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversQuadraticProgram measures the active-set QP on a +// strictly convex quadratic with 10 variables and 24 two-sided rows, +// tight enough that the working set moves several times. +func BenchmarkSolversQuadraticProgram(b *testing.B) { + rng := solversRNG() + h, c, cons := qpProblem(12, 30, rng) + x0 := solverArray(make([]float64, 12), 12) + opts := QPOptions{} + b.ReportAllocs() + for b.Loop() { + if _, _, _, err := MinimiseQP(h, c, cons, x0, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversCMAES measures the covariance-adaptation strategy on +// a four-dimensional bowl, one generation pair being cheap at that size. +func BenchmarkSolversCMAES(b *testing.B) { + rng := solversRNG() + f, x0 := solverBowl(4, rng) + opts := CMAESOptions{Sigma0: 0.5, Generations: 60, Tolerance: 1e-12, Seed: 7, AllowBudgetExit: true} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseCMAES(f, x0, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversLBFGSLeastSquares fits a six-term cosine series to +// samples of a fixed function with a consistent analytic gradient, the +// path whose history buffer rotates once the memory is full. +func BenchmarkSolversLBFGSLeastSquares(b *testing.B) { + const nPar, nObs = 6, 64 + truth := make([]float64, nPar) + for j := range truth { + truth[j] = 1 / float64(j+1) + } + tSamples := make([]float64, nObs) + y := make([]float64, nObs) + for i := range nObs { + t := 0.5 * float64(i) / float64(nObs) + tSamples[i] = t + for j := range nPar { + y[i] += truth[j] * math.Cos(float64(j)*t) + } + } + f := func(p *core.Array) (float64, error) { + total := 0.0 + for i := range nObs { + model := 0.0 + for j := range nPar { + model += p.FloatAt(j) * math.Cos(float64(j)*tSamples[i]) + } + d := model - y[i] + total += d * d + } + return total, nil + } + grad := func(p *core.Array) (*core.Array, error) { + g := make([]float64, nPar) + for i := range nObs { + model := 0.0 + for j := range nPar { + model += p.FloatAt(j) * math.Cos(float64(j)*tSamples[i]) + } + d := 2 * (model - y[i]) + for j := range nPar { + g[j] += d * math.Cos(float64(j)*tSamples[i]) + } + } + return core.FromFloats(g, nPar) + } + x0 := solverArray(make([]float64, nPar), nPar) + opts := LBFGSOptions{MaxIterations: 300, Memory: 4} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseLBFGS(f, grad, x0, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversLBFGSFiniteDiff measures the central-difference +// gradient path: two objective evaluations per coordinate per step. +func BenchmarkSolversLBFGSFiniteDiff(b *testing.B) { + rng := solversRNG() + f, x0 := solverBowl(24, rng) + opts := LBFGSOptions{MaxIterations: 60} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseLBFGS(f, nil, x0, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversLevenbergFiniteDiff measures Levenberg-Marquardt on a +// small polynomial fit with the Jacobian by central differences. +func BenchmarkSolversLevenbergFiniteDiff(b *testing.B) { + const nObs, nPar = 40, 6 + t := make([]float64, nObs) + obs := make([]float64, nObs) + truth := make([]float64, nPar) + for j := range truth { + truth[j] = 0.5 + 0.1*float64(j) + } + for i := range nObs { + t[i] = float64(i) / 8 + v := 0.0 + for j := range nPar { + v += truth[j] * math.Pow(t[i], float64(j)) + } + obs[i] = v + } + residual := func(p *core.Array) (*core.Array, error) { + r := make([]float64, nObs) + for i := range nObs { + v := 0.0 + for j := range nPar { + v += p.FloatAt(j) * math.Pow(t[i], float64(j)) + } + r[i] = v - obs[i] + } + return core.FromFloats(r, nObs) + } + p0 := solverArray(make([]float64, nPar), nPar) + opts := LMOptions{MaxIterations: 20} + b.ReportAllocs() + for b.Loop() { + if _, _, err := LevenbergMarquardt(residual, p0, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversRootSystemFiniteDiff measures the damped Newton +// iteration with a per-step central-difference Jacobian. +func BenchmarkSolversRootSystemFiniteDiff(b *testing.B) { + const n = 8 + r := func(x *core.Array) (*core.Array, error) { + out := make([]float64, n) + for i := range n { + v := x.FloatAt(i) + out[i] = v*v + 0.1*v - float64(i+1) + if i+1 < n { + out[i] += 0.05 * x.FloatAt(i+1) + } + } + return core.FromFloats(out, n) + } + start := make([]float64, n) + for i := range start { + start[i] = 1 + 0.1*float64(i) + } + x0 := solverArray(start, n) + opts := RootSystemOptions{MaxIterations: 20} + b.ReportAllocs() + for b.Loop() { + if _, _, err := FindRootSystem(r, x0, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolversNonlinearConstrained measures the augmented-Lagrangian +// outer loop over one equality and two inequality rows, whose stencil is +// the per-row, per-coordinate hot path. +func BenchmarkSolversNonlinearConstrained(b *testing.B) { + const n = 6 + f := func(p *core.Array) (float64, error) { return p.FloatAt(0), nil } + cons := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { + s := -1.0 + for i := range n { + s += p.FloatAt(i) * p.FloatAt(i) + } + return s, nil + }, + }, + Inequalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { + s := 0.0 + for i := range n { + s += p.FloatAt(i) + } + return s - 2, nil + }, + func(p *core.Array) (float64, error) { return -p.FloatAt(0) - 2, nil }, + }, + } + start := make([]float64, n) + for i := range start { + start[i] = 0.5 + } + x0 := solverArray(start, n) + b.ReportAllocs() + for b.Loop() { + if _, _, _, err := MinimiseNonlinearConstrained(f, nil, x0, cons, LBFGSOptions{}); err != nil { + b.Fatal(err) + } + } +} diff --git a/optim/bench_test.go b/optim/bench_test.go new file mode 100644 index 0000000..3d919e2 --- /dev/null +++ b/optim/bench_test.go @@ -0,0 +1,187 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the optimiser hot paths: the L-BFGS two-loop recursion +// with analytic and finite-difference gradients, the simplex method, +// Levenberg-Marquardt's normal equations and the damped Newton system +// solver. + +// benchQuadratic builds a separable convex quadratic +// f(x) = Σ (xᵢ − cᵢ)² + 0.01·Σ xᵢ² with cᵢ = i/n, whose minimum and +// gradient are closed form, so the L-BFGS runs are identical every +// iteration. +func benchQuadratic(n int) (f func(*core.Array) (float64, error), grad func(*core.Array) (*core.Array, error), x0 []float64) { + c := make([]float64, n) + for i := range c { + c[i] = float64(i) / float64(n) + } + f = func(a *core.Array) (float64, error) { + total := 0.0 + for i := range n { + d := a.FloatAt(i) - c[i] + total += d*d + 0.01*a.FloatAt(i)*a.FloatAt(i) + } + return total, nil + } + grad = func(a *core.Array) (*core.Array, error) { + out := core.New(core.Float, n) + vals := out.RawFloats() + for i := range n { + vals[i] = 2*(a.FloatAt(i)-c[i]) + 0.02*a.FloatAt(i) + } + return out, nil + } + x0 = make([]float64, n) + for i := range x0 { + x0[i] = 1 + } + return f, grad, x0 +} + +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 BenchmarkMinimiseLBFGS(b *testing.B) { + f, grad, x0 := benchQuadratic(64) + start := benchVector(b, x0) + opts := LBFGSOptions{MaxIterations: 200} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseLBFGS(f, grad, start, opts); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMinimiseLBFGSFiniteDiff(b *testing.B) { + f, _, x0 := benchQuadratic(64) + start := benchVector(b, x0) + opts := LBFGSOptions{MaxIterations: 200} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseLBFGS(f, nil, start, opts); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMinimiseLBFGSBounded(b *testing.B) { + f, grad, x0 := benchQuadratic(64) + lower := make([]float64, 64) + upper := make([]float64, 64) + for i := range upper { + lower[i] = -2 + upper[i] = 2 + } + start := benchVector(b, x0) + opts := LBFGSOptions{MaxIterations: 200, Lower: lower, Upper: upper} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseLBFGS(f, grad, start, opts); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMinimiseSimplex(b *testing.B) { + f, _, x0 := benchQuadratic(8) + start := benchVector(b, x0) + opts := MinimiseOptions{MaxIterations: 500} + b.ReportAllocs() + for b.Loop() { + if _, _, err := Minimise(f, start, opts); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLevenbergMarquardt(b *testing.B) { + // Fit y = p0·exp(−p1·t) on 40 noisy-free samples. + const nObs = 40 + t := make([]float64, nObs) + y := make([]float64, nObs) + for i := range nObs { + t[i] = float64(i) / 4 + y[i] = 2.5 * math.Exp(-0.7*t[i]) + } + residual := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, nObs) + vals := out.RawFloats() + for i := range nObs { + vals[i] = p.FloatAt(0)*math.Exp(-p.FloatAt(1)*t[i]) - y[i] + } + return out, nil + } + p0 := benchVector(b, []float64{1, 0.2}) + opts := LMOptions{MaxIterations: 30} + b.ReportAllocs() + for b.Loop() { + if _, _, err := LevenbergMarquardt(residual, p0, opts); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkFindRootSystem(b *testing.B) { + n := 6 + r := func(x *core.Array) (*core.Array, error) { + out := core.New(core.Float, n) + vals := out.RawFloats() + for i := range n { + vals[i] = x.FloatAt(i)*x.FloatAt(i) - float64(i+1) + } + return out, nil + } + start := make([]float64, n) + for i := range start { + start[i] = float64(i) + 1.5 + } + x0 := benchVector(b, start) + opts := RootSystemOptions{MaxIterations: 40} + b.ReportAllocs() + for b.Loop() { + if _, _, err := FindRootSystem(r, x0, opts); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMinimiseDifferentialEvolution(b *testing.B) { + f, _, _ := benchQuadratic(4) + lower := benchVector(b, []float64{-5, -5, -5, -5}) + upper := benchVector(b, []float64{5, 5, 5, 5}) + opts := DifferentialEvolutionOptions{Generations: 20} + b.ReportAllocs() + for b.Loop() { + if _, _, err := MinimiseDifferentialEvolution(f, lower, upper, opts); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkFindRoot guards the scalar Brent iteration the root +// benchmarks above surround. +func BenchmarkFindRoot(b *testing.B) { + f := math.Cos + b.ReportAllocs() + for b.Loop() { + if _, err := FindRoot(f, 0.5, 2, 0); err != nil { + b.Fatal(err) + } + } +} diff --git a/optim/bounded_test.go b/optim/bounded_test.go new file mode 100644 index 0000000..909915b --- /dev/null +++ b/optim/bounded_test.go @@ -0,0 +1,165 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestLBFGSBoundedQuadratic pins a coordinate onto each wall kind: the +// optimum of the separable bowl sits at (1, 1), the box drags the +// first coordinate to its lower wall and leaves the second free, which +// is the shape every constrained fit with physical parameter ranges +// takes. +func TestLBFGSBoundedQuadratic(t *testing.T) { + f := func(p *core.Array) (float64, error) { + total := 0.0 + for i := range p.Len() { + d := p.FloatAt(i) - 1 + total += d * d + } + return total, nil + } + start, _ := core.FromFloats([]float64{0, 0}, 2) + point, value, err := MinimiseLBFGS(f, nil, start, LBFGSOptions{ + Lower: []float64{2, math.Inf(-1)}, + Upper: []float64{math.Inf(1), 9}, + }) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + if math.Abs(point.FloatAt(0)-2) > 1e-6 { + t.Fatalf("first coordinate = %.10g, want 2 on the lower wall", point.FloatAt(0)) + } + if math.Abs(point.FloatAt(1)-1) > 1e-6 { + t.Fatalf("second coordinate = %.10g, want 1 free", point.FloatAt(1)) + } + if value > 1+1e-6 { + t.Fatalf("minimum value = %.10g, want 1", value) + } +} + +// TestLBFGSBoundedRosenbrock is the analytic case: with x forced past +// 1.5, the valley's unconstrained neck at (1, 1) is infeasible and the +// constrained minimum sits exactly on the wall at (1.5, 2.25) with +// value 0.25. The wall coordinate's gradient pushes outward, which is +// the KKT signature the optimiser must respect rather than project it +// away. +func TestLBFGSBoundedRosenbrock(t *testing.T) { + rosenbrock := func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + } + gradFn := func(p *core.Array) (*core.Array, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + out := core.New(core.Float, 2) + out.RawFloats()[0] = -2*(1-x) - 400*x*(y-x*x) + out.RawFloats()[1] = 200 * (y - x*x) + return out, nil + } + start, _ := core.FromFloats([]float64{-1.2, 1}, 2) + point, value, err := MinimiseLBFGS(rosenbrock, gradFn, start, LBFGSOptions{ + Lower: []float64{1.5, math.Inf(-1)}, + }) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + if math.Abs(point.FloatAt(0)-1.5) > 1e-6 { + t.Fatalf("first coordinate = %.10g, want 1.5 on the wall", point.FloatAt(0)) + } + if math.Abs(point.FloatAt(1)-2.25) > 1e-4 { + t.Fatalf("second coordinate = %.10g, want 2.25", point.FloatAt(1)) + } + if math.Abs(value-0.25) > 1e-6 { + t.Fatalf("value = %.10g, want 0.25", value) + } +} + +// TestLBFGSBoundedDomain proves the finite-difference gradient turns +// one-sided at a wall: log is undefined below the wall, so a central +// stencil would make the objective return an error and the run would +// fail outright. +func TestLBFGSBoundedDomain(t *testing.T) { + f := func(p *core.Array) (float64, error) { + x := p.FloatAt(0) + if x < 0.5 { + return 0, base.Errf("the objective is undefined below 0.5") + } + return math.Log(x), nil + } + start, _ := core.FromFloats([]float64{2}, 1) + point, _, err := MinimiseLBFGS(f, nil, start, LBFGSOptions{ + Lower: []float64{0.5}, + }) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + if math.Abs(point.FloatAt(0)-0.5) > 1e-6 { + t.Fatalf("coordinate = %.10g, want 0.5 on the wall", point.FloatAt(0)) + } +} + +// TestLBFGSBoundProjection checks that an infeasible start is +// projected onto the box and still converges, and that hostile bounds +// are refused rather than silently swapped or clamped. +func TestLBFGSBoundProjection(t *testing.T) { + f := func(p *core.Array) (float64, error) { + d := p.FloatAt(0) - 3 + return d * d, nil + } + start, _ := core.FromFloats([]float64{-10}, 1) + point, _, err := MinimiseLBFGS(f, nil, start, LBFGSOptions{Lower: []float64{1}}) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + if math.Abs(point.FloatAt(0)-3) > 1e-6 { + t.Fatalf("coordinate = %.10g, want 3", point.FloatAt(0)) + } + bad, _ := core.FromFloats([]float64{0}, 1) + if _, _, err := MinimiseLBFGS(f, nil, bad, LBFGSOptions{Lower: []float64{2}, Upper: []float64{1}}); err == nil { + t.Fatal("crossed walls accepted") + } + if _, _, err := MinimiseLBFGS(f, nil, bad, LBFGSOptions{Lower: []float64{1, 2}}); err == nil { + t.Fatal("bounds of the wrong length accepted") + } + if _, _, err := MinimiseLBFGS(f, nil, bad, LBFGSOptions{Lower: []float64{math.NaN()}}); err == nil { + t.Fatal("NaN wall accepted") + } +} + +// TestLBFGSBoundedAgainstEvolution cross-checks the local method +// against the global one on the same box: differential evolution +// clamps its population into the bounds too, and on a convex problem +// both must land on the same value. +func TestLBFGSBoundedAgainstEvolution(t *testing.T) { + f := func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return (x-2)*(x-2) + 10*(y+1)*(y+1), nil + } + start, _ := core.FromFloats([]float64{0, 0}, 2) + point, value, err := MinimiseLBFGS(f, nil, start, LBFGSOptions{ + Lower: []float64{3, math.Inf(-1)}, + Upper: []float64{math.Inf(1), 0.5}, + }) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + if math.Abs(point.FloatAt(0)-3) > 1e-6 || math.Abs(point.FloatAt(1)+1) > 1e-6 { + t.Fatalf("minimiser = (%.10g, %.10g), want (3, -1)", point.FloatAt(0), point.FloatAt(1)) + } + lower, _ := core.FromFloats([]float64{3, -1}, 2) + upper, _ := core.FromFloats([]float64{10, 5}, 2) + _, evoValue, err := MinimiseDifferentialEvolution(f, lower, upper, + DifferentialEvolutionOptions{Seed: 7, Generations: 200}) + if err != nil { + t.Fatalf("MinimiseDifferentialEvolution: %v", err) + } + if value > evoValue+1e-6 || evoValue > value+1e-4 { + t.Fatalf("L-BFGS value %.10g and evolution value %.10g disagree", value, evoValue) + } +} diff --git a/optim/brent.go b/optim/brent.go new file mode 100644 index 0000000..973e382 --- /dev/null +++ b/optim/brent.go @@ -0,0 +1,144 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// Scalar root bracketing by Brent's method with the iteration's +// budget and tolerance under the caller's control. FindRoot answers +// the common case with fixed settings; FindRootBrent exists for the +// caller who must pin the residual it accepts and the evaluations it +// pays, and who wants a bracket float64 can no longer refine reported +// as a failure instead of receiving a point that misses the +// tolerance. + +// BrentOptions tunes FindRootBrent, in the vocabulary of +// RootSystemOptions: Tolerance ≤ 0 means 1e-10, MaxIterations ≤ 0 +// means 100. The tolerance is a threshold on |f| at the returned +// point, the scalar counterpart of the residual infinity norm, and a +// run costs at most MaxIterations + 2 evaluations of f. +type BrentOptions struct { + Tolerance float64 + MaxIterations int +} + +// FindRootBrent returns a root of f inside the bracket [a, b] by +// Brent's method: inverse quadratic interpolation, with the secant +// step as its fallback and bisection as the guarantee. f(a) and f(b) +// must be finite with opposite signs, so a root is guaranteed inside; +// a bracket that does not change sign is an error. The method never +// evaluates f outside the bracket it was given, the returned point +// always lies inside the initial bracket with |f| at or below the +// tolerance, and an exhausted budget is an error, never a silent +// guess. +func FindRootBrent(f func(float64) float64, a, b float64, opts BrentOptions) (float64, error) { + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-10 + } + if opts.MaxIterations <= 0 { + opts.MaxIterations = 100 + } + fa, fb := f(a), f(b) + if math.IsNaN(fa) || math.IsNaN(fb) || math.IsInf(fa, 0) || math.IsInf(fb, 0) { + return 0, base.Errf("FindRootBrent: the bracket must evaluate to finite values, got f(%g)=%g, f(%g)=%g", a, fa, b, fb) + } + if fa == 0 { + return a, nil + } + if fb == 0 { + return b, nil + } + if sameSign(fa, fb) { + return 0, base.Errf("FindRootBrent: the bracket [%g, %g] does not change sign (f(a)=%g, f(b)=%g)", a, b, fa, fb) + } + // Brent's iteration: b is the best estimate, c the opposite-sign + // end of the bracket, e the width of the step before last. An + // interpolation step is accepted only when it lands safely inside + // the bracket; otherwise the midpoint is taken, so every + // evaluation falls in [b, c] and the sign change survives every + // round. + c, fc := a, fa + d, e := b-a, b-a + for range opts.MaxIterations { + if sameSign(fb, fc) { + // The last step swallowed the opposite sign: drop back to + // the previous iterate, which still brackets the root. + c, fc = a, fa + d, e = b-a, b-a + } + if math.Abs(fc) < math.Abs(fb) { + // Keep b the better end: the smaller residual. + a, b, c = b, c, b + fa, fb, fc = fb, fc, fb + } + // The width below which float64 can no longer separate two + // points of the bracket: the tolerance itself stays out of the + // floor, so the tolerance means |f| and nothing else. + tol1 := 2 * base.EpsF * math.Abs(b) + xm := 0.5 * (c - b) + if fb == 0 || math.Abs(fb) <= opts.Tolerance { + return b, nil + } + if math.Abs(xm) <= tol1 { + return 0, base.Errf("FindRootBrent: the bracket collapsed to machine width at x=%g with |f|=%g, above the tolerance %g", + b, math.Abs(fb), opts.Tolerance) + } + if math.Abs(e) >= tol1 && math.Abs(fa) > math.Abs(fb) { + s := fb / fa + var p, q float64 + if a == c { + // Secant. + p = 2 * xm * s + q = 1 - s + } else { + // Inverse quadratic interpolation through (a, fa), + // (b, fb) and (c, fc). + q = fa / fc + r := fb / fc + p = s * (2*xm*q*(q-r) - (b-a)*(r-1)) + q = (q - 1) * (r - 1) * (s - 1) + } + if p > 0 { + q = -q + } + p = math.Abs(p) + // Accept the interpolated step only when it stays well + // short of the bracket's far end and at least half the + // previous step; bisection otherwise. + if 2*p < min(3*xm*q-math.Abs(tol1*q), math.Abs(e*q)) { + e, d = d, p/q + } else { + d, e = xm, xm + } + } else { + d, e = xm, xm + } + a, fa = b, fb + if math.Abs(d) > tol1 { + b += d + } else { + b += tol1 * signOf(xm) + } + fb = f(b) + if math.IsNaN(fb) || math.IsInf(fb, 0) { + return 0, base.Errf("FindRootBrent: the objective left the real numbers at x=%g", b) + } + if fb == 0 { + return b, nil + } + } + return 0, base.Errf("FindRootBrent: no convergence in %d iterations", opts.MaxIterations) +} + +// sameSign reports whether two finite values carry the same sign. The +// comparison is by sign and not by product, so a pair of tiny values +// whose product underflows to zero is still recognised as +// same-signed, which the product test FindRoot uses would miss. +func sameSign(x, y float64) bool { + return (x > 0 && y > 0) || (x < 0 && y < 0) +} diff --git a/optim/brent_test.go b/optim/brent_test.go new file mode 100644 index 0000000..b0419e9 --- /dev/null +++ b/optim/brent_test.go @@ -0,0 +1,203 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// TestFindRootBrentTranscendentalRoots pins the documented contract on +// known transcendental roots: the answer sits inside the bracket and +// carries |f| at or below the default tolerance of 1e-10. +func TestFindRootBrentTranscendentalRoots(t *testing.T) { + cases := []struct { + name string + f func(float64) float64 + a, b float64 + root float64 + }{ + {"sin", math.Sin, 1, 4, math.Pi}, + {"cos x = x", func(x float64) float64 { return math.Cos(x) - x }, 0, 1, 0.7390851332151607}, + {"exp(x) = 10", func(x float64) float64 { return math.Exp(x) - 10 }, 0, 3, math.Log(10)}, + {"x³ = 2x + 5", func(x float64) float64 { return x*x*x - 2*x - 5 }, 1, 3, 2.0945514815423265}, + } + for _, tc := range cases { + root, err := FindRootBrent(tc.f, tc.a, tc.b, BrentOptions{}) + if err != nil { + t.Fatalf("%s: FindRootBrent: %v", tc.name, err) + } + if root <= tc.a || root >= tc.b { + t.Fatalf("%s: root %g outside the open bracket (%g, %g)", tc.name, root, tc.a, tc.b) + } + if math.Abs(root-tc.root) > 1e-9 { + t.Fatalf("%s: root = %.16g, want %.16g", tc.name, root, tc.root) + } + if fx := math.Abs(tc.f(root)); fx > 1e-10 { + t.Fatalf("%s: |f(root)| = %g, want ≤ 1e-10", tc.name, fx) + } + } +} + +// TestFindRootBrentRefusesSameSignBracket pins the bracket refusal: a +// pair of endpoint values of one sign carries no guarantee, and the +// solver answers with the package's error style instead of a guess. +func TestFindRootBrentRefusesSameSignBracket(t *testing.T) { + f := func(x float64) float64 { return x*x - 1 } + _, err := FindRootBrent(f, 2, 3, BrentOptions{}) + if err == nil || !strings.Contains(err.Error(), "does not change sign") { + t.Fatalf("err = %v, want the sign-change refusal", err) + } + if strings.Count(err.Error(), "FindRootBrent") != 1 { + t.Fatalf("err = %v, want exactly one entry-point prefix", err) + } + // The reversed orientation refuses for the same reason. + if _, err := FindRootBrent(f, 3, 2, BrentOptions{}); err == nil { + t.Fatal("reversed same-sign bracket: want an error") + } + // A non-finite endpoint value is refused before anything moves. + nan := func(float64) float64 { return math.NaN() } + if _, err := FindRootBrent(nan, 1, 2, BrentOptions{}); err == nil || !strings.Contains(err.Error(), "finite values") { + t.Fatalf("err = %v, want the finite-value refusal", err) + } +} + +// TestFindRootBrentAnswersEndpointRoot pins that a root sitting +// exactly on a bracket end is answered as it stands, in either +// orientation, and that an exact zero hit mid-run comes straight back. +func TestFindRootBrentAnswersEndpointRoot(t *testing.T) { + f := func(x float64) float64 { return x - 2 } + for _, ends := range [][2]float64{{2, 5}, {5, 2}} { + root, err := FindRootBrent(f, ends[0], ends[1], BrentOptions{}) + if err != nil { + t.Fatalf("FindRootBrent([%g, %g]): %v", ends[0], ends[1], err) + } + if root != 2 { + t.Fatalf("FindRootBrent([%g, %g]) = %g, want 2", ends[0], ends[1], root) + } + } + // A bisection whose midpoint is the root exactly. + root, err := FindRootBrent(f, 0, 4, BrentOptions{}) + if err != nil || root != 2 { + t.Fatalf("FindRootBrent([0, 4]) = (%g, %v), want (2, nil)", root, err) + } +} + +// TestFindRootBrentSteepSlopeRefuses pins the honest failure at the +// resolution limit: a sign change carried by a jump has no point where +// |f| can sit below the tolerance, and once the bracket closes around +// the jump the solver reports the collapse instead of a point that +// misses the documented contract. +func TestFindRootBrentSteepSlopeRefuses(t *testing.T) { + f := func(x float64) float64 { + if x < 0.5 { + return -1 + } + return 1e6 + } + _, err := FindRootBrent(f, 0, 1, BrentOptions{}) + if err == nil || !strings.Contains(err.Error(), "collapsed to machine width") { + t.Fatalf("err = %v, want the collapsed-bracket refusal", err) + } +} + +// TestFindRootBrentStaysInsideAndBeatsBisection pins two properties on +// the classic cubic x³ = 2x + 5: every evaluation lands inside the +// initial bracket, and the superlinear convergence shows in the +// evaluation count against pure bisection under the same residual +// stop. +func TestFindRootBrentStaysInsideAndBeatsBisection(t *testing.T) { + const a, b, tol = 1.0, 3.0, 1e-10 + fc := func(x float64) float64 { return x*x*x - 2*x - 5 } + var evals []float64 + root, err := FindRootBrent(func(x float64) float64 { + evals = append(evals, x) + return fc(x) + }, a, b, BrentOptions{Tolerance: tol}) + if err != nil { + t.Fatalf("FindRootBrent: %v", err) + } + if math.Abs(fc(root)) > tol { + t.Fatalf("|f(root)| = %g, want ≤ %g", math.Abs(fc(root)), tol) + } + for _, x := range evals { + if x < a || x > b { + t.Fatalf("evaluated f at %g, outside the bracket [%g, %g]", x, a, b) + } + } + // The documented bound: two endpoint evaluations plus one per + // iteration of the default budget. + if len(evals) > 100+2 { + t.Fatalf("%d evaluations, want at most MaxIterations + 2", len(evals)) + } + // Pure bisection under the same residual stop, counted the same + // way. The cubic's slope near the root makes |f| ≤ 1e-10 need a + // bracket about 1e-11 wide, which bisection reaches only after a + // long halving chain. + bisectionEvals := 0 + bisect := func(x float64) float64 { + bisectionEvals++ + return fc(x) + } + xl, xr := a, b + fxl := bisect(xl) + for { + mid := 0.5 * (xl + xr) + fm := bisect(mid) + if math.Abs(fm) <= tol || math.Abs(xr-xl) <= 2*base.EpsF*math.Abs(mid) { + break + } + if (fxl < 0) == (fm < 0) { + xl, fxl = mid, fm + } else { + xr = mid + } + } + if len(evals)*2 >= bisectionEvals { + t.Fatalf("Brent used %d evaluations against bisection's %d, want less than half", len(evals), bisectionEvals) + } + t.Logf("Brent %d evaluations, bisection %d for the same tolerance", len(evals), bisectionEvals) +} + +// TestFindRootBrentBudgetError pins the bounded iteration count: a run +// that exhausts MaxIterations reports the budget, never a guess. +func TestFindRootBrentBudgetError(t *testing.T) { + f := func(x float64) float64 { return x*x*x - 2*x - 5 } + _, err := FindRootBrent(f, 1, 3, BrentOptions{MaxIterations: 2}) + if err == nil || !strings.Contains(err.Error(), "no convergence in 2 iterations") { + t.Fatalf("err = %v, want the budget refusal", err) + } +} + +// TestFindRootBrentPoleInsideBracket pins the mid-run guard: a pole +// between the endpoints mimics a sign change, and the first bisection +// of [1, 3] for 1/(x − 2) lands exactly on the pole, whose infinite +// value is an error naming the point. +func TestFindRootBrentPoleInsideBracket(t *testing.T) { + f := func(x float64) float64 { return 1 / (x - 2) } + _, err := FindRootBrent(f, 1, 3, BrentOptions{}) + if err == nil || !strings.Contains(err.Error(), "left the real numbers") { + t.Fatalf("err = %v, want the non-finite refusal", err) + } +} + +// TestFindRootBrentNoAllocationsInLoop pins the allocation discipline: +// the iteration is scalar bookkeeping over the callback, and a +// converged run allocates nothing. +func TestFindRootBrentNoAllocationsInLoop(t *testing.T) { + f := func(x float64) float64 { return x*x*x - 2*x - 5 } + var err error + allocs := testing.AllocsPerRun(20, func() { + _, err = FindRootBrent(f, 1, 3, BrentOptions{}) + }) + if err != nil { + t.Fatalf("FindRootBrent: %v", err) + } + if allocs != 0 { + t.Fatalf("%g allocations per run, want 0", allocs) + } +} diff --git a/optim/broyden.go b/optim/broyden.go new file mode 100644 index 0000000..01c1595 --- /dev/null +++ b/optim/broyden.go @@ -0,0 +1,126 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// Broyden's quasi-Newton maintenance for FindRootSystem, enabled by +// RootSystemOptions.UseBroyden. The central-difference Jacobian is +// built once at the start and the iteration walks on its maintained +// inverse: after every accepted step the rank-one update corrects the +// inverse so that it satisfies the secant equation H·y = s for the +// step just taken. +// +// The update maintained here is the bad Broyden form in its inverse +// shape, H + (s − H·y)·yᵀ/(yᵀy) with H the maintained inverse: a +// rank-one correction along the residual-difference direction y that +// satisfies the secant equation exactly for the step just taken, and +// turns every later iteration into a single matrix-vector product, +// δ = −H·r. The good form, whose correction runs along sᵀH instead, +// satisfies the same equation with a different matrix; maintaining it +// needs the product sᵀH·y on top of the step, which spends back the +// per-iteration saving the option exists for. + +// broydenInvert fills h with the inverse of the central-difference +// Jacobian jac by solving jac·h = I column-wise through the library's +// LU solver: one factorisation, n right-hand columns, the same cost +// class as the Newton path's single solve. jacWork is consumed in +// place, exactly as the Newton path consumes it, so jac stays +// pristine for the steepest-descent fallback. A singular Jacobian is +// returned as the solver's own error for the caller to fall back on. +func broydenInvert(jac, jacWork, h [][]float64) error { + n := len(jac) + rhs := make([][]float64, n) + for i := range n { + rhs[i] = make([]float64, n) + rhs[i][i] = 1 + } + for i := range n { + copy(jacWork[i], jac[i]) + } + sol, err := base.SolveSystem("FindRootSystem", jacWork, rhs) + if err != nil { + return err + } + // sol[k] is jac⁻¹·eₖ, the kth column of the inverse. + for i := range n { + for k := range n { + h[i][k] = sol[k][i] + } + } + return nil +} + +// broydenMaintain applies the rank-one update +// +// H ← H + (s − H·y)·yᵀ/(yᵀy) +// +// to the maintained inverse, where s is the accepted step and y the +// residual change it produced, and reports whether the inverse is fit +// to carry forward, together with the running count of consecutive +// steps that failed to lower the residual infinity norm. A false +// report leaves h untouched and orders FindRootSystem to rebuild the +// Jacobian numerically before it steps again, the restart the option +// documents. Two observations order the restart, both a degradation of +// the rank-one model: +// +// - the update denominator yᵀy is zero, non-finite, or at rounding +// level against the residual's own scale (‖y‖∞ ≤ ε·max(1, ‖r‖∞)): +// the division would amplify cancellation noise into H, and a y +// that small carries no curvature information at all; +// - two consecutive accepted steps each failed to lower the +// residual infinity norm. The damping guarantees the residual sum +// of squares falls on every accepted step, so a flat infinity +// norm twice in a row means the maintained inverse has stopped +// predicting the landscape and a fresh Jacobian is cheaper than +// more crawling. +func broydenMaintain(h [][]float64, s, y, r []float64, res, resPrev float64, stalled int) (bool, int) { + den := 0.0 + ynorm := 0.0 + for i := range y { + den += y[i] * y[i] + if v := math.Abs(y[i]); v > ynorm { + ynorm = v + } + } + if den == 0 || math.IsNaN(den) || math.IsInf(den, 0) || ynorm <= base.EpsF*math.Max(1, normInfOfStep(r)) { + return false, 0 + } + if res < resPrev { + stalled = 0 + } else { + stalled++ + if stalled >= 2 { + return false, 0 + } + } + for i := range h { + hy := 0.0 + for k := range y { + hy += h[i][k] * y[k] + } + w := (s[i] - hy) / den + for k := range y { + h[i][k] += w * y[k] + } + } + return true, stalled +} + +// broydenStep writes the quasi-Newton step δ = −H·r into step: the +// whole per-iteration linear algebra the maintained inverse leaves, +// a single matrix-vector product. +func broydenStep(h [][]float64, r, step []float64) { + for i := range step { + s := 0.0 + for k := range r { + s += h[i][k] * r[k] + } + step[i] = -s + } +} diff --git a/optim/broyden_test.go b/optim/broyden_test.go new file mode 100644 index 0000000..5419654 --- /dev/null +++ b/optim/broyden_test.go @@ -0,0 +1,282 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// jacobianBuilds counts the central-difference Jacobian builds in a +// recorded residual trace. A build is n consecutive ± pairs, column j +// perturbed first, each pair one stencil width √ε·max(1, |xⱼ|) about +// its base point: the exact pattern FindRootSystem's sweep produces, +// which backtracking trials (single points, several moving +// coordinates at once) never match. +func jacobianBuilds(points [][]float64) int { + n := len(points[0]) + stencil := math.Sqrt(base.EpsF) + matchPair := func(p, q []float64, col int) bool { + diff := -1 + for k := range n { + if p[k] != q[k] { + if diff != -1 { + return false + } + diff = k + } + } + if diff != col { + return false + } + mid := (p[col] + q[col]) / 2 + eps := math.Abs(p[col]-q[col]) / 2 + want := stencil * math.Max(1, math.Abs(mid)) + return math.Abs(eps-want) <= 1e-6*want + } + count := 0 + i := 0 + for i+2*n <= len(points) { + built := true + for col := range n { + if !matchPair(points[i+2*col], points[i+2*col+1], col) { + built = false + break + } + } + if built { + count++ + i += 2 * n + continue + } + i++ + } + return count +} + +// traceResidual wraps a residual so every evaluation's point is +// recorded, for the stencil counter to walk. +func traceResidual(t *testing.T, trace *[][]float64, n int, r func(x []float64) []float64) func(*core.Array) (*core.Array, error) { + return func(x *core.Array) (*core.Array, error) { + v := make([]float64, n) + for i := range n { + v[i] = x.FloatAt(i) + } + *trace = append(*trace, v) + return mustFloats(t, r(v)), nil + } +} + +// TestFindRootSystemBroydenOneJacobian pins the option's promise on +// the analytic systems: with UseBroyden the same roots are reached +// within tolerance and exactly one numerical Jacobian is built, the +// one at the start. +func TestFindRootSystemBroydenOneJacobian(t *testing.T) { + cases := []struct { + name string + n int + start []float64 + r func(x []float64) []float64 + want []float64 + }{ + {"circle-line", 2, []float64{0.5, 0.5}, + func(x []float64) []float64 { return []float64{x[0]*x[0] + x[1]*x[1] - 4, x[0] - x[1]} }, + []float64{math.Sqrt2, math.Sqrt2}}, + {"circle-hyperbola", 2, []float64{0.4, 2.2}, + func(x []float64) []float64 { return []float64{x[0]*x[0] + x[1]*x[1] - 5, x[0]*x[1] - 2} }, + nil}, + {"trig", 2, []float64{0.3, 0.1}, + func(x []float64) []float64 { return []float64{math.Cos(x[0]) - x[1], math.Sin(x[0]) - x[1]} }, + []float64{math.Pi / 4, math.Sqrt2 / 2}}, + } + for _, tc := range cases { + var trace [][]float64 + residual := traceResidual(t, &trace, tc.n, tc.r) + x, res, err := FindRootSystem(residual, mustFloats(t, tc.start), RootSystemOptions{UseBroyden: true}) + if err != nil { + t.Fatalf("%s: FindRootSystem(UseBroyden): %v", tc.name, err) + } + if res > 1e-10 { + t.Fatalf("%s: residual %g, want ≤ 1e-10", tc.name, res) + } + if tc.want != nil { + for i := range tc.n { + if math.Abs(x.FloatAt(i)-tc.want[i]) > 1e-9 { + t.Fatalf("%s: x[%d] = %.12g, want %.12g", tc.name, i, x.FloatAt(i), tc.want[i]) + } + } + } else { + // The hyperbola's two roots are (1, 2) and (2, 1). + s1 := math.Abs(x.FloatAt(0)-1) < 1e-9 && math.Abs(x.FloatAt(1)-2) < 1e-9 + s2 := math.Abs(x.FloatAt(0)-2) < 1e-9 && math.Abs(x.FloatAt(1)-1) < 1e-9 + if !s1 && !s2 { + t.Fatalf("%s: solution = (%.12g, %.12g), want (1, 2) or (2, 1)", + tc.name, x.FloatAt(0), x.FloatAt(1)) + } + } + if got := jacobianBuilds(trace); got != 1 { + t.Fatalf("%s: %d numerical Jacobian builds, want 1", tc.name, got) + } + } +} + +// TestFindRootSystemBroydenEightUnknowns pins the option on a harder +// system: eight coupled nonlinear equations with the known root +// xᵢ = i+1, converged under the default budget with the single +// starting Jacobian. +func TestFindRootSystemBroydenEightUnknowns(t *testing.T) { + const n = 8 + var trace [][]float64 + residual := traceResidual(t, &trace, n, func(x []float64) []float64 { + r := make([]float64, n) + for i := range n { + r[i] = x[i]*x[i] - float64(i+1)*float64(i+1) + for j := range n { + if j != i { + r[i] += 0.05 * (x[j] - float64(j+1)) + } + } + } + return r + }) + start := make([]float64, n) + for i := range n { + start[i] = 0.5 * float64(i+1) + } + x, res, err := FindRootSystem(residual, mustFloats(t, start), RootSystemOptions{UseBroyden: true}) + if err != nil { + t.Fatalf("FindRootSystem(UseBroyden, 8 unknowns): %v", err) + } + if res > 1e-10 { + t.Fatalf("residual %g, want ≤ 1e-10", res) + } + for i := range n { + if math.Abs(x.FloatAt(i)-float64(i+1)) > 1e-9 { + t.Fatalf("x[%d] = %.12g, want %d", i, x.FloatAt(i), i+1) + } + } + if got := jacobianBuilds(trace); got != 1 { + t.Fatalf("%d numerical Jacobian builds, want 1", got) + } +} + +// TestFindRootSystemBroydenSingularJacobianDescends pins the singular +// escape under the option. The duplicated equation x² = 1 has a rank- +// one Jacobian at every point, so no Newton solve ever succeeds and +// the steepest-descent fallback carries the iteration. From (2, 2) the +// damped descent lands on a root at once; from (1.5, 1.5) the descent +// cannot reach one within the budget, and the run refuses with the +// budget error while the trace shows the Jacobian was rebuilt every +// single round, the same restart path a degraded update takes. +func TestFindRootSystemBroydenSingularJacobianDescends(t *testing.T) { + rankOne := func(x []float64) []float64 { + return []float64{x[0]*x[0] - 1, x[0]*x[0] - 1} + } + var trace [][]float64 + x, res, err := FindRootSystem(traceResidual(t, &trace, 2, rankOne), + mustFloats(t, []float64{2, 2}), RootSystemOptions{UseBroyden: true}) + if err != nil { + t.Fatalf("FindRootSystem(UseBroyden, singular Jacobian): %v", err) + } + if res > 1e-10 { + t.Fatalf("residual %g, want ≤ 1e-10", res) + } + if math.Abs(math.Abs(x.FloatAt(0))-1) > 1e-8 { + t.Fatalf("x[0] = %.12g, want a root of x² = 1", x.FloatAt(0)) + } + // The hopeless start: the refusal is honest and the rebuilds are + // visible in the trace, one per round while the inverse stays + // unfit. + trace = nil + if _, _, err := FindRootSystem(traceResidual(t, &trace, 2, rankOne), + mustFloats(t, []float64{1.5, 1.5}), RootSystemOptions{UseBroyden: true}); err == nil { + t.Fatal("a stalling singular system: want the budget refusal") + } + if got := jacobianBuilds(trace); got < 50 { + t.Fatalf("%d numerical Jacobian builds, want one per round: a singular inverse rebuilds every time", got) + } +} + +// TestFindRootSystemBroydenBudgetStillRefused pins that the shared +// error contract survives the option: an impossible tolerance under +// UseBroyden is a budget refusal, not a silent answer. +func TestFindRootSystemBroydenBudgetStillRefused(t *testing.T) { + residual := func(x *core.Array) (*core.Array, error) { + cx, cy := x.FloatAt(0), x.FloatAt(1) + return mustFloats(t, []float64{cx*cx + cy*cy - 4, cx - cy}), nil + } + _, _, err := FindRootSystem(residual, mustFloats(t, []float64{1, 1}), + RootSystemOptions{UseBroyden: true, MaxIterations: 1, Tolerance: 1e-20}) + if err == nil || !strings.Contains(err.Error(), "MaxIterations=1") { + t.Fatalf("err = %v, want the budget refusal", err) + } +} + +// TestBroydenMaintainSecantCondition pins the update formula itself: +// after the rank-one correction the maintained inverse satisfies the +// secant equation H·y = s exactly to rounding, which is what makes the +// following steps quasi-Newton at all. +func TestBroydenMaintainSecantCondition(t *testing.T) { + h := [][]float64{{2, 0.5}, {-1, 3}} + s := []float64{0.3, -0.7} + y := []float64{1.1, 0.4} + r := []float64{0.9, -0.2} + ok, stalled := broydenMaintain(h, s, y, r, 0.1, 1.0, 0) + if !ok || stalled != 0 { + t.Fatalf("healthy update rejected: ok = %v, stalled = %d", ok, stalled) + } + for i := range 2 { + hy := h[i][0]*y[0] + h[i][1]*y[1] + if math.Abs(hy-s[i]) > 1e-12 { + t.Fatalf("secant equation violated: (H·y)[%d] = %.17g, want %.17g", i, hy, s[i]) + } + } +} + +// TestBroydenMaintainRestartTriggers pins the documented restart +// triggers of the rank-one maintenance: a degenerate denominator +// orders a rebuild at once, a rounding-level residual change likewise, +// and two consecutive steps without a fall of the residual infinity +// norm do what one cannot. A refused update leaves the inverse +// untouched, so the rebuild starts from a Jacobian and not from a +// half-updated one. +func TestBroydenMaintainRestartTriggers(t *testing.T) { + h := [][]float64{{1, 0}, {0, 1}} + // y all zero: the denominator trigger. + if ok, _ := broydenMaintain(h, []float64{1, 1}, []float64{0, 0}, []float64{1, 1}, 1, 2, 0); ok { + t.Fatal("a zero residual change was folded into the inverse") + } + // y at rounding level against the residual's own scale. + if ok, _ := broydenMaintain(h, []float64{1, 1}, []float64{1e-20, 0}, []float64{1, 1}, 1, 2, 0); ok { + t.Fatal("a rounding-level residual change was folded into the inverse") + } + // One stalled step keeps the inverse fit but counts the stall. + ok, stalled := broydenMaintain(h, []float64{0.1, 0}, []float64{0.5, 0.5}, []float64{1, 1}, 2, 2, 0) + if !ok || stalled != 1 { + t.Fatalf("first stall: ok = %v, stalled = %d, want the inverse kept and the stall counted", ok, stalled) + } + // A falling step resets the count. + ok, stalled = broydenMaintain(h, []float64{0.1, 0}, []float64{0.5, 0.5}, []float64{1, 1}, 1, 2, stalled) + if !ok || stalled != 0 { + t.Fatalf("falling step: ok = %v, stalled = %d, want the count reset", ok, stalled) + } + // The second consecutive stall orders a rebuild and leaves the + // inverse untouched. + before := [2][2]float64{{h[0][0], h[0][1]}, {h[1][0], h[1][1]}} + ok, stalled = broydenMaintain(h, []float64{0.1, 0}, []float64{0.5, 0.5}, []float64{1, 1}, 2, 2, 1) + if ok || stalled != 0 { + t.Fatalf("second stall: ok = %v, stalled = %d, want a rebuild ordered", ok, stalled) + } + for i := range 2 { + for k := range 2 { + if h[i][k] != before[i][k] { + t.Fatalf("a refused update moved h[%d][%d] from %g to %g", i, k, before[i][k], h[i][k]) + } + } + } +} diff --git a/optim/budget_exit_pin_test.go b/optim/budget_exit_pin_test.go new file mode 100644 index 0000000..323af66 --- /dev/null +++ b/optim/budget_exit_pin_test.go @@ -0,0 +1,108 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +// TestLBFGSBudgetExitIsNotConvergence pins the last silent exit of the +// optimiser. A run that spends its iteration budget has not converged: +// on a stiff objective it stops with a projected gradient orders of +// magnitude above the tolerance, and reporting that point as the answer +// is the same silent wrongness the stall and direction exits refuse. +// AllowBudgetExit is the documented escape hatch, and it is what the +// augmented Lagrangian's inexact inner solves use. +func TestLBFGSBudgetExitIsNotConvergence(t *testing.T) { + // Mixed units: the y direction is 1e10 times stiffer, so five + // iterations cannot reach the default 1e-8 tolerance. + stiff := func(p *core.Array) (float64, error) { + dx := p.FloatAt(0) - 3 + dy := p.FloatAt(1) - 5 + return dx*dx + 1e10*dy*dy, nil + } + start := mustFloats(t, []float64{0, 0}, 2) + + point, value, err := MinimiseLBFGS(stiff, nil, start, LBFGSOptions{MaxIterations: 5}) + if err == nil { + t.Fatalf("a budget stop was reported as convergence: point = %v, value = %g", floatsOf(point), value) + } + if !strings.Contains(err.Error(), "iteration budget") { + t.Fatalf("error = %v, want the iteration-budget refusal", err) + } + if point != nil { + t.Fatalf("the refused run returned the point %v", floatsOf(point)) + } + + // The escape hatch is what the constrained wrapper relies on: the + // point comes back with no error, and the outer loop's feasibility + // check is what judges it. + point, value, err = MinimiseLBFGS(stiff, nil, start, LBFGSOptions{MaxIterations: 5, AllowBudgetExit: true}) + if err != nil { + t.Fatalf("AllowBudgetExit: %v", err) + } + if point == nil { + t.Fatal("AllowBudgetExit returned no point") + } + t.Logf("best effort after 5 iterations: %v, value %g", floatsOf(point), value) + + // A run that does converge is unaffected by either setting. + for _, allow := range []bool{false, true} { + point, _, err = MinimiseLBFGS(func(p *core.Array) (float64, error) { + dx := p.FloatAt(0) - 3 + return dx * dx, nil + }, nil, mustFloats(t, []float64{0}, 1), LBFGSOptions{AllowBudgetExit: allow}) + if err != nil { + t.Fatalf("AllowBudgetExit=%v on a converging run: %v", allow, err) + } + if math := point.FloatAt(0); math < 2.999999999 || math > 3.000000001 { + t.Fatalf("AllowBudgetExit=%v: point %g, want 3", allow, math) + } + } +} + +// TestMinimiseAndLMBudgetExitIsNotConvergence pins the same refusal on +// Minimise and LevenbergMarquardt: a run that spends its iteration +// budget must not publish its last point as a converged answer, and +// AllowBudgetExit is the documented escape hatch. +func TestMinimiseAndLMBudgetExitIsNotConvergence(t *testing.T) { + stiff := func(p *core.Array) (float64, error) { + dx := p.FloatAt(0) - 3 + dy := p.FloatAt(1) - 5 + return dx*dx + 1e10*dy*dy, nil + } + start := mustFloats(t, []float64{0, 0}, 2) + + point, _, err := Minimise(stiff, start, MinimiseOptions{MaxIterations: 5}) + if err == nil || !strings.Contains(err.Error(), "iteration budget") { + t.Fatalf("Minimise: a budget stop was reported as convergence (err = %v)", err) + } + if point != nil { + t.Fatal("Minimise: the refused run returned a point") + } + point, _, err = Minimise(stiff, start, MinimiseOptions{MaxIterations: 5, AllowBudgetExit: true}) + if err != nil || point == nil { + t.Fatalf("Minimise AllowBudgetExit: err = %v, point = %v", err, point) + } + + residual := func(p *core.Array) (*core.Array, error) { + r := core.New(core.Float, 2) + r.RawFloats()[0] = p.FloatAt(0) - 3 + r.RawFloats()[1] = p.FloatAt(1) - 5 + return r, nil + } + fp, _, err := LevenbergMarquardt(residual, start, LMOptions{MaxIterations: 1}) + if err == nil || !strings.Contains(err.Error(), "iteration budget") { + t.Fatalf("LevenbergMarquardt: a budget stop was reported as convergence (err = %v)", err) + } + if fp != nil { + t.Fatal("LevenbergMarquardt: the refused run returned a point") + } + fp, _, err = LevenbergMarquardt(residual, start, LMOptions{MaxIterations: 1, AllowBudgetExit: true}) + if err != nil || fp == nil { + t.Fatalf("LevenbergMarquardt AllowBudgetExit: err = %v, point = %v", err, fp) + } +} diff --git a/optim/callback_hygiene_pins_test.go b/optim/callback_hygiene_pins_test.go new file mode 100644 index 0000000..09341c3 --- /dev/null +++ b/optim/callback_hygiene_pins_test.go @@ -0,0 +1,142 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Callback-hygiene pins: callbacks that were handed arrays aliasing +// reused scratch, gradients that divided by a stencil clamped to the +// same wall by equal bounds, and NaN objectives that outran their +// diagnosis. + +// TestFindRootSystemCallbackImmutability: the residual callback +// received an aliasing view of the reused iterate, so a callback that +// retained its argument observed the solver's later writes into it. +func TestFindRootSystemCallbackImmutability(t *testing.T) { + var kept *core.Array + f := func(x *core.Array) (*core.Array, error) { + if kept == nil { + kept = x // retain the first argument, as a caching callback would + } + return core.FromFloats([]float64{x.FloatAt(0) - 1, x.FloatAt(1) - 2}, 2) + } + x0, err := core.FromFloats([]float64{0, 0}, 2) + if err != nil { + t.Fatal(err) + } + root, _, err := FindRootSystem(f, x0, RootSystemOptions{}) + if err != nil { + t.Fatal(err) + } + if root.FloatAt(0) != 1 || root.FloatAt(1) != 2 { + t.Fatalf("root = %v, %v", root.FloatAt(0), root.FloatAt(1)) + } + if kept.FloatAt(0) != 0 || kept.FloatAt(1) != 0 { + t.Fatalf("the retained argument was mutated to %v, %v", kept.FloatAt(0), kept.FloatAt(1)) + } +} + +// TestLBFGSEqualBoundsFiniteDifferences: a coordinate pinned by equal +// bounds clamped both stencil points to the same wall and the central +// difference divided 0/0 into a NaN gradient. +func TestLBFGSEqualBoundsFiniteDifferences(t *testing.T) { + f := func(x *core.Array) (float64, error) { + return (x.FloatAt(0)-3)*(x.FloatAt(0)-3) + (x.FloatAt(1)-1)*(x.FloatAt(1)-1), nil + } + x0, err := core.FromFloats([]float64{0, 0}, 2) + if err != nil { + t.Fatal(err) + } + opts := LBFGSOptions{Lower: []float64{1, -10}, Upper: []float64{1, 10}} + pt, val, err := MinimiseLBFGS(f, nil, x0, opts) + if err != nil { + t.Fatalf("MinimiseLBFGS with an equality-pinned coordinate: %v", err) + } + if pt.FloatAt(0) != 1 { + t.Fatalf("pinned coordinate = %g, want 1", pt.FloatAt(0)) + } + if math.Abs(pt.FloatAt(1)-1) > 1e-6 || math.Abs(val-4) > 1e-9 { + t.Fatalf("free coordinate = %g, value = %g (want 1 and 4)", pt.FloatAt(1), val) + } +} + +// TestLBFGSRejectsNaNGradient: a NaN gradient coordinate skipped the +// projected-gradient max and could read as converged. +func TestLBFGSRejectsNaNGradient(t *testing.T) { + f := func(x *core.Array) (float64, error) { return x.FloatAt(0) * x.FloatAt(0), nil } + grad := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{math.NaN()}, 1) + } + x0, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + if _, _, err := MinimiseLBFGS(f, grad, x0, LBFGSOptions{}); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("MinimiseLBFGS with a NaN gradient: err = %v", err) + } +} + +// TestMinimiseRejectsNaNObjective: a NaN vertex burned the budget and +// was reported as a budget problem, or with AllowBudgetExit came back +// as the answer. +func TestMinimiseRejectsNaNObjective(t *testing.T) { + f := func(x *core.Array) (float64, error) { return math.NaN(), nil } + x0, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + if _, _, err := Minimise(f, x0, MinimiseOptions{}); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("Minimise with a NaN objective: err = %v", err) + } +} + +// TestMinimiseConstrainedReturnsObjective: the returned value was the +// augmented Lagrangian the inner solver minimised, not f at the +// answer. +func TestMinimiseConstrainedReturnsObjective(t *testing.T) { + f := func(x *core.Array) (float64, error) { + return (x.FloatAt(0) - 2) * (x.FloatAt(0) - 2), nil + } + grad := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{2 * (x.FloatAt(0) - 2)}, 1) + } + x0, err := core.FromFloats([]float64{0}, 1) + if err != nil { + t.Fatal(err) + } + cons := LinearConstraints{ + A: mustFloats2D(t, []float64{1}, 1, 1), + Lower: []float64{1}, + Upper: []float64{1}, + } + pt, val, err := MinimiseConstrained(f, grad, x0, cons, LBFGSOptions{}) + if err != nil { + t.Fatal(err) + } + if math.Abs(pt.FloatAt(0)-1) > 1e-6 { + t.Fatalf("constrained minimum at %g, want 1", pt.FloatAt(0)) + } + direct, err := f(pt) + if err != nil { + t.Fatal(err) + } + if val != direct { + t.Fatalf("returned value = %.17g, want f at the answer = %.17g exactly", val, direct) + } +} + +func mustFloats2D(t *testing.T, vals []float64, rows, cols int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, rows, cols) + if err != nil { + t.Fatal(err) + } + return a +} diff --git a/optim/callback_pins_test.go b/optim/callback_pins_test.go new file mode 100644 index 0000000..76dc1f1 --- /dev/null +++ b/optim/callback_pins_test.go @@ -0,0 +1,73 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "strings" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Callback-contract pins: a gradient callback that returns a shorter +// array than the problem has dimensions used to index past it, and a +// differential-evolution population below four has no way to draw +// three distinct others. + +// TestLBFGSGradientLength pins the callback contract. +func TestLBFGSGradientLength(t *testing.T) { + f := func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return x*x + y*y, nil + } + short, err := core.FromFloats([]float64{1}, 1) + if err != nil { + t.Fatal(err) + } + grad := func(*core.Array) (*core.Array, error) { return short, nil } + start, err := core.FromFloats([]float64{1, 1}, 2) + if err != nil { + t.Fatal(err) + } + if _, _, err := MinimiseLBFGS(f, grad, start, LBFGSOptions{}); err == nil { + t.Fatal("expected an error for a gradient of the wrong length") + } else if !strings.Contains(err.Error(), "gradient callback") { + t.Fatalf("error = %v, want a callback-length refusal", err) + } +} + +// TestDifferentialEvolutionPopulationFloor pins the option contract: a +// population below four has no way to draw three distinct others, and +// the search used to spin in the picking loop for ever. +func TestDifferentialEvolutionPopulationFloor(t *testing.T) { + lo, err := core.FromFloats([]float64{-5}, 1) + if err != nil { + t.Fatal(err) + } + hi, err := core.FromFloats([]float64{5}, 1) + if err != nil { + t.Fatal(err) + } + f := func(p *core.Array) (float64, error) { v := p.FloatAt(0); return v * v, nil } + done := make(chan error, 1) + go func() { + _, _, err := MinimiseDifferentialEvolution(f, lo, hi, DifferentialEvolutionOptions{Population: 3, Generations: 5}) + done <- err + }() + select { + case err := <-done: + if err == nil { + t.Fatal("expected an error for a population of 3") + } else if !strings.Contains(err.Error(), "population") { + t.Fatalf("error = %v, want a population refusal", err) + } + case <-time.After(10 * time.Second): + t.Fatal("MinimiseDifferentialEvolution did not return: the picking loop spins") + } + // The floor itself still works. + if _, _, err := MinimiseDifferentialEvolution(f, lo, hi, DifferentialEvolutionOptions{Population: 4, Generations: 20}); err != nil { + t.Fatalf("population 4: %v", err) + } +} diff --git a/optim/devolution.go b/optim/devolution.go new file mode 100644 index 0000000..945fdce --- /dev/null +++ b/optim/devolution.go @@ -0,0 +1,186 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Differential evolution: the global optimiser for the +// landscapes the local methods cannot be trusted with: multimodal, +// discontinuous, derivative-free. The classic rand/1/bin scheme: +// every generation, each population member is challenged by a mutant +// built from three distinct others, mixed by binomial crossover, and +// kept only if it beats the incumbent. No gradient, no assumptions +// beyond the bounds. + +// DifferentialEvolutionOptions tunes MinimiseDifferentialEvolution. +// Population defaults to 15·d when unset (at least 4, the scheme's +// minimum); F is the differential weight (default 0.7), CR the +// crossover probability (default 0.9), Generations the budget +// (default 1000), Seed the generator seed (zero is replaced by 42, so +// unset runs are reproducible; every other value, negatives included, +// seeds the xoshiro stream directly). +type DifferentialEvolutionOptions struct { + Population int + F float64 + CR float64 + Generations int + Seed int64 +} + +// MinimiseDifferentialEvolution returns the point and value of the +// global minimum of f over the box [lower, upper] by differential +// evolution (rand/1/bin with reflection-free clamping to the bounds). +// f receives candidate points as rank-1 arrays; a non-finite value is +// an error, mismatched or degenerate bounds are errors, and the result +// is the best point ever evaluated, fresh array the caller owns. +// The generation budget is a tuning parameter, not a convergence +// budget: differential evolution keeps improving the population as +// long as it runs, so the budget running out returns the best point +// found without an error, unlike the local solvers whose AllowBudgetExit +// default refuses a budget stop. +func MinimiseDifferentialEvolution(f func(*core.Array) (float64, error), + lower, upper *core.Array, opts DifferentialEvolutionOptions) (*core.Array, float64, error) { + const name = "MinimiseDifferentialEvolution" + n := lower.Len() + if lower.NDim() != 1 || upper.NDim() != 1 || upper.Len() != n { + return nil, 0, base.Errf("%s: lower and upper must be equal-length rank-1 bounds", name) + } + if n == 0 { + return nil, 0, base.Errf("%s: the bounds must not be empty", name) + } + if lower.Dtype() == core.Complex || upper.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex bounds are not supported", name) + } + for i := range n { + if !(upper.FloatAt(i) > lower.FloatAt(i)) { + return nil, 0, base.Errf("%s: bound %d runs from %g to %g", name, i, lower.FloatAt(i), upper.FloatAt(i)) + } + } + pop := opts.Population + if pop <= 0 { + pop = max(15*n, 4) + } + // The scheme draws three distinct others besides the target, so a + // population below four has no admissible draw: the search would + // never leave the picking loop. + if pop < 4 { + return nil, 0, base.Errf("%s: the population must be at least 4 to draw three distinct others, got %d", + name, pop) + } + fw := opts.F + if fw <= 0 { + fw = 0.7 + } + cr := opts.CR + if cr <= 0 { + cr = 0.9 + } + gens := opts.Generations + if gens <= 0 { + gens = 1000 + } + seed := opts.Seed + if seed == 0 { + seed = 42 + } + g := core.NewGenerator(seed) + + // The bounds hoisted into plain slices: the generation loops read + // them twice per coordinate per trial, and the accessor walk would + // pay the dtype dispatch that many times. The elements are the + // ones FloatAt returned, so the run is bit for bit the same. + lo := make([]float64, n) + hi := make([]float64, n) + for i := range n { + lo[i] = lower.FloatAt(i) + hi[i] = upper.FloatAt(i) + } + + eval := func(x []float64) (float64, error) { + arr, err := core.FromFloats(x, n) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + v, err := f(arr) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, base.Errf("%s: the objective is non-finite (%g)", name, v) + } + return v, nil + } + + // Latin-square-ish start: uniform draws inside the box. + popX := make([][]float64, pop) + popF := make([]float64, pop) + for p := range pop { + x := make([]float64, n) + for i := range n { + x[i] = lo[i] + g.Unit()*(hi[i]-lo[i]) + } + v, err := eval(x) + if err != nil { + return nil, 0, err + } + popX[p], popF[p] = x, v + } + + // pick draws population indices until one avoids the excluded + // candidates; the excludes arrive as plain values, so the hot draw + // loop allocates nothing. A negative exclude never matches, which + // is how the callers drop the slots they do not need. + pick := func(avoid1, avoid2, avoid3 int) int { + for { + c := int(g.Unit() * float64(pop)) + if c >= 0 && c < pop && c != avoid1 && c != avoid2 && c != avoid3 { + return c + } + } + } + + trial := make([]float64, n) + for gen := 0; gen < gens; gen++ { + for target := range pop { + r1 := pick(target, -1, -1) + r2 := pick(target, r1, -1) + r3 := pick(target, r1, r2) + jr := int(g.Unit()*float64(n)) % n + for i := range n { + if g.Unit() < cr || i == jr { + m := popX[r1][i] + fw*(popX[r2][i]-popX[r3][i]) + // Clamp, never reflect: the bounds are the contract. + m = math.Min(math.Max(m, lo[i]), hi[i]) + trial[i] = m + } else { + trial[i] = popX[target][i] + } + } + v, err := eval(trial) + if err != nil { + return nil, 0, err + } + if v <= popF[target] { + copy(popX[target], trial) + popF[target] = v + } + } + } + best := 0 + for p := 1; p < pop; p++ { + if popF[p] < popF[best] { + best = p + } + } + out, err := core.FromFloats(popX[best], n) + if err != nil { + return nil, 0, base.Errf("%s: %w", name, err) + } + return out, popF[best], nil +} diff --git a/optim/devolution_test.go b/optim/devolution_test.go new file mode 100644 index 0000000..a9acb92 --- /dev/null +++ b/optim/devolution_test.go @@ -0,0 +1,114 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestDESphere pins the global minimum of the sphere from a wide box. +func TestDESphere(t *testing.T) { + lower, _ := core.FromFloats([]float64{-10, -10, -10}, 3) + upper, _ := core.FromFloats([]float64{10, 10, 10}, 3) + x, fv, err := MinimiseDifferentialEvolution(func(a *core.Array) (float64, error) { + s := 0.0 + for i := range 3 { + s += a.FloatAt(i) * a.FloatAt(i) + } + return s, nil + }, lower, upper, DifferentialEvolutionOptions{Generations: 300}) + if err != nil { + t.Fatalf("MinimiseDifferentialEvolution: %v", err) + } + if fv > 1e-10 { + t.Fatalf("sphere minimum = %.3e, want 0", fv) + } + for i := range 3 { + if math.Abs(x.FloatAt(i)) > 1e-5 { + t.Fatalf("x[%d] = %g, want 0", i, x.FloatAt(i)) + } + } +} + +// TestDERosenbrockRastrigin pins two multimodal classics: Rosenbrock's +// valley and Rastrigin's minefield of local minima. +func TestDERosenbrockRastrigin(t *testing.T) { + lower, _ := core.FromFloats([]float64{-2.5, -2.5}, 2) + upper, _ := core.FromFloats([]float64{2.5, 2.5}, 2) + _, fvR, err := MinimiseDifferentialEvolution(func(a *core.Array) (float64, error) { + x, y := a.FloatAt(0), a.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + }, lower, upper, DifferentialEvolutionOptions{Generations: 600}) + if err != nil { + t.Fatalf("Rosenbrock: %v", err) + } + if fvR > 1e-6 { + t.Fatalf("Rosenbrock minimum = %.3e, want 0", fvR) + } + + lowerR, _ := core.FromFloats([]float64{-5.12, -5.12, -5.12}, 3) + upperR, _ := core.FromFloats([]float64{5.12, 5.12, 5.12}, 3) + x, fv, err := MinimiseDifferentialEvolution(func(a *core.Array) (float64, error) { + s := 0.0 + for i := range 3 { + z := a.FloatAt(i) + s += z*z - 10*math.Cos(2*math.Pi*z) + 10 + } + return s, nil + }, lowerR, upperR, DifferentialEvolutionOptions{Generations: 800}) + if err != nil { + t.Fatalf("Rastrigin: %v", err) + } + if fv > 1e-6 { + t.Fatalf("Rastrigin minimum = %.3e, want 0 (global basin found at %v)", fv, x) + } +} + +// TestDEBoundClamping pins that the search stays inside the box. +func TestDEBoundClamping(t *testing.T) { + lower, _ := core.FromFloats([]float64{-1, -1}, 2) + upper, _ := core.FromFloats([]float64{1, 1}, 2) + x, _, err := MinimiseDifferentialEvolution(func(a *core.Array) (float64, error) { + // Push toward a corner outside the box. + return -a.FloatAt(0) - 2*a.FloatAt(1), nil + }, lower, upper, DifferentialEvolutionOptions{Generations: 200}) + if err != nil { + t.Fatalf("MinimiseDifferentialEvolution: %v", err) + } + for i := range 2 { + if x.FloatAt(i) < lower.FloatAt(i)-1e-12 || x.FloatAt(i) > upper.FloatAt(i)+1e-12 { + t.Fatalf("x[%d] = %g left the box", i, x.FloatAt(i)) + } + } +} + +// mustBound builds one bound vector. +func mustBound(vs ...float64) *core.Array { + a, _ := core.FromFloats(vs, len(vs)) + return a +} + +// TestDEErrors pins the input gates. +func TestDEErrors(t *testing.T) { + lower, _ := core.FromFloats([]float64{1, 1}, 2) + upper, _ := core.FromFloats([]float64{0, 2}, 2) + if _, _, err := MinimiseDifferentialEvolution(func(a *core.Array) (float64, error) { return 0, nil }, + lower, upper, DifferentialEvolutionOptions{}); err == nil { + t.Error("degenerate bound accepted") + } + lo, _ := core.FromFloats([]float64{0, 0}, 2) + hi, _ := core.FromFloats([]float64{1, 1, 1}, 3) + if _, _, err := MinimiseDifferentialEvolution(func(a *core.Array) (float64, error) { return 0, nil }, + lo, hi, DifferentialEvolutionOptions{}); err == nil { + t.Error("mismatched bounds accepted") + } + if _, _, err := MinimiseDifferentialEvolution(func(a *core.Array) (float64, error) { + return math.NaN(), nil + }, lo, mustBound(1, 1), DifferentialEvolutionOptions{}); err == nil { + t.Error("NaN objective accepted") + } +} diff --git a/optim/doc.go b/optim/doc.go new file mode 100644 index 0000000..eec8bb8 --- /dev/null +++ b/optim/doc.go @@ -0,0 +1,84 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package optim fits, minimises and solves: nonlinear least squares, +// local and global minimisation, constrained minimisation, linear and +// quadratic programming, and root finding for one equation or a system +// of them. +// +// # What it is +// +// Every entry point consumes and returns the library's Array. The +// caller's objective, residual, gradient or constraint function +// receives the candidate point as a rank-1 array of its own and returns +// a value and an error; the answer comes back as a fresh array the +// caller owns, and nothing the caller passed in is modified. +// +// Callbacks run one evaluation at a time on the goroutine that entered +// the solver. The one opt-out is ParallelJacobian on LMOptions and +// RootSystemOptions: setting it consents to the residual callback being +// read from several goroutines at once while a finite-difference +// Jacobian sweeps its columns, and the numbers come out identical +// either way. +// +// The methods, by entry point: +// +// - LevenbergMarquardt fits a model to data by damped Gauss-Newton on +// the residual vector, with a central-difference or an analytic +// Jacobian. +// - Minimise is the derivative-free Nelder-Mead simplex; +// MinimiseLBFGS is limited-memory BFGS with an Armijo backtracking +// line search and optional box walls. +// - MinimiseConstrained adds the linear rows l ≤ A·x ≤ u to an +// objective through an augmented Lagrangian; the same machinery +// carries the functional rows of MinimiseNonlinearConstrained. +// - MinimiseLinear and MinimiseLinearRows solve a linear program by +// the two-phase revised simplex under Bland's rule; MinimiseQP +// solves the strictly convex quadratic program by a primal +// active-set method and returns the row multipliers. +// - MinimiseDifferentialEvolution, MinimiseCMAES and +// MinimiseSimulatedAnnealing search a landscape with several +// basins without derivatives. +// - FindRoot, FindRootBrent and FindRootNewton solve one equation in +// one unknown; FindRootSystem solves as many equations as unknowns. +// +// # Tolerances +// +// The tolerances are absolute in the units of the quantity they +// measure, the caller's own: a gradient coordinate, a row violation, +// an objective spread, a reduced cost. An objective or a constraint +// whose natural scale sits many orders of magnitude away from one +// should be rescaled to O(1) before it is handed to a solver, because +// a solution that is converged in a small unit is reported as +// converged at the point the solver started from. +// +// # Budgets and honesty +// +// No entry point returns a point it did not earn. A local solver that +// spends its iteration budget without meeting its tolerance is refused +// with an error naming the figure it reached and the tolerance it fell +// short of; a line search that stalls, a damping that collapses and a +// search direction that vanishes are refused the same way. Setting +// AllowBudgetExit reports the best point reached instead, which is the +// documented escape hatch and never the default. The constraint +// entries use it for their inner solves on purpose: the outer loop +// judges those points by the rows' feasibility, so an inexact inner +// solve still carries the iteration forward. Differential evolution is +// the exception that needs no flag, because its generation budget +// tunes the search rather than deciding convergence. MinimiseCMAES +// runs once: the restart schemes are the caller's loop, and the +// result of one run is reported as such. +// +// # What it does not do +// +// The package is real-valued: a complex starting point, cost, bound, +// constraint matrix or callback payload is refused with an error. It +// also leaves the modelling to the caller. MinimiseLinear requires the +// standard form A·x = b with x ≥ 0 exactly, and MinimiseLinearRows is +// the wrapper that converts the house two-sided rows into it. +// MinimiseDifferentialEvolution clamps to its box and requires one, +// while MinimiseCMAES has no bounds at all and expects a caller who +// needs them to reparametrise. A local minimum is a local minimum: +// Minimise and MinimiseLBFGS answer for the basin they started in, and +// multistart or a global searcher is what finds the others. +package optim diff --git a/optim/dtypes_census_test.go b/optim/dtypes_census_test.go new file mode 100644 index 0000000..89ab082 --- /dev/null +++ b/optim/dtypes_census_test.go @@ -0,0 +1,348 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The dtype census for optim: the callback machinery's caller-supplied +// arrays (starting points, bounds, constraint tables) and the +// callbacks' own output arrays, probed with Bool, the narrow integers +// and the Int anchor against a float64 baseline carrying exactly the +// widened probe values. The optimisers widen through accessor walks by +// design and refuse complex operands by name; narrow inputs follow Int +// bit for bit. Nothing panics or silently misreads. + +var opDtypes = []core.Dtype{core.Bool, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32, core.Int} + +type opMaker func(vals []float64, shape ...int) *core.Array + +func opCast(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 opMakers(t *testing.T, dt core.Dtype) (probe, base opMaker) { + t.Helper() + castOf := func(vals []float64) []float64 { + out := make([]float64, len(vals)) + for i, v := range vals { + out[i] = opCast(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 opElems(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 opPoint(t *testing.T, label string, dt core.Dtype, pp *core.Array, pv float64, perr error, + bp *core.Array, bv float64, 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 pp == nil || bp == nil { + t.Fatalf("%s(%s): nil point (probe %v, base %v)", label, dt, pp, bp) + } + if pp.Dtype() != bp.Dtype() { + t.Fatalf("%s(%s): point dtype %s, want the baseline dtype %s", label, dt, pp.Dtype(), bp.Dtype()) + } + pe, be := opElems(t, pp), opElems(t, bp) + for i := range pe { + if pe[i] != be[i] { + t.Fatalf("%s(%s): point %d = %v, want %v (all %v vs %v)", label, dt, i, pe[i], be[i], pe, be) + } + } + if pv != bv { + t.Fatalf("%s(%s): value %v, want the baseline %v", label, dt, pv, bv) + } +} + +func opFloats(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 opWantErr(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) + } + } +} + +// quadAround returns the objective Σ(x_i − target_i)², dtype +// independent: the callbacks always see the solver's own float64 +// arrays whatever the caller handed in. +func quadAround(target []float64) func(p *core.Array) (float64, error) { + return func(p *core.Array) (float64, error) { + s := 0.0 + for i := range p.Len() { + d := p.FloatAt(i) - target[i%len(target)] + s += d * d + } + return s, nil + } +} + +// gradAround builds the analytic gradient, answered in the maker's +// dtype: the callback-output surface the census probes. +func gradAround(mk opMaker, target []float64) func(p *core.Array) (*core.Array, error) { + return func(p *core.Array) (*core.Array, error) { + out := make([]float64, p.Len()) + for i := range out { + out[i] = 2 * (p.FloatAt(i) - target[i%len(target)]) + } + return mk(out, len(out)), nil + } +} + +// TestDtypesCensusOptim probes every array-taking public entry. +func TestDtypesCensusOptim(t *testing.T) { + rows := []struct { + name string + run func(t *testing.T, probe, base opMaker, dt core.Dtype) + }{ + {"Minimise", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + f := quadAround([]float64{1, 1}) + pp, pv, perr := Minimise(f, probe([]float64{2, 3}, 2), MinimiseOptions{}) + bp, bv, berr := Minimise(f, base([]float64{2, 3}, 2), MinimiseOptions{}) + opPoint(t, "Minimise", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseLBFGS", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + target := []float64{1, 1} + pp, pv, perr := MinimiseLBFGS(quadAround(target), gradAround(probe, target), probe([]float64{2, 3}, 2), LBFGSOptions{}) + bp, bv, berr := MinimiseLBFGS(quadAround(target), gradAround(base, target), base([]float64{2, 3}, 2), LBFGSOptions{}) + opPoint(t, "MinimiseLBFGS", dt, pp, pv, perr, bp, bv, berr) + }}, + {"LevenbergMarquardt", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + resid := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{p.FloatAt(0) - 1, p.FloatAt(1) - 2}, 2) + } + pp, pv, perr := LevenbergMarquardt(resid, probe([]float64{0, 0}, 2), LMOptions{}) + bp, bv, berr := LevenbergMarquardt(resid, base([]float64{0, 0}, 2), LMOptions{}) + opPoint(t, "LevenbergMarquardt", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseConstrained", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + target := []float64{1, 1} + pcons := LinearConstraints{A: probe([]float64{1, 1}, 1, 2), Lower: []float64{2}, Upper: []float64{2}} + bcons := LinearConstraints{A: base([]float64{1, 1}, 1, 2), Lower: []float64{2}, Upper: []float64{2}} + pp, pv, perr := MinimiseConstrained(quadAround(target), gradAround(probe, target), probe([]float64{2, 3}, 2), pcons, LBFGSOptions{}) + bp, bv, berr := MinimiseConstrained(quadAround(target), gradAround(base, target), base([]float64{2, 3}, 2), bcons, LBFGSOptions{}) + opPoint(t, "MinimiseConstrained", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseNonlinearConstrained", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + target := []float64{1, 1} + cons := NonlinearConstraints{Equalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { return p.FloatAt(0) + p.FloatAt(1) - 2, nil }, + }} + pp, pv, _, perr := MinimiseNonlinearConstrained(quadAround(target), gradAround(probe, target), probe([]float64{2, 3}, 2), cons, LBFGSOptions{}) + bp, bv, _, berr := MinimiseNonlinearConstrained(quadAround(target), gradAround(base, target), base([]float64{2, 3}, 2), cons, LBFGSOptions{}) + opPoint(t, "MinimiseNonlinearConstrained", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseQP", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + pcons := LinearConstraints{A: probe([]float64{1, 1}, 1, 2), Lower: []float64{1}, Upper: []float64{1}} + bcons := LinearConstraints{A: base([]float64{1, 1}, 1, 2), Lower: []float64{1}, Upper: []float64{1}} + pp, pv, _, perr := MinimiseQP(probe([]float64{2, 0, 0, 2}, 2, 2), probe([]float64{2, 2}, 2), pcons, probe([]float64{0, 0}, 2), QPOptions{}) + bp, bv, _, berr := MinimiseQP(base([]float64{2, 0, 0, 2}, 2, 2), base([]float64{2, 2}, 2), bcons, base([]float64{0, 0}, 2), QPOptions{}) + opPoint(t, "MinimiseQP", dt, pp, pv, perr, bp, bv, berr) + }}, + {"FindRootSystem", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + resid := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{p.FloatAt(0) + p.FloatAt(1) - 3, p.FloatAt(0) - p.FloatAt(1) - 1}, 2) + } + pp, pv, perr := FindRootSystem(resid, probe([]float64{1, 1}, 2), RootSystemOptions{}) + bp, bv, berr := FindRootSystem(resid, base([]float64{1, 1}, 2), RootSystemOptions{}) + opPoint(t, "FindRootSystem", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseLinear", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + pp, pv, perr := MinimiseLinear(probe([]float64{2, 3}, 2), probe([]float64{1, 1}, 1, 2), probe([]float64{4}, 1), LinearProgramOptions{}) + bp, bv, berr := MinimiseLinear(base([]float64{2, 3}, 2), base([]float64{1, 1}, 1, 2), base([]float64{4}, 1), LinearProgramOptions{}) + opPoint(t, "MinimiseLinear", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseLinearRows", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + pcons := LinearConstraints{A: probe([]float64{1, 1}, 1, 2), Lower: []float64{4}, Upper: []float64{4}} + bcons := LinearConstraints{A: base([]float64{1, 1}, 1, 2), Lower: []float64{4}, Upper: []float64{4}} + pp, pv, perr := MinimiseLinearRows(probe([]float64{2, 3}, 2), pcons, LinearProgramOptions{}) + bp, bv, berr := MinimiseLinearRows(base([]float64{2, 3}, 2), bcons, LinearProgramOptions{}) + opPoint(t, "MinimiseLinearRows", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseCMAES", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + f := quadAround([]float64{1, 1}) + opts := CMAESOptions{Seed: 11, Generations: 60} + pp, pv, perr := MinimiseCMAES(f, probe([]float64{2, 3}, 2), opts) + bp, bv, berr := MinimiseCMAES(f, base([]float64{2, 3}, 2), opts) + opPoint(t, "MinimiseCMAES", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseSimulatedAnnealing", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + f := quadAround([]float64{1, 1}) + opts := SimulatedAnnealingOptions{Seed: 11, Steps: 300} + pp, pv, perr := MinimiseSimulatedAnnealing(f, probe([]float64{2, 3}, 2), opts) + bp, bv, berr := MinimiseSimulatedAnnealing(f, base([]float64{2, 3}, 2), opts) + opPoint(t, "MinimiseSimulatedAnnealing", dt, pp, pv, perr, bp, bv, berr) + }}, + {"MinimiseDifferentialEvolution", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + f := quadAround([]float64{1, 1}) + opts := DifferentialEvolutionOptions{Seed: 11, Population: 10, Generations: 40} + pp, pv, perr := MinimiseDifferentialEvolution(f, probe([]float64{0, 0}, 2), probe([]float64{4, 4}, 2), opts) + bp, bv, berr := MinimiseDifferentialEvolution(f, base([]float64{0, 0}, 2), base([]float64{4, 4}, 2), opts) + opPoint(t, "MinimiseDifferentialEvolution", dt, pp, pv, perr, bp, bv, berr) + }}, + // The standing complex refusals keep their wording; every + // probe dtype passes the same gates Int passes. + {"Minimise complex refusal preserved", func(t *testing.T, probe, base opMaker, dt core.Dtype) { + if dt != core.Int8 { + return + } + cx, err := core.FromComplexes([]complex128{1, 1}, 2) + if err != nil { + t.Fatal(err) + } + _, _, cerr := Minimise(quadAround([]float64{1, 1}), cx, MinimiseOptions{}) + opWantErr(t, "Minimise complex", cerr, "Minimise", "complex starting points are not supported") + _, _, lerr := MinimiseLBFGS(quadAround([]float64{1, 1}), nil, cx, LBFGSOptions{}) + opWantErr(t, "MinimiseLBFGS complex", lerr, "MinimiseLBFGS", "complex starting points are not supported") + _, _, merr := LevenbergMarquardt(func(p *core.Array) (*core.Array, error) { return p, nil }, cx, LMOptions{}) + opWantErr(t, "LevenbergMarquardt complex", merr, "LevenbergMarquardt", "complex parameters are not supported") + }}, + } + for _, row := range rows { + for _, dt := range opDtypes { + t.Run(row.name+"/"+dt.String(), func(t *testing.T) { + probe, base := opMakers(t, dt) + row.run(t, probe, base, dt) + }) + } + } +} diff --git a/optim/example_test.go b/optim/example_test.go new file mode 100644 index 0000000..bc5f35c --- /dev/null +++ b/optim/example_test.go @@ -0,0 +1,219 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim_test + +import ( + "fmt" + "math" + + "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/optim" +) + +// ExampleFindRoot brackets the root of a cubic between 2 and 3. The +// bracket changes sign, which is all Brent's method needs. +func ExampleFindRoot() { + f := func(x float64) float64 { return x*x*x - 2*x - 5 } + root, err := optim.FindRoot(f, 2, 3, 1e-12) + if err != nil { + fmt.Println("root finding failed:", err) + return + } + fmt.Printf("root = %.10f\n", root) + fmt.Printf("residual = %.3e\n", math.Abs(f(root))) + // Output: + // root = 2.0945514815 + // residual = 3.553e-15 +} + +// ExampleFindRootSystem solves a coupled pair of equations, +// x0 + x1 = 3 and x0² − x1 = 1, from the starting guess (1.5, 1.5). +func ExampleFindRootSystem() { + system := func(x *tensor.Array) (*tensor.Array, error) { + x0, x1 := x.FloatAt(0), x.FloatAt(1) + return tensor.FromFloats([]float64{x0 + x1 - 3, x0*x0 - x1 - 1}, 2) + } + x0, _ := tensor.FromFloats([]float64{1.5, 1.5}, 2) + x, residual, err := optim.FindRootSystem(system, x0, optim.RootSystemOptions{}) + if err != nil { + fmt.Println("system solve failed:", err) + return + } + fmt.Printf("x = (%.6f, %.6f)\n", x.FloatAt(0), x.FloatAt(1)) + fmt.Printf("residual = %.3e\n", residual) + // Output: + // x = (1.561553, 1.438447) + // residual = 4.841e-14 +} + +// ExampleLevenbergMarquardt fits the exponential decay y = a·exp(−b·t) +// to seven perturbed samples, recovering both parameters from the +// samples alone. The generator is seeded, so the data and the fit are +// reproduced exactly on every run. +func ExampleLevenbergMarquardt() { + const wantA, wantB = 2.5, 1.3 + g := tensor.NewGenerator(7) + ts := make([]float64, 7) + ys := make([]float64, 7) + for i := range ts { + ts[i] = 0.5 * float64(i) + ys[i] = wantA*math.Exp(-wantB*ts[i]) + 0.02*g.NormalUnit() + } + + // The residual the fit drives to zero holds one entry per + // observation: the model at the parameters minus the sample. + 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) - ys[i] + } + return tensor.FromFloats(out, len(out)) + } + + p0, _ := tensor.FromFloats([]float64{1, 1}, 2) + p, chi2, err := optim.LevenbergMarquardt(residual, p0, optim.LMOptions{}) + if err != nil { + fmt.Println("fit failed:", err) + return + } + fmt.Printf("a = %.4f (want %.4f)\n", p.FloatAt(0), wantA) + fmt.Printf("b = %.4f (want %.4f)\n", p.FloatAt(1), wantB) + fmt.Printf("chi2 = %.6f\n", chi2) + // Output: + // a = 2.5325 (want 2.5000) + // b = 1.2978 (want 1.3000) + // chi2 = 0.000404 +} + +// ExampleMinimiseLBFGS minimises the Rosenbrock function inside a box, +// from a start that lies outside it: the iterate is projected onto the +// box rather than refused. +func ExampleMinimiseLBFGS() { + rosenbrock := func(x *tensor.Array) (float64, error) { + xx, yy := x.FloatAt(0), x.FloatAt(1) + return 100*(yy-xx*xx)*(yy-xx*xx) + (1-xx)*(1-xx), nil + } + gradient := func(x *tensor.Array) (*tensor.Array, error) { + xx, yy := x.FloatAt(0), x.FloatAt(1) + return tensor.FromFloats([]float64{ + -400*xx*(yy-xx*xx) - 2*(1-xx), + 200 * (yy - xx*xx), + }, 2) + } + x0, _ := tensor.FromFloats([]float64{-3, 5}, 2) + opts := optim.LBFGSOptions{ + Lower: []float64{-2, -1}, + Upper: []float64{2, 3}, + } + x, f, err := optim.MinimiseLBFGS(rosenbrock, gradient, x0, opts) + if err != nil { + fmt.Println("minimisation failed:", err) + return + } + fmt.Printf("x = (%.4f, %.4f)\n", x.FloatAt(0), x.FloatAt(1)) + fmt.Printf("f = %.3e\n", f) + // Output: + // x = (1.0000, 1.0000) + // f = 3.369e-21 +} + +// ExampleMinimiseConstrained minimises the distance to (1, 1) subject +// to the equality row x0 + x1 = 1, which the unconstrained answer +// violates. The augmented Lagrangian pulls the iterate onto the row. +func ExampleMinimiseConstrained() { + objective := func(x *tensor.Array) (float64, error) { + d0, d1 := x.FloatAt(0)-1, x.FloatAt(1)-1 + return d0*d0 + d1*d1, nil + } + gradient := func(x *tensor.Array) (*tensor.Array, error) { + return tensor.FromFloats([]float64{2 * (x.FloatAt(0) - 1), 2 * (x.FloatAt(1) - 1)}, 2) + } + a, _ := tensor.FromFloats([]float64{1, 1}, 1, 2) + cons := optim.LinearConstraints{ + A: a, + Lower: []float64{1}, + Upper: []float64{1}, + } + x0, _ := tensor.FromFloats([]float64{0, 0}, 2) + x, f, err := optim.MinimiseConstrained(objective, gradient, x0, cons, optim.LBFGSOptions{}) + if err != nil { + fmt.Println("minimisation failed:", err) + return + } + fmt.Printf("x = (%.4f, %.4f)\n", x.FloatAt(0), x.FloatAt(1)) + fmt.Printf("x0+x1 = %.4f\n", x.FloatAt(0)+x.FloatAt(1)) + fmt.Printf("f = %.4f\n", f) + // Output: + // x = (0.5000, 0.5000) + // x0+x1 = 1.0000 + // f = 0.5000 +} + +// ExampleMinimiseLinearRows solves a linear program in the house +// two-sided form: maximise 3x0 + 2x1, written as the minimisation of +// its negation, subject to x0 + x1 ≤ 4, x0 + 3x1 ≤ 6 and x0, x1 ≥ 0. +func ExampleMinimiseLinearRows() { + cost, _ := tensor.FromFloats([]float64{-3, -2}, 2) + a, _ := tensor.FromFloats([]float64{ + 1, 1, + 1, 3, + 1, 0, + 0, 1, + }, 4, 2) + cons := optim.LinearConstraints{ + A: a, + Lower: []float64{ + math.Inf(-1), // -∞ ≤ x0 + x1 + math.Inf(-1), // -∞ ≤ x0 + 3x1 + 0, // 0 ≤ x0 + 0, // 0 ≤ x1 + }, + Upper: []float64{ + 4, 6, // x0 + x1 ≤ 4, x0 + 3x1 ≤ 6 + math.Inf(1), math.Inf(1), + }, + } + x, value, err := optim.MinimiseLinearRows(cost, cons, optim.LinearProgramOptions{}) + if err != nil { + fmt.Println("linear program failed:", err) + return + } + fmt.Printf("x = (%.4f, %.4f)\n", x.FloatAt(0), x.FloatAt(1)) + fmt.Printf("c·x = %.4f\n", value) + // Output: + // x = (4.0000, 0.0000) + // c·x = -12.0000 +} + +// ExampleMinimiseDifferentialEvolution finds the global minimum of the +// two-dimensional Rastrigin function, a landscape of many local minima +// around one global minimum at the origin. The seed makes the run +// reproducible. +func ExampleMinimiseDifferentialEvolution() { + rastrigin := func(x *tensor.Array) (float64, error) { + sum := 20.0 + for i := range 2 { + v := x.FloatAt(i) + sum += v*v - 10*math.Cos(2*math.Pi*v) + } + return sum, nil + } + lower, _ := tensor.FromFloats([]float64{-5.12, -5.12}, 2) + upper, _ := tensor.FromFloats([]float64{5.12, 5.12}, 2) + opts := optim.DifferentialEvolutionOptions{Seed: 7, Population: 60, Generations: 400} + x, f, err := optim.MinimiseDifferentialEvolution(rastrigin, lower, upper, opts) + if err != nil { + fmt.Println("global search failed:", err) + return + } + // The search reaches the origin. The residual coordinate is a + // round-off of the population's spread, so its sign is not fixed + // by the algorithm; the distance to the minimum is what the run + // guarantees. + fmt.Printf("f = %.6f\n", f) + fmt.Printf("|x| < 1e-6: %v\n", math.Hypot(x.FloatAt(0), x.FloatAt(1)) < 1e-6) + // Output: + // f = 0.000000 + // |x| < 1e-6: true +} diff --git a/optim/fallback_guard_pins_test.go b/optim/fallback_guard_pins_test.go new file mode 100644 index 0000000..4721591 --- /dev/null +++ b/optim/fallback_guard_pins_test.go @@ -0,0 +1,194 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "go/ast" + "go/parser" + "go/token" + "math" + "os" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Fallback and diagnosis pins: the singular-Jacobian fallback of +// FindRootSystem computing its direction from the factored matrix, a +// NaN residual dying as a bogus damping diagnosis in +// LevenbergMarquardt, the duplicate ArrayFromFloatsSafe export +// shadowed by the facade's linalg copy, the doubled "Minimise:" error +// prefix and the dead wrapVector helper. + +// TestFindRootSystemSingularFallbackDescends runs the case: +// r = [x−1, (x−1)−3(x−1)²] has a rank-one Jacobian with a zero second +// column at every point, so every Newton system is singular and the +// steepest-descent fallback carries the whole iteration. The true +// −Jᵀr at the start (2, 0) is (−11, 0), straight downhill to the root +// (1, 0); a direction formed from the factored matrix points the other +// way and the run used to burn its budget with a frozen iterate. +func TestFindRootSystemSingularFallbackDescends(t *testing.T) { + residual := func(x *core.Array) (*core.Array, error) { + d := x.FloatAt(0) - 1 + return mustFloats(t, []float64{d, d - 3*d*d}), nil + } + x, res, err := FindRootSystem(residual, mustFloats(t, []float64{2, 0}), RootSystemOptions{}) + if err != nil { + t.Fatalf("FindRootSystem: %v", err) + } + if res > 1e-10 { + t.Fatalf("residual %g, want ≤ 1e-10", res) + } + if math.Abs(x.FloatAt(0)-1) > 1e-9 || math.Abs(x.FloatAt(1)) > 1e-9 { + t.Fatalf("solution = (%.12g, %.12g), want (1, 0)", x.FloatAt(0), x.FloatAt(1)) + } +} + +// TestSteepestDescentStepDescent checks the extracted fallback at unit +// level: the direction it writes is a descent direction of ‖r‖², which +// the linearised model makes exact, and the same formula applied to the +// matrix base.Factor has consumed is not one, which is the defect the +// pristine copy removes. +func TestSteepestDescentStepDescent(t *testing.T) { + // The Jacobian at (2, 0): rank one, columns [1, −5]ᵀ + // and zero; r = (1, −2). + jac := [][]float64{{1, 0}, {-5, 0}} + r := []float64{1, -2} + step := make([]float64, 2) + steepestDescentStep(jac, r, step) + + // The directional derivative of ‖r + α·J·step‖² at α = 0 is + // 2·rᵀ(J·step): negative means the direction descends the + // residual norm, and one small α then really lowers it. + apply := func(dir []float64) []float64 { + out := make([]float64, len(r)) + for k := range jac { + s := 0.0 + for i := range dir { + s += jac[k][i] * dir[i] + } + out[k] = s + } + return out + } + slope := 0.0 + for k, jv := range apply(step) { + slope += r[k] * jv + } + if slope >= 0 { + t.Fatalf("fallback direction has residual slope %g, want < 0", slope) + } + const alpha = 0.05 + before, after := 0.0, 0.0 + jstep := apply(step) + for k := range r { + before += r[k] * r[k] + v := r[k] + alpha*jstep[k] + after += v * v + } + if after >= before { + t.Fatalf("linearised residual norm rose from %g to %g, want a fall", before, after) + } + + // The factored matrix must not be mistaken for the Jacobian: + // Factor pivots the rows and stores the multipliers on the + // subdiagonal in place, and the direction from that matrix points + // uphill here. + factored := [][]float64{{1, 0}, {-5, 0}} + base.Factor(factored) + bad := make([]float64, 2) + for i := range bad { + s := 0.0 + for k := range r { + s += factored[k][i] * r[k] + } + bad[i] = -s + } + badSlope := 0.0 + for k, jv := range apply(bad) { + badSlope += r[k] * jv + } + if badSlope <= 0 { + t.Fatalf("factored-matrix direction has slope %g, want > 0 (the defect the copy fixes)", badSlope) + } +} + +// TestLevenbergMarquardtNonFiniteResidual pins the diagnosis: a +// residual that returns NaN or Inf is reported as the non-finite value +// it is, with its index, and never misread as a damping collapse. +func TestLevenbergMarquardtNonFiniteResidual(t *testing.T) { + nan := func(*core.Array) (*core.Array, error) { + return mustFloats(t, []float64{math.NaN(), math.NaN()}), nil + } + _, _, err := LevenbergMarquardt(nan, mustFloats(t, []float64{0, 0}), LMOptions{}) + if err == nil { + t.Fatal("NaN residual: want an error") + } + if msg := err.Error(); !strings.Contains(msg, "the residual returned the non-finite value NaN at 0") { + t.Fatalf("NaN residual misdiagnosed: %v", err) + } else if strings.Contains(msg, "damping") { + t.Fatalf("NaN residual reported as a damping collapse: %v", err) + } + inf := func(*core.Array) (*core.Array, error) { + return mustFloats(t, []float64{1, math.Inf(1)}), nil + } + _, _, err = LevenbergMarquardt(inf, mustFloats(t, []float64{0, 0}), LMOptions{}) + if err == nil { + t.Fatal("Inf residual: want an error") + } + if msg := err.Error(); !strings.Contains(msg, "the residual returned the non-finite value +Inf at 1") { + t.Fatalf("Inf residual misdiagnosed: %v", err) + } +} + +// TestMinimiseNonFiniteSinglePrefix pins the exact wording: the +// non-finite objective carries one "Minimise:" prefix, because the +// inner error is bare and each call site of eval wraps once. +func TestMinimiseNonFiniteSinglePrefix(t *testing.T) { + f := func(*core.Array) (float64, error) { return math.NaN(), nil } + _, _, err := Minimise(f, mustFloats(t, []float64{0, 0}), MinimiseOptions{}) + if err == nil { + t.Fatal("NaN objective: want an error") + } + want := "tensor: Minimise: f returned the non-finite value NaN at [0 0]" + if err.Error() != want { + t.Fatalf("error = %q, want %q", err.Error(), want) + } +} + +// TestOptimSourcesDeclareNoMovedOrDeadHelpers guards the public +// surface: ArrayFromFloatsSafe lives in linalg alone (the facade +// forwards the name there), and the dead wrapVector wrapper is gone. +func TestOptimSourcesDeclareNoMovedOrDeadHelpers(t *testing.T) { + banned := map[string]string{ + "ArrayFromFloatsSafe": "a bit-identical duplicate of linalg.ArrayFromFloatsSafe, which the facade already exports", + "wrapVector": "a dead helper with no caller in the package", + } + entries, err := os.ReadDir(".") + if err != nil { + t.Fatalf("ReadDir: %v", err) + } + fset := token.NewFileSet() + for _, e := range entries { + name := e.Name() + if e.IsDir() || !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") { + continue + } + file, perr := parser.ParseFile(fset, name, nil, 0) + if perr != nil { + t.Fatalf("parse %s: %v", name, perr) + } + for _, decl := range file.Decls { + fn, ok := decl.(*ast.FuncDecl) + if !ok || fn.Recv != nil { + continue + } + if why, bad := banned[fn.Name.Name]; bad { + t.Errorf("%s declares %s again: %s", name, fn.Name.Name, why) + } + } + } +} diff --git a/optim/global2.go b/optim/global2.go new file mode 100644 index 0000000..645c81d --- /dev/null +++ b/optim/global2.go @@ -0,0 +1,600 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "cmp" + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/linalg" +) + +// Two derivative-free global searchers beside the differential +// evolution in devolution.go, both drawing from the house xoshiro +// generator so a fixed seed fixes the whole trajectory. +// +// CMA-ES (the covariance matrix adaptation evolution strategy) is the +// choice for smooth continuous landscapes where a local method cannot +// be trusted to find the right basin: it adapts a full covariance +// matrix from the successful offspring, which turns it along the +// valley whatever orientation the valley has. The implementation is a +// bounded single run of the standard algorithm with the default +// parameters of N. Hansen, The CMA Evolution Strategy: A Tutorial, +// arXiv:1604.00772 (2016): rank-one and rank-mu updates, the evolution +// paths p_sigma and p_c with their tag for the stalled-path case, and +// the tutorial's equations (48) to (53) for the weights and rates. +// No restarts: the restart schemes (IPOP and friends) are the +// caller's loop, and this entry reports one run honestly. +// +// Simulated annealing is the choice for landscapes too rough, too +// discrete-like or too deceptive for covariance adaptation: a random +// walk that accepts uphill steps with the Metropolis probability and +// cools the acceptance threshold geometrically. It is a basin finder, +// not a precision optimiser. + +// CMAESOptions tunes MinimiseCMAES. Sigma0 ≤ 0 means 0.3 (the +// tutorial's typical starting step for problems scaled to O(1)), +// Generations ≤ 0 means 500, Tolerance ≤ 0 means 1e-12, Seed 0 is +// replaced by 42 as in MinimiseDifferentialEvolution (any other +// value, negatives included, seeds the xoshiro stream directly). +// +// The tolerance ends the run when either the distribution has +// collapsed or the landscape has gone flat: the largest principal axis +// of the search distribution, sigma·sqrt(max C_ii), has fallen to +// Tolerance·max(1, ‖mean‖∞), or the objective spread over one +// generation has fallen to Tolerance·max(1, |best|). The axis is the +// diagonal of C, a proxy for the largest eigenvalue that avoids a +// second decomposition; a run that stops on the spread criterion on a +// genuinely flat landscape reports what it saw. +type CMAESOptions struct { + Sigma0 float64 + Generations int + Tolerance float64 + Seed int64 + // AllowBudgetExit makes a run that exhausts Generations report its + // best point instead of an error. The default is false, so a + // budget stop is never mistaken for a converged answer; the flag + // mirrors LBFGSOptions.AllowBudgetExit. + AllowBudgetExit bool +} + +// MinimiseCMAES returns the point and value of the best solution f the +// strategy found. f receives candidate points as rank-1 arrays; a +// non-finite value or an error is fatal for the run. There are no +// bounds: CMA-ES is an unconstrained method, and the caller who needs +// a box should reparametrise (or use the differential-evolution entry, +// which clamps). +// +// Termination, budget and honesty: a run ends when the tolerance +// criteria above are met (converged) or when the generation budget is +// spent. Following the house budget policy, an exhausted budget is an +// error naming the best value reached and the tolerance it fell short +// of, never a silent answer; AllowBudgetExit is the documented escape +// hatch. +func MinimiseCMAES(f func(*core.Array) (float64, error), x0 *core.Array, opts CMAESOptions) (*core.Array, float64, error) { + const name = "MinimiseCMAES" + if x0.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex starting points are not supported", name) + } + n := x0.Len() + if n == 0 { + return nil, 0, base.Errf("%s: the starting point must have at least one element", name) + } + sigma := opts.Sigma0 + if sigma <= 0 { + sigma = 0.3 + } + generations := opts.Generations + if generations <= 0 { + generations = 500 + } + tol := opts.Tolerance + if tol <= 0 { + tol = 1e-12 + } + seed := opts.Seed + if seed == 0 { + seed = 42 + } + g := core.NewGenerator(seed) + + // The tutorial's default parameters, equations (48) to (53). + lambda := 4 + int(3*math.Log(float64(n))) + mu := lambda / 2 + weights := make([]float64, mu) + wSum := 0.0 + for i := range mu { + weights[i] = math.Log(float64(lambda)/2+0.5) - math.Log(float64(i+1)) + wSum += weights[i] + } + for i := range mu { + weights[i] /= wSum + } + muEff := 0.0 + for _, w := range weights { + muEff += w * w + } + muEff = 1 / muEff + cSigma := (muEff + 2) / (float64(n) + muEff + 5) + dSigma := 1 + 2*math.Max(0, math.Sqrt((muEff-1)/(float64(n)+1))-1) + cSigma + cC := (4 + muEff/float64(n)) / (4 + float64(n) + 2*muEff/float64(n)) + c1 := 2 / ((float64(n)+1.3)*(float64(n)+1.3) + muEff) + cMu := math.Min(1-c1, 2*(muEff-2+1/muEff)/((float64(n)+2)*(float64(n)+2)+muEff)) + chiN := math.Sqrt(float64(n)) * (1 - 1/(4*float64(n)) + 1/(21*float64(n)*float64(n))) + + mean := cloneDense(x0) + cov := make([]float64, n*n) + for i := range n { + cov[i*n+i] = 1 + } + pathC := make([]float64, n) + pathS := make([]float64, n) + + eval := func(p []float64) (float64, error) { + v, err := f(linalg.ArrayFromFloatsSafe(p, len(p))) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, base.Errf("%s: the objective is non-finite (%g)", name, v) + } + return v, nil + } + + bestF, bestX := math.Inf(1), make([]float64, n) + offX := make([]float64, lambda*n) + offY := make([]float64, lambda*n) + offF := make([]float64, lambda) + order := make([]int, lambda) + yW := make([]float64, n) + yC := make([]float64, n) + newCov := make([]float64, n*n) + // Per-generation scratch, owned by the run: the sampling scale per + // principal axis, the unit normal drawn per offspring, the + // eigenvector projection rootInverse accumulates in, the rank-mu + // outer-product accumulator with its weighted axis vector, and the + // diagonalisation's own working set. Each is fully overwritten + // before it is read. + sd := make([]float64, n) + z := make([]float64, n) + yInv := make([]float64, n) + rankMu := make([]float64, n*n) + wy := make([]float64, n) + var eig jacobiScratch + + for gen := range generations { + vals, vecs := eig.eigen(cov, n) + // Sampling scale per principal axis, floored at zero: a + // numerically degenerate axis contributes nothing rather than + // a complex square root. + for i := range n { + sd[i] = math.Sqrt(math.Max(vals[i], 0)) + } + worstF := math.Inf(-1) + bestGen := math.Inf(1) + for k := range lambda { + // One unit normal per principal direction: the draw is + // B·diag(sd)·z, whose covariance is exactly C. Sharing a + // single scalar across the directions would sample a + // diagonal distribution scaled by one fixed vector and + // leave the adapted covariance unused. + for i := range n { + z[i] = g.NormalUnit() + } + cmaDraw(vecs, sd, z, offY[k*n:k*n+n]) + for i := range n { + offX[k*n+i] = mean[i] + sigma*offY[k*n+i] + } + v, err := eval(offX[k*n : k*n+n]) + if err != nil { + return nil, 0, err + } + offF[k] = v + worstF = math.Max(worstF, v) + bestGen = math.Min(bestGen, v) + } + for k := range lambda { + order[k] = k + } + slices.SortStableFunc(order, func(a, b int) int { return cmp.Compare(offF[a], offF[b]) }) + if offF[order[0]] < bestF { + bestF = offF[order[0]] + copy(bestX, offX[order[0]*n:(order[0]+1)*n]) + } + clear(yW) + for k := range mu { + w := weights[k] + for i := range n { + yW[i] += w * offY[order[k]*n+i] + } + } + for i := range n { + mean[i] += sigma * yW[i] + } + + // Conjugate evolution path: C^{-1/2} y_w through the + // eigendecomposition, then the sigma update from its length. + rootInverse(yC, yW, vals, vecs, yInv, n) + psNorm := 0.0 + for i := range n { + pathS[i] = (1-cSigma)*pathS[i] + math.Sqrt(cSigma*(2-cSigma)*muEff)*yC[i] + psNorm += pathS[i] * pathS[i] + } + psNorm = math.Sqrt(psNorm) + sigma *= math.Exp((cSigma / dSigma) * (psNorm/chiN - 1)) + hs := 0.0 + if psNorm/math.Sqrt(1-math.Pow(1-cSigma, 2*float64(gen+1))) < (1.4+2/(float64(n)+1))*chiN { + hs = 1 + } + pcNormScale := math.Sqrt(cC * (2 - cC) * muEff) + for i := range n { + pathC[i] = (1-cC)*pathC[i] + hs*pcNormScale*yW[i] + } + // The rank-one and rank-mu updates; delta(hs) keeps the + // covariance from growing along p_c across a stall of the + // conjugate path. The rank-mu sum accumulates as contiguous + // outer products, one per selected offspring walked in order: + // every cell still sums (w·yᵢ)·yⱼ over ascending k, so the + // bits are the per-cell walk's and the inner loop stays on + // unit stride. + clear(rankMu) + for k := range mu { + base := order[k] * n + w := weights[k] + for i := range n { + wy[i] = w * offY[base+i] + } + for i := range n { + wi := wy[i] + row := rankMu[i*n : i*n+n] + yk := offY[base : base+n] + for j := range n { + row[j] += wi * yk[j] + } + } + } + factor := 1 + c1*(1-hs) - c1 - cMu + for i := range n { + ci := c1 * pathC[i] + off := i * n + for j := range n { + newCov[off+j] = factor*cov[off+j] + ci*pathC[j] + cMu*rankMu[off+j] + } + } + for i := range n { + for j := i + 1; j < n; j++ { + avg := (newCov[i*n+j] + newCov[j*n+i]) / 2 + newCov[i*n+j], newCov[j*n+i] = avg, avg + } + } + copy(cov, newCov) + + axis := 0.0 + for i := range n { + axis = math.Max(axis, sigma*math.Sqrt(math.Max(cov[i*n+i], 0))) + } + if math.IsNaN(axis) || sigma > 1e12 { + return nil, 0, base.Errf("%s: the search diverged (sigma %g, largest axis %g) at f = %g", name, sigma, axis, bestF) + } + if axis <= tol*math.Max(1, maxAbs(mean)) || worstF-bestGen <= tol*math.Max(1, math.Abs(bestF)) { + out, fv := packResult(bestX, bestF) + return out, fv, nil + } + } + if !opts.AllowBudgetExit { + return nil, 0, base.Errf("%s: the generation budget of %d ran out at f = %g, above the tolerance %g", name, generations, bestF, tol) + } + out, fv := packResult(bestX, bestF) + return out, fv, nil +} + +// rootInverse writes C^{-1/2} y into dst through the eigendecomposition +// the caller already paid for: B·diag(1/sqrt(d))·Bᵀ·y with the +// eigenvalues floored away from zero so a numerically flat direction +// cannot divide by nothing. tmp is scratch of length at least n, fully +// overwritten before it is read. +func rootInverse(dst, y, vals, vecs, tmp []float64, n int) { + scale := 0.0 + for i := range n { + scale = math.Max(scale, vals[i]) + } + floor := 1e-20 * math.Max(1, scale) + for j := range n { // tmp = Bᵀ y + s := 0.0 + for i := range n { + s += vecs[j*n+i] * y[i] + } + tmp[j] = s / math.Sqrt(math.Max(vals[j], floor)) + } + clear(dst) + for j := range n { // dst = B tmp + c := tmp[j] + if c == 0 { + continue + } + for i := range n { + dst[i] += vecs[j*n+i] * c + } + } +} + +// eigenPair is one eigenvalue with the column its eigenvector occupies +// in the accumulator the diagonalisation carries. +type eigenPair struct { + val float64 + index int +} + +// jacobiScratch is one diagonalisation's working set: the cyclically +// rotated copy of the matrix, the eigenvector accumulator, the +// value/index pairs the descending sort walks and the two blocks handed +// back. The driver keeps one instance for the whole run, so the +// per-generation decomposition allocates nothing; every buffer is fully +// overwritten before it is read. +type jacobiScratch struct { + work []float64 + acc []float64 + pairs []eigenPair + vals []float64 + sorted []float64 + evecs []float64 +} + +// jacobiEigen diagonalises the symmetric row-major n×n matrix into +// freshly allocated blocks, the allocating entry over +// (*jacobiScratch).eigen. +func jacobiEigen(a []float64, n int) (vals, vecs []float64) { + var s jacobiScratch + return s.eigen(a, n) +} + +// eigen diagonalises the symmetric row-major n×n matrix by cyclic +// Jacobi rotations: the eigenvalues come back in descending order and +// vecs[j*n+i] is component i of the eigenvector that belongs to +// vals[j]. The sweeps stop once the off-diagonal mass has fallen to +// the working precision of the matrix's own Frobenius norm. Both +// returned blocks belong to the scratch and stay valid until its next +// call. +func (s *jacobiScratch) eigen(a []float64, n int) (vals, vecs []float64) { + if cap(s.work) < n*n { + s.work = make([]float64, n*n) + } + work := s.work[:n*n] + copy(work, a) + if cap(s.acc) < n*n { + s.acc = make([]float64, n*n) + } + acc := s.acc[:n*n] + clear(acc) + for i := range n { + acc[i*n+i] = 1 + } + frob := 0.0 + for _, v := range work { + frob += v * v + } + const maxSweeps = 60 + for range maxSweeps { + off := 0.0 + for i := range n { + for j := i + 1; j < n; j++ { + off += work[i*n+j] * work[i*n+j] + } + } + if off <= 1e-30*math.Max(frob, 1) { + break + } + for p := range n { + for q := p + 1; q < n; q++ { + apq := work[p*n+q] + if apq == 0 { + continue + } + theta := (work[q*n+q] - work[p*n+p]) / (2 * apq) + t := 1 / (math.Abs(theta) + math.Sqrt(theta*theta+1)) + if theta < 0 { + t = -t + } + c := 1 / math.Sqrt(t*t+1) + s := t * c + for k := range n { + akp, akq := work[k*n+p], work[k*n+q] + work[k*n+p] = c*akp - s*akq + work[k*n+q] = s*akp + c*akq + } + for k := range n { + apk, aqk := work[p*n+k], work[q*n+k] + work[p*n+k] = c*apk - s*aqk + work[q*n+k] = s*apk + c*aqk + } + work[p*n+q], work[q*n+p] = 0, 0 + for k := range n { + vkp, vkq := acc[k*n+p], acc[k*n+q] + acc[k*n+p] = c*vkp - s*vkq + acc[k*n+q] = s*vkp + c*vkq + } + } + } + } + if cap(s.vals) < n { + s.vals = make([]float64, n) + } + vals = s.vals[:n] + for i := range n { + vals[i] = work[i*n+i] + } + // Sort the pairs descending by value. The accumulator carries + // eigenvector k in its column k, and the documented layout here is + // row j for vector j, so the extraction transposes. + if cap(s.pairs) < n { + s.pairs = make([]eigenPair, n) + } + pairs := s.pairs[:n] + for i := range n { + pairs[i] = eigenPair{vals[i], i} + } + slices.SortFunc(pairs, func(a, b eigenPair) int { return cmp.Compare(b.val, a.val) }) + if cap(s.sorted) < n { + s.sorted = make([]float64, n) + } + if cap(s.evecs) < n*n { + s.evecs = make([]float64, n*n) + } + sortedVals := s.sorted[:n] + sortedVecs := s.evecs[:n*n] + for j, pr := range pairs { + sortedVals[j] = pr.val + for i := range n { + sortedVecs[j*n+i] = acc[i*n+pr.index] + } + } + return sortedVals, sortedVecs +} + +// SimulatedAnnealingOptions tunes MinimiseSimulatedAnnealing. +// Steps ≤ 0 means 20000 proposals, Temperature0 ≤ 0 means 1, +// CoolingRate ≤ 0 means 0.9995 (the geometric factor per proposal), +// StepScale ≤ 0 means 0.1 (the relative Gaussian proposal scale), +// Tolerance ≤ 0 means 1e-6, Seed 0 is replaced by 42 as in +// MinimiseDifferentialEvolution. +// +// The schedule is the algorithm: the temperature runs from +// Temperature0 down by the factor CoolingRate at every proposal, so +// the chain freezes exponentially and the tail of the schedule is a +// local polish. Tolerance is the convergence test on that tail: the +// best value may not improve by more than Tolerance·max(1, |best|) +// over the final quarter of the schedule. A run that is still +// improving at the end of the schedule has not converged, and the +// budget refusal says so with both figures, exactly as the local +// solvers refuse an unfinished run; AllowBudgetExit is the escape +// hatch that reports the best point anyway. +type SimulatedAnnealingOptions struct { + Steps int + Temperature0 float64 + CoolingRate float64 + StepScale float64 + Seed int64 + Tolerance float64 + // AllowBudgetExit makes a run whose schedule ended while the best + // value was still improving report its best point instead of an + // error. The flag mirrors LBFGSOptions.AllowBudgetExit. + AllowBudgetExit bool +} + +// MinimiseSimulatedAnnealing returns the best point and value the +// chain found: a Metropolis walk over the landscape with a Gaussian +// proposal of the given relative scale (coordinate i moves by +// StepScale·max(1, |x_i|) standard normals) and the geometric cooling +// schedule above. f receives candidate points as rank-1 arrays; a +// non-finite value or an error is fatal for the run. Simulated +// annealing finds basins, not minima to full precision: run +// MinimiseLBFGS from the returned point when a polished answer is +// wanted, which is the standard composition. +func MinimiseSimulatedAnnealing(f func(*core.Array) (float64, error), x0 *core.Array, opts SimulatedAnnealingOptions) (*core.Array, float64, error) { + const name = "MinimiseSimulatedAnnealing" + if x0.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex starting points are not supported", name) + } + n := x0.Len() + if n == 0 { + return nil, 0, base.Errf("%s: the starting point must have at least one element", name) + } + steps := opts.Steps + if steps <= 0 { + steps = 20000 + } + t0 := opts.Temperature0 + if t0 <= 0 { + t0 = 1 + } + cooling := opts.CoolingRate + if cooling > 1 { + return nil, 0, base.Errf("%s: the cooling rate %g would heat the chain; want a factor in (0, 1]", name, cooling) + } + if cooling <= 0 { + cooling = 0.9995 + } + scale := opts.StepScale + if scale <= 0 { + scale = 0.1 + } + tol := opts.Tolerance + if tol <= 0 { + tol = 1e-6 + } + seed := opts.Seed + if seed == 0 { + seed = 42 + } + g := core.NewGenerator(seed) + + x := cloneDense(x0) + cur, err := f(linalg.ArrayFromFloatsSafe(x, n)) + if err != nil { + return nil, 0, base.Errf("%s: %w", name, err) + } + if math.IsNaN(cur) || math.IsInf(cur, 0) { + return nil, 0, base.Errf("%s: the objective is non-finite (%g) at the start", name, cur) + } + bestF, bestX := cur, slices.Clone(x) + proposal := make([]float64, n) + quarterF := math.Inf(1) + quarterAt := 3 * steps / 4 + + for k := range steps { + temperature := t0 * math.Pow(cooling, float64(k)) + for i := range n { + proposal[i] = x[i] + scale*math.Max(1, math.Abs(x[i]))*g.NormalUnit() + } + fv, ferr := f(linalg.ArrayFromFloatsSafe(proposal, n)) + if ferr != nil { + return nil, 0, base.Errf("%s: %w", name, ferr) + } + if math.IsNaN(fv) || math.IsInf(fv, 0) { + return nil, 0, base.Errf("%s: the objective is non-finite (%g) at proposal %d", name, fv, k+1) + } + delta := fv - cur + if delta <= 0 || g.Unit() < math.Exp(-delta/temperature) { + copy(x, proposal) + cur = fv + } + if cur < bestF { + bestF = cur + copy(bestX, x) + } + if k == quarterAt { + quarterF = bestF + } + } + // The schedule ended: convergence is the frozen tail, judged by + // the improvement the final quarter still bought. + if quarterF-bestF > tol*math.Max(1, math.Abs(bestF)) && !opts.AllowBudgetExit { + return nil, 0, base.Errf("%s: the schedule of %d steps ended with the best value still improving (%g to %g); raise Steps or set AllowBudgetExit", + name, steps, quarterF, bestF) + } + out, fv := packResult(bestX, bestF) + return out, fv, nil +} + +// cmaDraw fills y with B·diag(sd)·z: one unit normal per principal +// direction transformed into the covariance's own axes, the draw the +// tutorial's sampling equation defines, whose covariance is the +// adapted C itself. The walk goes one principal direction at a time so +// the eigenvector rows are read on unit stride; the per-element +// grouping (v·sdⱼ)·zⱼ and the ascending j order are the i-outer walk's, +// so the draw is bit for bit the same. +func cmaDraw(vecs, sd, z, y []float64) { + clear(y) + n := len(y) + for j := range z { + sdJ, zJ := sd[j], z[j] + row := vecs[j*n : j*n+n] + for i := range n { + y[i] += row[i] * sdJ * zJ + } + } +} diff --git a/optim/global2_test.go b/optim/global2_test.go new file mode 100644 index 0000000..b4ffae2 --- /dev/null +++ b/optim/global2_test.go @@ -0,0 +1,378 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// rotEllipsoid builds the test landscape: f(x) = Σ i·(Rx)ᵢ² for a +// fixed orthonormal R (the Householder reflection through the +// normalised all-ones vector), a rotated ellipsoid with condition 6ⁿ +// along axes the coordinate system does not see. It returns the +// objective and the apply function so tests can check coordinates. +func rotEllipsoid(t *testing.T, n int) (func(*core.Array) (float64, error), func(x []float64) []float64) { + t.Helper() + v := make([]float64, n) + for i := range n { + v[i] = float64(i + 1) + } + norm := 0.0 + for _, a := range v { + norm += a * a + } + norm = math.Sqrt(norm) + apply := func(x []float64) []float64 { + dot := 0.0 + for i := range n { + dot += v[i] * x[i] + } + dot *= 2 / (norm * norm) + rx := make([]float64, n) + for i := range n { + rx[i] = x[i] - dot*v[i] + } + return rx + } + f := func(a *core.Array) (float64, error) { + total := 0.0 + for i, rx := range apply(floatsOf(a)) { + total += float64(i+1) * rx * rx + } + return total, nil + } + return f, apply +} + +// TestJacobiEigen pins the eigendecomposition the CMA-ES loop leans +// on: a known symmetric matrix with irrational eigenvalues, verified +// through the reconstruction C = B·D·Bᵀ and the orthogonality of B. +func TestJacobiEigen(t *testing.T) { + // [[2, 1], [1, 3]]: eigenvalues (5 ± √5)/2. + vals, vecs := jacobiEigen([]float64{2, 1, 1, 3}, 2) + want1 := (5 + math.Sqrt(5)) / 2 + want2 := (5 - math.Sqrt(5)) / 2 + if math.Abs(vals[0]-want1) > 1e-12 || math.Abs(vals[1]-want2) > 1e-12 { + t.Fatalf("eigenvalues = (%g, %g), want (%g, %g)", vals[0], vals[1], want1, want2) + } + // Orthonormality: B·Bᵀ = I. + for i := range 2 { + for j := range 2 { + s := 0.0 + for k := range 2 { + s += vecs[k*2+i] * vecs[k*2+j] + } + want := 0.0 + if i == j { + want = 1 + } + if math.Abs(s-want) > 1e-12 { + t.Fatalf("B·Bᵀ[%d][%d] = %g, want %g", i, j, s, want) + } + } + } + // Reconstruction: Σ_j vals[j]·vec_j·vec_jᵀ = C. + for i := range 2 { + for j := range 2 { + s := 0.0 + for k := range 2 { + s += vals[k] * vecs[k*2+i] * vecs[k*2+j] + } + want := 1.0 + if i == j { + want = float64(2 + i) + } + if math.Abs(s-want) > 1e-12 { + t.Fatalf("reconstruction[%d][%d] = %g, want %g", i, j, s, want) + } + } + } +} + +// TestMinimiseCMAESRotatedEllipsoid pins the strategy on a rotated +// ellipsoid, the landscape its covariance adaptation exists for: seed +// 7, sigma0 0.5 and a budget of 300 generations carry the run to the +// origin within 1e−6 per coordinate and 1e−12 in value. +func TestMinimiseCMAESRotatedEllipsoid(t *testing.T) { + f, _ := rotEllipsoid(t, 6) + start := mustFloats(t, []float64{1, -1, 0.5, 2, 0, -0.5}) + x, fv, err := MinimiseCMAES(f, start, CMAESOptions{Seed: 7, Sigma0: 0.5, Generations: 300}) + if err != nil { + t.Fatalf("MinimiseCMAES: %v", err) + } + if fv > 1e-12 { + t.Fatalf("value = %.3e, want <= 1e-12", fv) + } + for i := range 6 { + if math.Abs(x.FloatAt(i)) > 1e-5 { + t.Fatalf("x[%d] = %.3e, want within 1e-5 of the origin", i, x.FloatAt(i)) + } + } +} + +// TestMinimiseCMAESDeterministic pins reproducibility: two runs on one +// seed walk the same landscape with the same draws and must return +// bit-identical trajectories. +func TestMinimiseCMAESDeterministic(t *testing.T) { + f, _ := rotEllipsoid(t, 4) + start := mustFloats(t, []float64{0.5, -0.5, 1, -1}) + // Bit-identical trajectories under one seed: 80 generations do not + // exhaust to convergence, so the runs use the escape hatch and the + // pin is on the trajectories themselves. + xa, fa, errA := MinimiseCMAES(f, start, CMAESOptions{Seed: 7, Sigma0: 0.5, Generations: 80, AllowBudgetExit: true}) + xb, fb, errB := MinimiseCMAES(f, start, CMAESOptions{Seed: 7, Sigma0: 0.5, Generations: 80, AllowBudgetExit: true}) + if errA != nil || errB != nil { + t.Fatalf("MinimiseCMAES: %v, %v", errA, errB) + } + if math.Float64bits(fa) != math.Float64bits(fb) { + t.Fatalf("values differ across identical runs: %.20g vs %.20g", fa, fb) + } + for i := range 4 { + if math.Float64bits(xa.FloatAt(i)) != math.Float64bits(xb.FloatAt(i)) { + t.Fatalf("x[%d] differs across identical runs", i) + } + } +} + +// TestMinimiseCMAESBudgetAndGates pins the honest budget refusal, the +// escape hatch, the divergence refusal on an unbounded-below +// objective, and the input and objective gates. +func TestMinimiseCMAESBudgetAndGates(t *testing.T) { + f, _ := rotEllipsoid(t, 4) + start := mustFloats(t, []float64{0.5, -0.5, 1, -1}) + // Two generations cannot collapse the distribution: the budget + // stop is refused with the evidence, not reported as an answer. + _, _, err := MinimiseCMAES(f, start, CMAESOptions{Seed: 7, Sigma0: 0.5, Generations: 2}) + if err == nil || !strings.Contains(err.Error(), "budget") { + t.Fatalf("error = %v, want the budget refusal", err) + } + // The escape hatch returns the best point with no error. + xb, _, err := MinimiseCMAES(f, start, CMAESOptions{Seed: 7, Sigma0: 0.5, Generations: 2, AllowBudgetExit: true}) + if err != nil || xb == nil { + t.Fatalf("AllowBudgetExit: err = %v, x = %v", err, xb) + } + // An objective unbounded below drives sigma past the guard. + _, _, err = MinimiseCMAES(func(a *core.Array) (float64, error) { + s := 0.0 + for i := range a.Len() { + s += a.FloatAt(i) * a.FloatAt(i) + } + return -math.Sqrt(s), nil + }, start, CMAESOptions{Seed: 7, Sigma0: 0.5, Generations: 200}) + if err == nil || !strings.Contains(err.Error(), "diverged") { + t.Fatalf("error = %v, want the divergence refusal", err) + } + // Empty and complex starts are refused. + if _, _, err := MinimiseCMAES(f, core.New(core.Float, 0), CMAESOptions{}); err == nil { + t.Fatal("an empty starting point was accepted") + } + if _, _, err := MinimiseCMAES(f, mustComplexPoint(t), CMAESOptions{}); err == nil { + t.Fatal("a complex starting point was accepted") + } + // A non-finite objective is fatal. + if _, _, err := MinimiseCMAES(func(*core.Array) (float64, error) { return math.NaN(), nil }, + start, CMAESOptions{Seed: 7, Generations: 5}); err == nil { + t.Fatal("a NaN objective was accepted") + } + // An objective's own error propagates. + if _, _, err := MinimiseCMAES(func(*core.Array) (float64, error) { + return 0, base.Errf("the model exploded") + }, start, CMAESOptions{Seed: 7, Generations: 5}); err == nil { + t.Fatal("the objective's error did not propagate") + } + // The defaults carry the run when the caller passes nothing but + // the escape hatch: sigma0 0.3, 500 generations, tolerance + // 1e-12 and seed 42. + _, _, err = MinimiseCMAES(f, start, CMAESOptions{AllowBudgetExit: true}) + if err != nil { + t.Fatalf("MinimiseCMAES with defaults: %v", err) + } +} + +// TestMinimiseAnnealRotatedEllipsoid pins simulated annealing on the +// same rotated ellipsoid at basin precision: 30000 proposals reach the +// origin's basin, which for this landscape means every coordinate +// within 0.5 and a value below 6. +func TestMinimiseAnnealRotatedEllipsoid(t *testing.T) { + f, _ := rotEllipsoid(t, 6) + start := mustFloats(t, []float64{1, -1, 0.5, 2, 0, -0.5}) + x, fv, err := MinimiseSimulatedAnnealing(f, start, SimulatedAnnealingOptions{Seed: 7, Steps: 30000}) + if err != nil { + t.Fatalf("MinimiseSimulatedAnnealing: %v", err) + } + if fv > 6 { + t.Fatalf("value = %.3e, want <= 6 (the basin of the origin)", fv) + } + for i := range 6 { + if math.Abs(x.FloatAt(i)) > 0.5 { + t.Fatalf("x[%d] = %.3e, want within the unit basin", i, x.FloatAt(i)) + } + } +} + +// TestMinimiseAnnealDeterministicAndBudget pins reproducibility, the +// budget refusal of a chain that is still improving, the escape hatch +// and the input gates. +func TestMinimiseAnnealDeterministicAndBudget(t *testing.T) { + f, _ := rotEllipsoid(t, 4) + start := mustFloats(t, []float64{0.5, -0.5, 1, -1}) + // Bit-identical trajectories under one seed: 500 proposals end + // mid-polish, so the runs use the escape hatch and the pin is on + // the trajectories themselves. + xa, fa, errA := MinimiseSimulatedAnnealing(f, start, SimulatedAnnealingOptions{Seed: 7, Steps: 500, AllowBudgetExit: true}) + xb, fb, errB := MinimiseSimulatedAnnealing(f, start, SimulatedAnnealingOptions{Seed: 7, Steps: 500, AllowBudgetExit: true}) + if errA != nil || errB != nil { + t.Fatalf("MinimiseSimulatedAnnealing: %v, %v", errA, errB) + } + if math.Float64bits(fa) != math.Float64bits(fb) { + t.Fatalf("values differ across identical runs: %.20g vs %.20g", fa, fb) + } + for i := range 4 { + if math.Float64bits(xa.FloatAt(i)) != math.Float64bits(xb.FloatAt(i)) { + t.Fatalf("x[%d] differs across identical runs", i) + } + } + // A linear objective keeps producing new bests to the last + // proposal: the schedule ends unfinished and is refused. + line := func(a *core.Array) (float64, error) { + s := 0.0 + for i := range a.Len() { + s -= a.FloatAt(i) + } + return s, nil + } + _, _, err := MinimiseSimulatedAnnealing(line, start, SimulatedAnnealingOptions{Seed: 7, Steps: 2000, Tolerance: 1e-300}) + if err == nil || !strings.Contains(err.Error(), "still improving") { + t.Fatalf("error = %v, want the unfinished-schedule refusal", err) + } + // The escape hatch reports the best point anyway. + xbest, fv, err := MinimiseSimulatedAnnealing(line, start, SimulatedAnnealingOptions{Seed: 7, Steps: 2000, Tolerance: 1e-300, AllowBudgetExit: true}) + if err != nil || xbest == nil { + t.Fatalf("AllowBudgetExit: err = %v, x = %v", err, xbest) + } + if fv >= 0 { + t.Fatalf("value = %.3e, want the linear objective's negative value", fv) + } + // Empty and complex starts are refused. + if _, _, err := MinimiseSimulatedAnnealing(f, core.New(core.Float, 0), SimulatedAnnealingOptions{}); err == nil { + t.Fatal("an empty starting point was accepted") + } + if _, _, err := MinimiseSimulatedAnnealing(f, mustComplexPoint(t), SimulatedAnnealingOptions{}); err == nil { + t.Fatal("a complex starting point was accepted") + } + // A non-finite objective at the start is refused. + if _, _, err := MinimiseSimulatedAnnealing(func(*core.Array) (float64, error) { return math.Inf(-1), nil }, + start, SimulatedAnnealingOptions{}); err == nil { + t.Fatal("a non-finite start value was accepted") + } + // An objective's own error propagates. + if _, _, err := MinimiseSimulatedAnnealing(func(*core.Array) (float64, error) { + return 0, base.Errf("the model exploded") + }, start, SimulatedAnnealingOptions{}); err == nil { + t.Fatal("the objective's error did not propagate") + } + // A proposal that wanders where the objective is undefined is + // fatal, and so is one that errors: the chain climbs a linear + // slope until it crosses the model's domain. + climb := func(limit float64, verdict func() (float64, error)) func(*core.Array) (float64, error) { + return func(a *core.Array) (float64, error) { + s := 0.0 + for i := range a.Len() { + s -= a.FloatAt(i) + } + if s < limit { + return verdict() + } + return s, nil + } + } + if _, _, err := MinimiseSimulatedAnnealing(climb(-30, func() (float64, error) { return math.NaN(), nil }), + start, SimulatedAnnealingOptions{Seed: 7, Steps: 4000}); err == nil { + t.Fatal("a NaN proposal value was accepted") + } + if _, _, err := MinimiseSimulatedAnnealing(climb(-30, func() (float64, error) { + return 0, base.Errf("the proposal exploded") + }), start, SimulatedAnnealingOptions{Seed: 7, Steps: 4000}); err == nil { + t.Fatal("the proposal's error did not propagate") + } +} + +// TestCMADrawSamplesTheAdaptedCovariance pins the sampling equation: +// the candidates are B·diag(sd)·z with one unit normal per principal +// direction, so the empirical covariance of many draws reproduces the +// adapted C even after a rotation that leaves no axis aligned. A +// shared scalar across the directions sampled a diagonal distribution +// instead and failed this probe loudly. +func TestCMADrawSamplesTheAdaptedCovariance(t *testing.T) { + n := 4 + // C = Q·D·Qᵀ for a Householder reflection through (1, 2, 3, 4) and + // a spread of eigenvalues, so no axis survives the rotation. + q := make([]float64, n*n) + { + v := []float64{1, 2, 3, 4} + norm := 0.0 + for _, x := range v { + norm += x * x + } + norm = math.Sqrt(norm) + for i := range n { + for j := range n { + h := 0.0 + if i == j { + h = 1 + } + q[i*n+j] = h - 2*v[i]*v[j]/(norm*norm) + } + } + } + eigen := []float64{4, 1, 0.5, 0.25} + sd := make([]float64, n) + for i := range n { + sd[i] = math.Sqrt(eigen[i]) + } + c := make([]float64, n*n) + for i := range n { + for j := range n { + for k := range n { + c[i*n+j] += q[k*n+i] * eigen[k] * q[k*n+j] + } + } + } + g := core.NewGenerator(9) + draws := 300000 + z := make([]float64, n) + var y []float64 + sum := make([]float64, n) + cov := make([]float64, n*n) + for range draws { + for i := range n { + z[i] = g.NormalUnit() + } + y = make([]float64, n) + cmaDraw(q, sd, z, y) + for i := range n { + sum[i] += y[i] + for j := range n { + cov[i*n+j] += y[i] * y[j] + } + } + } + froC, froErr := 0.0, 0.0 + for i := range n { + for j := range n { + cov[i*n+j] = cov[i*n+j]/float64(draws) - sum[i]*sum[j]/(float64(draws)*float64(draws)) + d := cov[i*n+j] - c[i*n+j] + froErr += d * d + froC += c[i*n+j] * c[i*n+j] + } + } + if math.Sqrt(froErr/froC) > 0.02 { + t.Fatalf("empirical covariance off the adapted C by %.3g relative Frobenius", math.Sqrt(froErr/froC)) + } +} diff --git a/optim/helpers_test.go b/optim/helpers_test.go new file mode 100644 index 0000000..98c4e39 --- /dev/null +++ b/optim/helpers_test.go @@ -0,0 +1,24 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +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 +} diff --git a/optim/jacobianparallel_test.go b/optim/jacobianparallel_test.go new file mode 100644 index 0000000..a37c574 --- /dev/null +++ b/optim/jacobianparallel_test.go @@ -0,0 +1,341 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "fmt" + "math" + "runtime" + "sync/atomic" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Tests and benchmarks for the opt-in parallel sweep of the +// central-difference Jacobian, LMOptions.ParallelJacobian and +// RootSystemOptions.ParallelJacobian. The contract under test: the +// default false never calls the callback from more than one goroutine +// and answers bit for bit what it answered before; the true value may +// overlap the callback's evaluations across columns and still answers +// bit for bit the same numbers, because every column is differenced by +// the same stencil against the same point. + +// concurrencyTracker records the largest number of callbacks it has +// seen inside at once. The counter and the peak are atomic: the tracker +// itself must not be the thing that serialises the calls it measures. +type concurrencyTracker struct { + cur atomic.Int64 + peak atomic.Int64 +} + +func (c *concurrencyTracker) enter() { + n := c.cur.Add(1) + for { + p := c.peak.Load() + if n <= p || c.peak.CompareAndSwap(p, n) { + return + } + } +} + +func (c *concurrencyTracker) leave() { c.cur.Add(-1) } + +// lmParallelProblem builds an eight-parameter two-exponential plus +// sinusoid fit with a known generating model, hard enough that the fit +// walks a number of damped steps before it converges. +func lmParallelProblem() (residual func(*core.Array) (*core.Array, error), start []float64) { + const nObs = 48 + t := make([]float64, nObs) + y := make([]float64, nObs) + truth := []float64{2.0, 0.7, 1.5, 1.9, 0.8, 3.1, 0.4, 0.05} + for i := range nObs { + t[i] = float64(i) / 6 + y[i] = truth[0]*math.Exp(-truth[1]*t[i]) + truth[2]*math.Exp(-truth[3]*t[i]) + + truth[4]*math.Sin(truth[5]*t[i]+truth[6]) + truth[7]*t[i] + } + residual = func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, nObs) + vals := out.RawFloats() + for i := range nObs { + vals[i] = p.FloatAt(0)*math.Exp(-p.FloatAt(1)*t[i]) + p.FloatAt(2)*math.Exp(-p.FloatAt(3)*t[i]) + + p.FloatAt(4)*math.Sin(p.FloatAt(5)*t[i]+p.FloatAt(6)) + p.FloatAt(7)*t[i] - y[i] + } + return out, nil + } + return residual, []float64{1.6, 0.5, 1.2, 1.6, 0.6, 2.7, 0.2, 0.02} +} + +// rootParallelProblem builds eight coupled quadratic equations whose +// root sits near xᵢ = i+1, hard enough that Newton walks several rounds. +func rootParallelProblem() (f func(*core.Array) (*core.Array, error), start []float64) { + const n = 8 + f = func(x *core.Array) (*core.Array, error) { + out := core.New(core.Float, n) + vals := out.RawFloats() + for i := range n { + vals[i] = x.FloatAt(i)*x.FloatAt(i) - float64(i+1)*float64(i+1) + 0.1*x.FloatAt((i+1)%n) + } + return out, nil + } + start = make([]float64, n) + for i := range start { + start[i] = float64(i+1) + 0.5 + } + return f, start +} + +func TestParallelJacobianBitIdenticalFit(t *testing.T) { + residual, start := lmParallelProblem() + run := func(parallel bool) (*core.Array, float64) { + t.Helper() + p, chi2, err := LevenbergMarquardt(residual, mustFloats(t, start), LMOptions{ + MaxIterations: 100, + ParallelJacobian: parallel, + }) + if err != nil { + t.Fatalf("LevenbergMarquardt(ParallelJacobian=%v): %v", parallel, err) + } + return p, chi2 + } + pSerial, chi2Serial := run(false) + pParallel, chi2Parallel := run(true) + if math.Float64bits(chi2Serial) != math.Float64bits(chi2Parallel) { + t.Fatalf("chi2 differs: serial %v, parallel %v", chi2Serial, chi2Parallel) + } + serial, parallel := pSerial.RawFloats(), pParallel.RawFloats() + for i := range serial { + if math.Float64bits(serial[i]) != math.Float64bits(parallel[i]) { + t.Fatalf("parameter %d differs: serial %.17g, parallel %.17g", i, serial[i], parallel[i]) + } + } +} + +func TestParallelJacobianBitIdenticalRoot(t *testing.T) { + f, start := rootParallelProblem() + run := func(parallel, broyden bool) (*core.Array, float64) { + t.Helper() + x, res, err := FindRootSystem(f, mustFloats(t, start), RootSystemOptions{ + MaxIterations: 100, + UseBroyden: broyden, + ParallelJacobian: parallel, + }) + if err != nil { + t.Fatalf("FindRootSystem(UseBroyden=%v, ParallelJacobian=%v): %v", broyden, parallel, err) + } + return x, res + } + for _, broyden := range []bool{false, true} { + xSerial, resSerial := run(false, broyden) + xParallel, resParallel := run(true, broyden) + if math.Float64bits(resSerial) != math.Float64bits(resParallel) { + t.Fatalf("UseBroyden=%v: residual norm differs: serial %v, parallel %v", broyden, resSerial, resParallel) + } + sSerial, sParallel := xSerial.RawFloats(), xParallel.RawFloats() + for i := range sSerial { + if math.Float64bits(sSerial[i]) != math.Float64bits(sParallel[i]) { + t.Fatalf("UseBroyden=%v: root coordinate %d differs: serial %.17g, parallel %.17g", + broyden, i, sSerial[i], sParallel[i]) + } + } + } +} + +// TestParallelJacobianBitIdenticalBroydenRebuilds drives the rebuild +// path: the duplicated equation x² = 1 has a rank-one Jacobian, so a +// Broyden run from the hopeless start rebuilds the numerical Jacobian +// every round until the budget refuses it. Serial and parallel must +// refuse with the same words, which pins every rebuilt Jacobian's bits. +func TestParallelJacobianBitIdenticalBroydenRebuilds(t *testing.T) { + rankOne := func(x *core.Array) (*core.Array, error) { + v := x.FloatAt(0)*x.FloatAt(0) - 1 + out := core.New(core.Float, 2) + vals := out.RawFloats() + vals[0], vals[1] = v, v + return out, nil + } + run := func(parallel bool) string { + t.Helper() + _, _, err := FindRootSystem(rankOne, mustFloats(t, []float64{1.5, 1.5}), RootSystemOptions{ + MaxIterations: 12, + UseBroyden: true, + ParallelJacobian: parallel, + }) + if err == nil { + t.Fatalf("FindRootSystem(ParallelJacobian=%v): want the budget refusal", parallel) + } + return err.Error() + } + if serial, parallel := run(false), run(true); serial != parallel { + t.Fatalf("the refusal differs: serial %q, parallel %q", serial, parallel) + } +} + +// TestParallelJacobianCallbackConcurrency pins the consent boundary: +// with the option off the residual callback never runs inside more +// than one goroutine at once, and with it on the sweeps do overlap. +func TestParallelJacobianCallbackConcurrency(t *testing.T) { + prev := engine.SetNumWorkers(4) + defer engine.SetNumWorkers(prev) + + const nP = 8 + const nObs = 8 + peak := func(t *testing.T, parallel bool) int64 { + t.Helper() + var tr concurrencyTracker + residual := func(p *core.Array) (*core.Array, error) { + tr.enter() + defer tr.leave() + if parallel { + // The caller holds its ground until a second one is + // inside, so the overlap the option promises is + // demonstrated by construction instead of left to the + // scheduler to interleave two calls this short. A + // sweep that collapsed to one goroutine would wait + // out the grace alone and fail the peak below. + deadline := time.Now().Add(250 * time.Millisecond) + for tr.peak.Load() < 2 && time.Now().Before(deadline) { + runtime.Gosched() + } + } + s := 0.0 + for k := range 800 { + s += math.Sin(float64(k)*7e-4+p.FloatAt(k&(nP-1))) * math.Cos(float64(k)*3e-4) + } + out := core.New(core.Float, nObs) + vals := out.RawFloats() + for i := range nObs { + vals[i] = s + p.FloatAt(i&(nP-1)) + } + return out, nil + } + start := make([]float64, nP) + if _, _, err := LevenbergMarquardt(residual, mustFloats(t, start), LMOptions{ + MaxIterations: 2, + AllowBudgetExit: true, + ParallelJacobian: parallel, + }); err != nil { + t.Fatalf("LevenbergMarquardt(ParallelJacobian=%v): %v", parallel, err) + } + return tr.peak.Load() + } + t.Run("serial", func(t *testing.T) { + if got := peak(t, false); got != 1 { + t.Fatalf("the serial run reached %d concurrent callback calls, want exactly 1", got) + } + }) + t.Run("parallel", func(t *testing.T) { + if got := peak(t, true); got < 2 { + t.Fatalf("the parallel run reached %d concurrent callback calls, want more than 1", got) + } + }) +} + +// TestParallelJacobianErrorMatchesSerial fails one column's stencil +// and requires the parallel sweep to report the same error the serial +// walk reports: the lowest failing column, not whichever worker got +// there first. +func TestParallelJacobianErrorMatchesSerial(t *testing.T) { + const nP = 6 + const nObs = 10 + // Coordinate 3 is nonzero exactly while column 3's stencils run, + // so only that column's residual evaluations fail. + residual := func(p *core.Array) (*core.Array, error) { + if p.FloatAt(3) != 0 { + return nil, fmt.Errorf("the stencil touched column 3") + } + out := core.New(core.Float, nObs) + for i := range nObs { + out.RawFloats()[i] = p.FloatAt(0) + float64(i) + } + return out, nil + } + run := func(parallel bool) string { + t.Helper() + start := make([]float64, nP) + _, _, err := LevenbergMarquardt(residual, mustFloats(t, start), LMOptions{ + MaxIterations: 10, + ParallelJacobian: parallel, + }) + if err == nil { + t.Fatalf("LevenbergMarquardt(ParallelJacobian=%v): want the stencil error", parallel) + } + return err.Error() + } + serial := run(false) + parallel := run(true) + if serial != parallel { + t.Fatalf("the error differs: serial %q, parallel %q", serial, parallel) + } +} + +// benchJacResidual builds a sixteen-parameter least squares problem +// whose residual costs roughly twenty microseconds of real arithmetic +// per call, and whose data sit far from every model the fit can reach, +// so a run spends a fixed number of rounds, almost all of it inside +// the Jacobian sweep, and never converges. +func benchJacResidual(b *testing.B) (residual func(*core.Array) (*core.Array, error), start *core.Array) { + b.Helper() + const nP = 16 + const nObs = 64 + const spin = 512 + t := make([]float64, nObs) + y := make([]float64, nObs) + for i := range nObs { + t[i] = float64(i) / 8 + y[i] = 1e3 + float64(i) + } + residual = func(p *core.Array) (*core.Array, error) { + s := 0.0 + for k := range spin { + s += math.Sin(float64(k)*7e-4+p.FloatAt(k&(nP-1))) * math.Cos(float64(k)*3e-4) + } + out := core.New(core.Float, nObs) + vals := out.RawFloats() + for i := range nObs { + v := s + for k := range nP / 2 { + v += p.FloatAt(2*k) * math.Exp(-p.FloatAt(2*k+1)*t[i]) + } + vals[i] = v - y[i] + } + return out, nil + } + p0 := make([]float64, nP) + for i := range p0 { + p0[i] = 0.5 + } + arr, err := core.FromFloats(p0, nP) + if err != nil { + b.Fatal(err) + } + return residual, arr +} + +// BenchmarkLevenbergMarquardtJacobianSweep times the central-difference +// Jacobian sweep serial against parallel in one binary, the same fit +// walked both ways. The worker count mirrors the GOMAXPROCS=8 the run +// is read under: eight workers, one per thread. +func BenchmarkLevenbergMarquardtJacobianSweep(b *testing.B) { + restore := engine.SetNumWorkers(8) + defer engine.SetNumWorkers(restore) + residual, start := benchJacResidual(b) + for _, parallel := range []bool{false, true} { + name := "serial" + if parallel { + name = "parallel" + } + b.Run(name, func(b *testing.B) { + opts := LMOptions{MaxIterations: 12, AllowBudgetExit: true, ParallelJacobian: parallel} + b.ReportAllocs() + for b.Loop() { + if _, _, err := LevenbergMarquardt(residual, start, opts); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/optim/lbfgs.go b/optim/lbfgs.go new file mode 100644 index 0000000..dcd4086 --- /dev/null +++ b/optim/lbfgs.go @@ -0,0 +1,551 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/linalg" +) + +import "math" + +// Limited-memory BFGS optimisation. The plain `Minimise` (Nelder-Mead) +// is derivative-free and robust but needs O(n²) memory for the simplex +// and many objective evaluations per step. L-BFGS is the standard +// quasi-Newton method for smooth objectives on moderate-to-large +// problems: it approximates the inverse Hessian from the last `memory` +// position/gradient pairs by the two-loop recursion, which costs +// O(memory·n) per step, far cheaper than the O(n²) simplex or the +// O(n³) dense Hessian, and far faster to converge than gradient +// descent on ill-conditioned objectives. +// +// The gradient may be supplied by the caller or computed by central +// finite differences, which costs one extra objective evaluation per +// coordinate per step. + +// LBFGSOptions tunes the optimiser. MaxIterations ≤ 0 means 10000, +// Tolerance ≤ 0 means 1e-8 (L∞ norm of the projected gradient, see +// below), Memory ≤ 0 means 10. +// +// The tolerance is absolute in the gradient's own scale: an objective +// whose gradient sits many orders of magnitude below one is reported +// converged at its start point. Rescale the objective to O(1) before +// calling MinimiseLBFGS when the natural units are not of that size. +// +// Lower and Upper, when non-nil, bound the search coordinate-wise: +// an infinite entry leaves that side open and a crossed or NaN pair of +// walls is an error. The iterate is then kept feasible by projecting +// the line search onto the box, and a coordinate sitting on a wall the +// gradient pushes against is pinned: the two-loop recursion and the +// curvature pairs run in the free subspace, and the convergence +// measure is the projected gradient, which is the KKT residual of the +// box problem. A gradient by finite differences turns one-sided at a +// wall, so the objective is never evaluated outside the box. x0 is +// projected onto the box rather than refused. +type LBFGSOptions struct { + MaxIterations int + Tolerance float64 + Memory int + Lower []float64 + Upper []float64 + // AllowBudgetExit makes a run that exhausts MaxIterations report + // the best point it reached with a nil error instead of refusing + // it. MinimiseConstrained sets it for its inner solves, whose + // accuracy the outer loop's feasibility check judges; a direct + // caller leaves it false, so a budget stop is never mistaken for a + // converged answer. A run that converges normally is unaffected. + AllowBudgetExit bool +} + +// MinimiseLBFGS returns the point and value of a local minimum of f +// near x0 by the limited-memory BFGS method with backtracking Armijo +// line search. The gradient grad may be nil, in which case central +// finite differences are used (costing two extra objective evaluations +// per coordinate per step). +func MinimiseLBFGS(f func(*core.Array) (float64, error), grad func(*core.Array) (*core.Array, error), + x0 *core.Array, opts LBFGSOptions) (*core.Array, float64, error) { + n := x0.Len() + if n == 0 { + return nil, 0, base.Errf("MinimiseLBFGS: the starting point must have at least one element") + } + if x0.Dtype() == core.Complex { + return nil, 0, base.Errf("MinimiseLBFGS: complex starting points are not supported") + } + if opts.MaxIterations <= 0 { + opts.MaxIterations = 10000 + } + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-8 + } + if opts.Memory <= 0 { + opts.Memory = 10 + } + mem := min(opts.Memory, n) + + lower, upper, err := boundsOf(opts, n) + if err != nil { + return nil, 0, base.Errf("MinimiseLBFGS: %w", err) + } + + x := make([]float64, n) + for i := range n { + x[i] = min(max(x0.FloatAt(i), lower[i]), upper[i]) + } + fx, err := f(linalg.ArrayFromFloatsSafe(x, n)) + if err != nil { + return nil, 0, base.Errf("MinimiseLBFGS: %w", err) + } + if math.IsNaN(fx) || math.IsInf(fx, 0) { + return nil, 0, base.Errf("MinimiseLBFGS: the objective returned the non-finite value %g at the starting point", fx) + } + g := make([]float64, n) + xp := make([]float64, n) + xm := make([]float64, n) + if err := evalGrad(f, grad, x, g, lower, upper, xp, xm); err != nil { + return nil, 0, base.Errf("MinimiseLBFGS: %w", err) + } + pinnedMask := make([]bool, n) + projected := make([]float64, n) + project := func() float64 { + // The projected gradient zeroes a coordinate that sits on a + // wall the gradient pushes against: its norm is the KKT + // residual of the box problem, its support the free subspace + // the two-loop recursion runs in. The true gradient stays + // untouched in g. + norm := 0.0 + for i := range n { + pinnedMask[i] = pinnedAt(lower, upper, x, g, i) + if pinnedMask[i] { + projected[i] = 0 + } else { + projected[i] = g[i] + } + if v := math.Abs(projected[i]); v > norm { + norm = v + } + } + return norm + } + gNorm := project() + if gNorm <= opts.Tolerance { + out, fv := packResult(x, fx) + return out, fv, nil + } + + // Limited-memory correction pairs: (s_i, y_i) with ρ_i = 1/(y_iᵀs_i). + // The window is a ring of pairs, each pair's vectors allocated with + // the first pair that needs them: count is how many of + // its slots hold a pair, head is the oldest of them, and the pair at + // position d of the window sits at (head+d) % mem, position 0 being + // the oldest. A step that arrives once the window is full overwrites + // the slot it drops, so an accepted step allocates nothing. + history := make([]lmPair, mem) + head, count := 0, 0 + + // Scratch reused across iterations: the two-loop direction, its + // per-pair store, the accepted-step buffers and the finite-difference + // step and its two stencils. Each is fully overwritten before it is + // read, and nothing escapes the iteration by reference. + q := make([]float64, n) + store := make([]float64, mem) + dir := q + trial := make([]float64, n) + xNew := make([]float64, n) + gNew := make([]float64, n) + sVec := make([]float64, n) + yVec := make([]float64, n) + + converged := false + for iter := 0; iter < opts.MaxIterations; iter++ { + gNorm = project() + if gNorm <= opts.Tolerance { + converged = true + break + } + // Two-loop recursion for the search direction, driven by the + // projected gradient and confined to the free subspace: the + // dots skip pinned coordinates and q is re-zeroed on them + // after every update, because history pairs predate the + // current active set and their stale components would drag a + // pinned coordinate off its wall or corrupt the slope test. + twoLoopRecursion(projected, history, head, count, mem, pinnedMask, q, store) + // The search direction is −q (the negative approximate + // inverse Hessian times the gradient); q is dead past this + // point, so the negation happens in place. + for j := range n { + dir[j] = -q[j] + } + dirNorm := maxAbs(dir) + if dirNorm == 0 { + return nil, 0, base.Errf("MinimiseLBFGS: the search direction vanished with a projected gradient of %g", gNorm) + } + + // Backtracking Armijo line search along the free subspace, + // projected onto the box so every trial point stays feasible. + // The trial step starts at one unit, the standard choice, and + // halves until the objective's sufficient decrease is met: the + // step that actually reduces a quadratic is ≈ 1/L for a + // curvature L, so a stiff objective (penalty weights, mixed + // units) needs a long halving chain before any trial is + // acceptable. A search that cannot find an acceptable point + // within the budget is a failure, never a silent success: the + // caller would otherwise read the start point as the answer. + step := 1.0 + slope := 0.0 + for j := range n { + slope += g[j] * dir[j] + } + if slope >= 0 { + return nil, 0, base.Errf("MinimiseLBFGS: the search direction is not a descent direction (slope %g) at a projected gradient of %g", + slope, gNorm) + } + const ( + armijo = 1e-4 + backtrack = 60 + ) + accepted := false + var ft float64 + var ferr error + for range backtrack { + for j := range n { + trial[j] = min(max(x[j]+step*dir[j], lower[j]), upper[j]) + } + ft, ferr = f(linalg.ArrayFromFloatsSafe(trial, n)) + if ferr != nil { + step *= 0.5 + continue + } + if ft <= fx+armijo*step*slope { + accepted = true + break + } + step *= 0.5 + } + if !accepted { + if ferr != nil { + // The last trial failed by error, not by value: the + // caller sees the cause instead of a bare stall. + return nil, 0, base.Errf("MinimiseLBFGS: the line search stalled at a step of %g with a projected gradient of %g: %w", + step, gNorm, ferr) + } + return nil, 0, base.Errf("MinimiseLBFGS: the line search stalled at a step of %g with a projected gradient of %g", + step, gNorm) + } + if math.IsNaN(ft) || math.IsInf(ft, 0) { + return nil, 0, base.Errf("MinimiseLBFGS: the objective returned the non-finite value %g during the line search", ft) + } + // The accepted point is the projected trial, not the unprojected + // step: the displacement s must be what actually moved. Its + // value is the line search's own: the trial already sat at + // xNew, so a second evaluation would ask the objective for the + // number it just returned. + copy(xNew, trial) + fxNew := ft + if err := evalGrad(f, grad, xNew, gNew, lower, upper, xp, xm); err != nil { + return nil, 0, base.Errf("MinimiseLBFGS: %w", err) + } + ys := 0.0 + for j := range n { + sVec[j] = xNew[j] - x[j] + yVec[j] = gNew[j] - g[j] + // A coordinate the projection held still carries no + // curvature information; its gradient change is noise the + // second loop would otherwise mix into the direction. + if sVec[j] == 0 { + yVec[j] = 0 + } + ys += yVec[j] * sVec[j] + } + if ys > 1e-16 { + if count < mem { + p := &history[(head+count)%mem] + if p.s == nil { + // The vectors arrive with the first pair that needs + // them, so a solve that takes few steps does not pay + // for the whole window up front. + p.s = make([]float64, n) + p.y = make([]float64, n) + } + copy(p.s, sVec) + copy(p.y, yVec) + p.rho = 1 / ys + count++ + } else { + p := &history[head] + copy(p.s, sVec) + copy(p.y, yVec) + p.rho = 1 / ys + head = (head + 1) % mem + } + } + copy(x, xNew) + copy(g, gNew) + fx = fxNew + } + // The last accepted step updated x, g and fx 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. + gNorm = project() + if gNorm <= opts.Tolerance { + converged = true + } + // Falling out of the loop means the budget ran out, not that a + // minimum was found: the projected gradient is still above the + // tolerance, and reporting the point as a converged answer is the + // silent-wrongness the other exits refuse. The best point is not + // returned, exactly as the stall and direction exits do not return + // theirs. + if !converged { + if !opts.AllowBudgetExit { + return nil, 0, base.Errf("MinimiseLBFGS: the iteration budget of %d ran out with a projected gradient of %g, above the tolerance %g", + opts.MaxIterations, gNorm, opts.Tolerance) + } + out, fv := packResult(x, fx) + return out, fv, nil + } + out, fv := packResult(x, fx) + return out, fv, nil +} + +// lmPair is one limited-memory correction pair: the position +// difference s, the gradient difference y and ρ = 1/(yᵀs). +type lmPair struct { + s, y []float64 + rho float64 +} + +// twoLoopRecursion drives the search direction through the stored +// pairs, driven by the projected gradient q starts from. When no +// coordinate is pinned the four sweeps run mask-free, which is the +// common case for an interior iterate; with pins the masked walk +// skips them in the dots and re-zeroes q on them after every update. +// Both walks produce the same bits for the same state: the mask-free +// path is the masked one with an always-false branch removed. +func twoLoopRecursion(projected []float64, history []lmPair, head, count, mem int, pinned []bool, q, store []float64) { + n := len(projected) + copy(q, projected) + store = store[:count] + free := 0 + for j := range n { + if !pinned[j] { + free++ + } + } + if free == n { + for i := count - 1; i >= 0; i-- { + p := &history[(head+i)%mem] + dot := 0.0 + for j := range n { + dot += p.s[j] * q[j] + } + store[i] = dot * p.rho + for j := range n { + q[j] -= store[i] * p.y[j] + } + } + } else { + for i := count - 1; i >= 0; i-- { + p := &history[(head+i)%mem] + dot := 0.0 + for j := range n { + if !pinned[j] { + dot += p.s[j] * q[j] + } + } + store[i] = dot * p.rho + for j := range n { + q[j] -= store[i] * p.y[j] + if pinned[j] { + q[j] = 0 + } + } + } + } + gamma := 1.0 + if count > 0 { + last := &history[(head+count-1)%mem] + num, den := 0.0, 0.0 + for j := range n { + num += last.s[j] * last.y[j] + den += last.y[j] * last.y[j] + } + if den != 0 { + gamma = num / den + } + } + for j := range n { + q[j] *= gamma + } + if free == n { + for i := range count { + p := &history[(head+i)%mem] + dot := 0.0 + for j := range n { + dot += p.y[j] * q[j] + } + coeff := store[i] - p.rho*dot + for j := range n { + q[j] += coeff * p.s[j] + } + } + return + } + for i := range count { + p := &history[(head+i)%mem] + dot := 0.0 + for j := range n { + if !pinned[j] { + dot += p.y[j] * q[j] + } + } + coeff := store[i] - p.rho*dot + for j := range n { + q[j] += coeff * p.s[j] + if pinned[j] { + q[j] = 0 + } + } + } +} + +// boundsOf resolves the optional box walls: a nil slice opens every +// side, an infinite entry opens one side, and NaN walls or a crossed +// pair are errors. +func boundsOf(opts LBFGSOptions, n int) (lower, upper []float64, err error) { + lower = make([]float64, n) + upper = make([]float64, n) + for i := range n { + lower[i], upper[i] = math.Inf(-1), math.Inf(1) + } + if opts.Lower != nil { + if len(opts.Lower) != n { + return nil, nil, base.Errf("the lower bounds hold %d entries for %d variables", len(opts.Lower), n) + } + copy(lower, opts.Lower) + } + if opts.Upper != nil { + if len(opts.Upper) != n { + return nil, nil, base.Errf("the upper bounds hold %d entries for %d variables", len(opts.Upper), n) + } + copy(upper, opts.Upper) + } + for i := range n { + if math.IsNaN(lower[i]) || math.IsNaN(upper[i]) || lower[i] > upper[i] { + return nil, nil, base.Errf("variable %d has bounds [%g, %g]", i, lower[i], upper[i]) + } + } + return lower, upper, nil +} + +// pinnedAt reports whether coordinate i sits on a wall the gradient +// pushes against: at a lower wall with a positive gradient the descent +// direction -g points out of the box, so the KKT condition holds and +// the coordinate is pinned. The wall test carries a relative tolerance +// wide enough to cover the drift of accepted steps that landed beside +// the wall without crossing it; release needs only the gradient to +// turn favourable, so a generous tolerance cannot trap a coordinate. +func pinnedAt(lower, upper, x, g []float64, i int) bool { + const wall = 1e-7 + // An equality-bounded coordinate can never move, whatever the + // gradient says: pinning it keeps the two-loop direction and the + // projected-gradient norm consistent with the frozen wall. + if !math.IsInf(lower[i], 0) && lower[i] == upper[i] { + return true + } + if g[i] > 0 && !math.IsInf(lower[i], 0) && x[i]-lower[i] <= wall*math.Max(1, math.Abs(lower[i])) { + return true + } + if g[i] < 0 && !math.IsInf(upper[i], 0) && upper[i]-x[i] <= wall*math.Max(1, math.Abs(upper[i])) { + return true + } + return false +} + +// evalGrad fills g with the gradient of f at x: the caller-supplied +// gradient if available, else finite differences with step +// √ε·max(1,|xᵢ|). A coordinate whose central stencil would leave the +// box falls back to the one-sided stencil that stays inside it, so a +// bounded objective is never evaluated outside its domain. xp and xm +// are scratch of length len(x): the perturbed copy carries the offset +// on one coordinate only, which is restored once the coordinate is +// done, so neither a copy of the whole point nor a fresh slice per +// coordinate is needed. linalg.ArrayFromFloatsSafe copies into the +// array handed to f, so the objective never observes later mutation. +func evalGrad(f func(*core.Array) (float64, error), grad func(*core.Array) (*core.Array, error), + x, g, lower, upper, xp, xm []float64) error { + if grad != nil { + ga, err := grad(linalg.ArrayFromFloatsSafe(x, len(x))) + if err != nil { + return err + } + if err := requireReal("MinimiseLBFGS", "gradients", ga); err != nil { + return err + } + if ga.Len() != len(x) { + return base.Errf("MinimiseLBFGS: the gradient callback returned %d elements for %d variables", + ga.Len(), len(x)) + } + for i := range x { + g[i] = ga.FloatAt(i) + // A non-finite gradient coordinate must be an error, not a + // NaN that skips the projected-gradient max and reads as + // converged. + if math.IsNaN(g[i]) || math.IsInf(g[i], 0) { + return base.Errf("MinimiseLBFGS: the gradient callback returned the non-finite value %g at coordinate %d", g[i], i) + } + } + return nil + } + copy(xp, x) + copy(xm, x) + for i := range x { + eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(x[i])) + xp[i] = min(x[i]+eps, upper[i]) + xm[i] = max(x[i]-eps, lower[i]) + if xp[i] == xm[i] { + // An equality-bounded coordinate clamps both stencil + // points to the same wall; the central difference would + // divide 0/0 into a NaN gradient. + g[i] = 0 + xp[i], xm[i] = x[i], x[i] + continue + } + fp, err := f(linalg.ArrayFromFloatsSafe(xp, len(x))) + if err != nil { + return err + } + fm, err := f(linalg.ArrayFromFloatsSafe(xm, len(x))) + if err != nil { + return err + } + g[i] = (fp - fm) / (xp[i] - xm[i]) + xp[i], xm[i] = x[i], x[i] + // The finite-difference path validates what the analytic one + // does: a NaN the stencil produces is invisible to the + // projected-gradient max and would read as converged. + if math.IsNaN(g[i]) || math.IsInf(g[i], 0) { + return base.Errf("MinimiseLBFGS: the objective returned the non-finite values %g and %g around coordinate %d", fp, fm, i) + } + } + return nil +} + +// maxAbs returns the maximum absolute value of a slice. +func maxAbs(v []float64) float64 { + m := 0.0 + for _, a := range v { + if x := math.Abs(a); x > m { + m = x + } + } + return m +} + +// packResult wraps a flat vector as a rank-1 core.Array. +func packResult(x []float64, fx float64) (*core.Array, float64) { + out := core.New(core.Float, len(x)) + copy(out.RawFloats(), x) + return out, fx +} diff --git a/optim/lbfgs_test.go b/optim/lbfgs_test.go new file mode 100644 index 0000000..3249617 --- /dev/null +++ b/optim/lbfgs_test.go @@ -0,0 +1,109 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestLBFGSRosenbrock checks convergence on the Rosenbrock valley, +// the standard test for quasi-Newton methods whose narrow curved +// valley defeats naive gradient descent. +func TestLBFGSRosenbrock(t *testing.T) { + rosenbrock := func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + } + gradFn := func(p *core.Array) (*core.Array, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + out := core.New(core.Float, 2) + out.RawFloats()[0] = -2*(1-x) - 400*x*(y-x*x) + out.RawFloats()[1] = 200 * (y - x*x) + return out, nil + } + start, _ := core.FromFloats([]float64{-1.2, 1}, 2) + point, value, err := MinimiseLBFGS(rosenbrock, gradFn, start, LBFGSOptions{}) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + if math.Abs(point.FloatAt(0)-1) > 1e-5 || math.Abs(point.FloatAt(1)-1) > 1e-5 { + t.Fatalf("minimiser = (%.10g, %.10g), want (1, 1)", point.FloatAt(0), point.FloatAt(1)) + } + if value > 1e-10 { + t.Fatalf("minimum value = %v, want ≈ 0", value) + } +} + +// TestLBFGSNumericGradient checks that finite-difference gradients +// converge on the same problem without a caller-supplied derivative. +func TestLBFGSNumericGradient(t *testing.T) { + f := func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + } + start, _ := core.FromFloats([]float64{-1.2, 1}, 2) + point, _, err := MinimiseLBFGS(f, nil, start, LBFGSOptions{Tolerance: 1e-6}) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + if math.Abs(point.FloatAt(0)-1) > 1e-3 || math.Abs(point.FloatAt(1)-1) > 1e-3 { + t.Fatalf("minimiser = (%.10g, %.10g), want (1, 1)", point.FloatAt(0), point.FloatAt(1)) + } +} + +// TestLBFGSConvergenceSpeed checks that L-BFGS converges in +// substantially fewer iterations than gradient descent on the same +// ill-conditioned problem. +func TestLBFGSConvergenceSpeed(t *testing.T) { + const n = 10 + // Diagonal quadratic with eigenvalues 1..n. + f := func(p *core.Array) (float64, error) { + s := 0.0 + for i := range p.Len() { + d := p.FloatAt(i) - float64(i+1) + s += float64(i+1) * d * d + } + return s, nil + } + gradFn := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, n) + for i := range n { + out.RawFloats()[i] = 2 * float64(i+1) * (p.FloatAt(i) - float64(i+1)) + } + return out, nil + } + start, _ := core.FromFloats(make([]float64, n), n) + point, _, err := MinimiseLBFGS(f, gradFn, start, LBFGSOptions{}) + if err != nil { + t.Fatalf("MinimiseLBFGS: %v", err) + } + for i := range n { + want := float64(i + 1) + if math.Abs(point.FloatAt(i)-want) > 1e-6 { + t.Fatalf("x[%d] = %.10g, want %.10g", i, point.FloatAt(i), want) + } + } +} + +// TestLBFGSInvalid pins the error contract. +func TestLBFGSInvalid(t *testing.T) { + f := func(p *core.Array) (float64, error) { return 0, nil } + empty, _ := core.FromFloats(nil, 0) + if _, _, err := MinimiseLBFGS(f, nil, empty, LBFGSOptions{}); err == nil { + t.Fatal("expected an error for an empty starting point") + } + cx, _ := core.FromComplexes([]complex128{1}, 1) + if _, _, err := MinimiseLBFGS(f, nil, cx, LBFGSOptions{}); err == nil { + t.Fatal("expected an error for a complex starting point") + } + // Objective that errors must propagate. + failing := func(p *core.Array) (float64, error) { return 0, base.Errf("objective exploded") } + start, _ := core.FromFloats([]float64{1}, 1) + if _, _, err := MinimiseLBFGS(failing, nil, start, LBFGSOptions{}); err == nil { + t.Fatal("expected the objective's error to propagate") + } +} diff --git a/optim/leastsq.go b/optim/leastsq.go new file mode 100644 index 0000000..47d94e3 --- /dev/null +++ b/optim/leastsq.go @@ -0,0 +1,694 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" + "sourcedock.dev/petrbalvin/tensor/linalg" +) + +// FitStatus states how an iterative fit ended. +type FitStatus int + +const ( + // FitConverged marks a run that met one of the tolerances: the + // relative chi2 improvement, the gradient norm GradTol or the + // step size StepTol. + FitConverged FitStatus = iota + // FitStalled marks a run whose step died: the normal equations + // turned singular or the damping collapsed without the residual + // meeting the tolerance. The returned point is the best one the + // run reached, never a partial step past it. + FitStalled + // FitBudget marks a run that spent its iteration budget before + // any tolerance or stall fired. The returned point is the last + // iterate. + FitBudget +) + +// FitResult carries everything a fit reports: the point it ended on, +// the (weighted) residual sum of squares there, how the run ended +// and, when requested, the parameter covariance at that point. +type FitResult struct { + Parameters *core.Array + Chi2 float64 + Status FitStatus + Covariance *core.Array +} + +// LMOptions tunes LevenbergMarquardt and LevenbergMarquardtFit. +// Lambda ≤ 0 means 1e-3 (the initial damping factor), Tolerance ≤ 0 +// means 1e-10 (the relative χ² improvement threshold) and +// MaxIterations ≤ 0 means 200. +type LMOptions struct { + MaxIterations int + Tolerance float64 + Lambda float64 + Jacobian func(p *core.Array) (*core.Array, error) + // AllowBudgetExit makes a run that exhausts MaxIterations report + // its last point instead of an error. The default is false, so a + // budget stop is never mistaken for a converged answer; the flag + // mirrors LBFGSOptions.AllowBudgetExit. LevenbergMarquardtFit + // needs no flag: it reports the budget stop as FitBudget. + AllowBudgetExit bool + // ParallelJacobian lets the central-difference Jacobian sweep its + // columns on several goroutines. Setting it is the caller's + // consent that the residual callback may run concurrently from + // more than one goroutine: the default false keeps every + // evaluation on the caller's goroutine, and the fit is bit for bit + // the same either way, because the columns are independent and + // each one is differenced by the same stencil. The field does + // nothing while Jacobian supplies the analytic matrix. + ParallelJacobian bool + // GradTol converges the fit once the infinity norm of the + // gradient Jᵀr falls to it, the test that catches the flat + // optimum where chi2 still falls in slivers while the step + // directions carry no information. A value ≤ 0 disables the test, + // which is the default: a caller who sets it picks the scale. + GradTol float64 + // StepTol converges the fit once an accepted step's infinity norm + // falls to StepTol·(‖p‖∞ + StepTol), the relative step test that + // stops a fit whose parameters have stopped moving meaningfully. + // A value ≤ 0 disables the test, which is the default. + StepTol float64 + // Sigma weights the residuals by the measurement covariance: a + // vector holds one positive variance per residual, a square + // matrix holds the full nR×nR covariance and must be exactly + // symmetric and positive definite. The fit whitens the residuals + // and the Jacobian through the factor once, chi2 becomes rᵀC⁻¹r + // and a requested covariance becomes (JᵀC⁻¹J)⁻¹. Nil, the + // default, leaves every residual unweighted. + Sigma *core.Array + // RequestCovariance fills FitResult.Covariance with (JᵀJ)⁻¹ at + // the returned point, or (JᵀC⁻¹J)⁻¹ under Sigma. The answer costs + // one more Jacobian at the final point, which on the + // difference route is two residual evaluations per parameter. A + // Jacobian that is rank-deficient at the returned point has no + // covariance to report and the run fails naming that, so a + // caller asking for a covariance accepts the trade on a fit it + // expects to stall. + RequestCovariance bool +} + +// LevenbergMarquardt minimises ‖r(p)‖² by LM damping of the +// Gauss-Newton step, with a central-difference or user-supplied +// analytic Jacobian and the library's LU solver for the normal +// equations. +// +// The historical error contract stays: a run that stalls (a singular +// solve, a collapsed damping) or spends its budget without +// AllowBudgetExit is an error, not a point. LevenbergMarquardtFit +// reports those conditions as a FitResult instead and carries the +// gradient, step, weight and covariance options. +func LevenbergMarquardt(residual func(*core.Array) (*core.Array, error), p0 *core.Array, opts LMOptions) (*core.Array, float64, error) { + res, err := runLevenbergMarquardt(residual, p0, opts, true) + if err != nil { + return nil, 0, err + } + return res.Parameters, res.Chi2, nil +} + +// LevenbergMarquardtFit is LevenbergMarquardt with the full report: +// the status that says how the run ended and, on request, the +// parameter covariance. Where LevenbergMarquardt keeps its historical +// error contract, this one reports every ended run as a result: a +// stalled step and a spent budget come back as FitStalled and +// FitBudget on the best point reached, never as an error. The errors +// here are the model's own fault (a residual that fails or turns +// non-finite at a state the fit adopts, a malformed Sigma) and, with +// RequestCovariance, a rank-deficient Jacobian at the answer. +func LevenbergMarquardtFit(residual func(*core.Array) (*core.Array, error), p0 *core.Array, opts LMOptions) (*FitResult, error) { + return runLevenbergMarquardt(residual, p0, opts, false) +} + +// denseFloats returns the array's float64 payload when a is a dense +// float64 array and nil otherwise: hot loops branch once on the result +// and sweep the payload directly, falling back to the widening +// accessor for views and other dtypes. The elements are identical +// either way, so a dense sweep computes the same bits as the accessor +// walk it replaces. +func denseFloats(a *core.Array) []float64 { + if !a.Strided() && a.Dtype() == core.Float { + return a.RawFloats() + } + return nil +} + +// whitener carries the Sigma factor: the per-residual divisors of a +// variance vector, or the lower Cholesky factor of a full covariance. +// A nil whitener is the unweighted fit. +type whitener struct { + diag []float64 + factor [][]float64 +} + +// sigmaWhitener validates Sigma against the residual count nR and +// factors it. The matrix form is checked for exact symmetry before +// the factorisation reads one triangle: an asymmetric partner would +// silently weight by a matrix the caller did not pass. +func sigmaWhitener(sigma *core.Array, nR int) (*whitener, error) { + if sigma == nil { + return nil, nil + } + if err := requireReal("LevenbergMarquardt", "Sigma", sigma); err != nil { + return nil, err + } + if sigma.NDim() == 1 { + if sigma.Len() != nR { + return nil, base.Errf("LevenbergMarquardt: Sigma must hold one variance per residual (%d), got %d", nR, sigma.Len()) + } + w := &whitener{diag: make([]float64, nR)} + for i := range nR { + v := sigma.FloatAt(i) + if math.IsNaN(v) || v <= 0 { + return nil, base.Errf("LevenbergMarquardt: Sigma must hold positive variances, got %g at %d", v, i) + } + w.diag[i] = math.Sqrt(v) + } + return w, nil + } + if sigma.NDim() != 2 || sigma.Shape()[0] != nR || sigma.Shape()[1] != nR { + return nil, base.Errf("LevenbergMarquardt: Sigma must be a %d×%d covariance or a vector of %d variances, got shape %s", + nR, nR, nR, base.ShapeText(sigma.Shape())) + } + for i := range nR { + for j := i + 1; j < nR; j++ { + up, lo := sigma.FloatAt(i*nR+j), sigma.FloatAt(j*nR+i) + if up != lo { + return nil, base.Errf("LevenbergMarquardt: Sigma must be symmetric, got %g and %g at (%d, %d)", up, lo, i, j) + } + } + } + l, err := linalg.Cholesky(sigma) + if err != nil { + return nil, base.Errf("LevenbergMarquardt: Sigma is not positive definite: %w", err) + } + w := &whitener{factor: make([][]float64, nR)} + for i := range nR { + row := make([]float64, i+1) + for j := range i + 1 { + row[j] = l.FloatAt(i*nR + j) + } + w.factor[i] = row + } + return w, nil +} + +// vector whitens a residual in place: y solves L y = r. +func (w *whitener) vector(r []float64) { + if w == nil { + return + } + if w.diag != nil { + for i := range r { + r[i] /= w.diag[i] + } + return + } + for i := range r { + s := r[i] + li := w.factor[i] + for j := range i { + s -= li[j] * r[j] + } + r[i] = s / li[i] + } +} + +// matrix whitens a Jacobian in place: the rows solve L J' = J, so the +// downstream normal equations accumulate JᵀC⁻¹J without knowing a +// weight exists. +func (w *whitener) matrix(jac [][]float64) { + if w == nil { + return + } + for i := range jac { + if w.diag != nil { + for j := range jac[i] { + jac[i][j] /= w.diag[i] + } + continue + } + li := w.factor[i] + for j := range jac[i] { + s := jac[i][j] + for k := range i { + s -= li[k] * jac[k][j] + } + jac[i][j] = s / li[i] + } + } +} + +// covarianceFromJac inverts the unweighted normal equations of the +// (whitened) Jacobian, which is the parameter covariance. The solve +// runs column by column against the identity and the answer is +// symmetrised explicitly: a pivoted LU on a symmetric matrix may +// leave last-bit asymmetry the covariance must not carry. +func covarianceFromJac(name string, jac [][]float64, nP int) (*core.Array, error) { + a := make([][]float64, nP) + for i := range a { + a[i] = make([]float64, nP) + } + for k := range len(jac) { + row := jac[k] + for i := range nP { + xi := row[i] + ai := a[i] + for j := range nP { + ai[j] += xi * row[j] + } + } + } + rhs := make([][]float64, nP) + for i := range nP { + rhs[i] = make([]float64, nP) + rhs[i][i] = 1 + } + x, err := base.SolveSystem(name, a, rhs) + if err != nil { + return nil, base.Errf("%s: the Jacobian is rank-deficient at the answer, so no covariance exists: %w", name, err) + } + out := core.New(core.Float, nP, nP) + v := out.RawFloats() + for i := range nP { + for j := range nP { + v[i*nP+j] = (x[i][j] + x[j][i]) / 2 + } + } + return out, nil +} + +// runLevenbergMarquardt carries the fit. The legacy flag restores the +// historical error contract of LevenbergMarquardt: the same stalls +// the FitResult reports come back as errors with the messages the +// package has always published, so existing callers see nothing move. +func runLevenbergMarquardt(residual func(*core.Array) (*core.Array, error), p0 *core.Array, opts LMOptions, legacy bool) (*FitResult, error) { + if p0.Dtype() == core.Complex { + return nil, base.Errf("LevenbergMarquardt: complex parameters are not supported") + } + nP := p0.Len() + if nP == 0 { + return nil, base.Errf("LevenbergMarquardt: the parameter vector must not be empty") + } + if opts.MaxIterations <= 0 { + opts.MaxIterations = 200 + } + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-10 + } + if opts.Lambda <= 0 { + opts.Lambda = 1e-3 + } + + // cloneDense promotes through FloatAt, so Int and Float32 starting + // vectors behave exactly like Float64 ones (a RawFloats copy would + // silently start the fit from zeros for those dtypes). + p := cloneDense(p0) + nR := 0 + // whiten carries the Sigma factor; it is still nil for the very + // first evaluation, and the residual it returns is whitened by + // hand right after the factor is built. + var whiten *whitener + // evalR reads the residual at pp. The finiteness gate is strict + // for the states the fit adopts (the start point and every + // accepted iterate): a non-finite residual there poisons chi2 and + // every comparison against it, and the fit would die later as a + // bogus "the damping collapsed" diagnosis instead of the model's + // own fault. Backtracking trials take the lenient variant: a step + // into a saturating model is a candidate to damp past, not a dead + // run, the same recovery FindRootSystem's trials make. The Sigma + // whitening lands here, so every downstream consumer (chi2, the + // difference stencil, the trial comparison) works on the whitened + // residual and the weighted fit is the unweighted one on whitened + // data. + evalR := func(pp []float64, strict bool, dst []float64) ([]float64, error) { + a := linalg.ArrayFromFloatsSafe(pp, nP) + r, err := residual(a) + if err != nil { + return nil, err + } + if r.NDim() != 1 { + return nil, base.Errf("LevenbergMarquardt: the residual must be a vector") + } + if err := requireReal("LevenbergMarquardt", "residuals", r); err != nil { + return nil, err + } + if nR != 0 && r.Len() != nR { + return nil, base.Errf("LevenbergMarquardt: the residual length changed from %d to %d mid-fit", nR, r.Len()) + } + // dst carries a buffer the stencil reuses across columns; the + // states the fit keeps come back freshly allocated. Every entry + // of the buffer is written before it is read. + res := dst + if cap(res) < r.Len() { + res = make([]float64, r.Len()) + } + res = res[:r.Len()] + // Both branches fill res with the identical elements: the dense + // sweep reads the payload the accessor walk would widen. + if fs := denseFloats(r); fs != nil { + if strict { + for i, v := range fs { + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("LevenbergMarquardt: the residual returned the non-finite value %g at %d", v, i) + } + } + } + copy(res, fs) + } else { + for i := range r.Len() { + v := r.FloatAt(i) + if strict && (math.IsNaN(v) || math.IsInf(v, 0)) { + return nil, base.Errf("LevenbergMarquardt: the residual returned the non-finite value %g at %d", v, i) + } + res[i] = v + } + } + whiten.vector(res) + return res, nil + } + r, rerr := evalR(p, true, nil) + if rerr != nil { + return nil, base.Errf("LevenbergMarquardt: %w", rerr) + } + nR = len(r) + if nR < nP { + return nil, base.Errf("LevenbergMarquardt: underdetermined (%d obs, %d params)", nR, nP) + } + whiten, werr := sigmaWhitener(opts.Sigma, nR) + if werr != nil { + return nil, werr + } + whiten.vector(r) + chi2 := 0.0 + for i := range nR { + chi2 += r[i] * r[i] + } + + // buildJacobian assembles the row-major Jacobian at p, either from + // the caller's analytic callback or by central differences on the + // residual, one column per parameter. Its storage is allocated + // once and refilled per iteration: the sweep writes every entry. + jac := make([][]float64, nR) + for i := range nR { + jac[i] = make([]float64, nP) + } + // The difference stencils and the two residual vectors the columns + // are differenced from, allocated on first use and carried across + // the whole fit: a fit with an analytic Jacobian pays for none of + // them. The stencil carries the offset on one parameter at a time, + // restored as soon as the column is done, so neither a copy of the + // whole parameter vector nor a residual slice per column is needed. + var pp, pm []float64 + var resPlus, resMinus []float64 + buildJacobian := func(p []float64) error { + if opts.Jacobian != nil { + jm, err := opts.Jacobian(linalg.ArrayFromFloatsSafe(p, nP)) + if err != nil { + return base.Errf("LevenbergMarquardt: %w", err) + } + if jm.NDim() != 2 || jm.Shape()[0] != nR || jm.Shape()[1] != nP { + return base.Errf("LevenbergMarquardt: the Jacobian must be a %d×%d matrix, got shape %s", + nR, nP, base.ShapeText(jm.Shape())) + } + if err := requireReal("LevenbergMarquardt", "Jacobians", jm); err != nil { + return err + } + if fs := denseFloats(jm); fs != nil { + for i := range nR { + copy(jac[i], fs[i*nP:(i+1)*nP]) + } + } else { + for i := range nR { + ji := jac[i] + for j := range nP { + ji[j] = jm.FloatAt(i*nP + j) + } + } + } + whiten.matrix(jac) + return nil + } + // Central differences: evalR copies into the array handed to the + // callback, so nothing observes later mutation. column walks one + // parameter's stencil and writes that column of jac, and nothing + // else, so the bits it produces do not depend on which driver + // walks the columns. + column := func(j int, sp, sm, rp, rm []float64) ([]float64, []float64, error) { + eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(p[j])) + sp[j] += eps + sm[j] -= eps + rp, re1 := evalR(sp, true, rp) + rm, re2 := evalR(sm, true, rm) + sp[j], sm[j] = p[j], p[j] + if re1 != nil || re2 != nil { + return rp, rm, firstError(re1, re2) + } + for i := range nR { + jac[i][j] = (rp[i] - rm[i]) / (2 * eps) + } + return rp, rm, nil + } + if opts.ParallelJacobian { + // The consent the option records lets the columns go to the + // engine's workers: each goroutine owns a disjoint run of + // columns, writes only into those columns of jac and reads + // only p, so no two workers write the same address and the + // sweep needs no locks. A failing column records its error + // and abandons the rest of its run; the reported one is the + // lowest failing column, the one the serial walk would hit + // first. The residual buffers live per worker instead of + // being carried across columns: the option exists for + // expensive callbacks, where the carry buys nothing. + colErrs := make([]error, nP) + engine.ParallelMin(nP, 1, func(start, end int) { + sp, sm := make([]float64, nP), make([]float64, nP) + copy(sp, p) + copy(sm, p) + var rp, rm []float64 + for j := start; j < end; j++ { + var err error + rp, rm, err = column(j, sp, sm, rp, rm) + if err != nil { + colErrs[j] = err + return + } + } + }) + for _, err := range colErrs { + if err != nil { + return err + } + } + return nil + } + if cap(pp) < nP { + pp, pm = make([]float64, nP), make([]float64, nP) + } + pp, pm = pp[:nP], pm[:nP] + copy(pp, p) + copy(pm, p) + for j := range nP { + var err error + resPlus, resMinus, err = column(j, pp, pm, resPlus, resMinus) + if err != nil { + return err + } + } + return nil + } + + // The normal equations' storage, reused across iterations: the + // upper triangle of a is refilled by accumulation from an explicit + // zero and its lower one is mirrored back, bv is cleared likewise, + // and every other buffer is fully overwritten before it is read. + a := make([][]float64, nP) + for i := range nP { + a[i] = make([]float64, nP) + } + bv := make([]float64, nP) + pNew := make([]float64, nP) + // One right-hand-side header for the whole fit: the solve writes + // the step through bv in place, so the wrapper never changes. + solveRHS := [][]float64{bv} + lambda := opts.Lambda + status := FitBudget + var solveErr error + var collapseAt float64 + + // buildResult packs the fit's answer at the current p. The + // covariance rebuilds the Jacobian there: the loop's last one + // belongs to the point the last accepted step left behind, and the + // covariance must describe the point it is published beside. + buildResult := func() (*FitResult, error) { + out := core.New(core.Float, nP) + copy(out.RawFloats(), p) + res := &FitResult{Parameters: out, Chi2: chi2, Status: status} + if opts.RequestCovariance { + if jerr := buildJacobian(p); jerr != nil { + return nil, jerr + } + cov, cerr := covarianceFromJac("LevenbergMarquardt", jac, nP) + if cerr != nil { + return nil, cerr + } + res.Covariance = cov + } + return res, nil + } + + if chi2 == 0 { + // A start whose residual cancels exactly is already the perfect + // fit: the improvement test below is strict and cannot accept + // the zero step it produces, so the run would die in the + // damping collapse for being perfect. + status = FitConverged + return buildResult() + } + + for iter := 0; iter < opts.MaxIterations; iter++ { + if jerr := buildJacobian(p); jerr != nil { + return nil, jerr + } + + // a = JᵀJ + λ·diag(JᵀJ), bv = −Jᵀr. The accumulation walks the + // rows in ascending order, so each entry sums the same + // products in the same order the column-wise walk visited. + // The off-diagonal pair (i, j) and (j, i) sums the same + // products in the same order, the product commuting bitwise, + // so the pass runs the upper triangle alone and the mirror + // below reproduces the lower one exactly. + for i := range nP { + ai := a[i] + for j := i; j < nP; j++ { + ai[j] = 0 + } + } + clear(bv) + for k := range nR { + row := jac[k] + rk := r[k] + for i := range nP { + bv[i] -= row[i] * rk + } + for i := range nP { + xi := row[i] + ai := a[i] + for j := i; j < nP; j++ { + ai[j] += xi * row[j] + } + } + } + for i := range nP { + ai := a[i] + for j := i + 1; j < nP; j++ { + a[j][i] = ai[j] + } + } + for i := range nP { + a[i][i] *= (1 + lambda) + } + if opts.GradTol > 0 { + // The gradient test runs on bv before the solve: ‖bv‖∞ is + // ‖Jᵀr‖∞, and a gradient this small says the parameter + // directions carry nothing the step could spend, which is + // exactly the flat optimum the chi2 test alone never + // reaches. + gInf := 0.0 + for i := range nP { + gInf = max(gInf, math.Abs(bv[i])) + } + if gInf <= opts.GradTol { + status = FitConverged + return buildResult() + } + } + + delta, derr := base.SolveSystem("LevenbergMarquardt", a, solveRHS) + if derr != nil { + // A singular solve leaves the current point standing: it + // was good enough to build normal equations from, and no + // step replaced it. The fit reports it and stops. + status = FitStalled + solveErr = derr + break + } + + for j := range nP { + pNew[j] = p[j] + delta[0][j] + } + // A trial point: a saturating model here is damped past, and a + // finite chi2New admits only finite components, so an accepted + // trial never carries poison into the fit state. + rNew, rerr := evalR(pNew, false, nil) + if rerr != nil { + return nil, base.Errf("LevenbergMarquardt: %w", rerr) + } + chi2New := 0.0 + for i := range nR { + chi2New += rNew[i] * rNew[i] + } + + if chi2New < chi2 { + // The relative step test runs on the step just accepted, + // against the scale of the point it left: a step this small + // says the parameters have stopped moving meaningfully, + // whatever the residual still promises. + var stepSmall bool + if opts.StepTol > 0 { + stepInf, pInf := 0.0, 0.0 + for j := range nP { + stepInf = max(stepInf, math.Abs(delta[0][j])) + pInf = max(pInf, math.Abs(p[j])) + } + stepSmall = stepInf <= opts.StepTol*(pInf+opts.StepTol) + } + // The tolerance break must accept the better point too: + // reporting chi2New beside the old p publishes a fit quality + // the returned parameters do not achieve. + chi2Old := chi2 + copy(p, pNew) + r = rNew + chi2 = chi2New + if chi2Old-chi2 < opts.Tolerance*(1+chi2Old) { + status = FitConverged + return buildResult() + } + lambda *= 0.3 + if stepSmall { + status = FitConverged + return buildResult() + } + } else { + lambda *= 10 + if lambda > 1e20 { + // A damping that collapsed has no step left to take: + // the last point is a stall, never a converged answer. + status = FitStalled + collapseAt = lambda + break + } + } + } + + if legacy { + switch status { + case FitStalled: + if solveErr != nil { + return nil, base.Errf("LevenbergMarquardt: %w", solveErr) + } + return nil, base.Errf("LevenbergMarquardt: the damping collapsed to %g without the residual meeting the tolerance", collapseAt) + case FitBudget: + if !opts.AllowBudgetExit { + return nil, base.Errf("LevenbergMarquardt: the iteration budget of %d ran out without the residual meeting the tolerance", opts.MaxIterations) + } + case FitConverged: + } + } + return buildResult() +} diff --git a/optim/leastsq_mirror_pin_test.go b/optim/leastsq_mirror_pin_test.go new file mode 100644 index 0000000..e106c17 --- /dev/null +++ b/optim/leastsq_mirror_pin_test.go @@ -0,0 +1,51 @@ +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// A residual linear in the parameters makes one Gauss-Newton step the +// exact answer: the normal equations solve to the least-squares +// solution in a single iteration, so the fit's first step must land on +// it. Every off-diagonal entry of the three-parameter normal matrix is +// load-bearing there; a step computed from an asymmetric matrix lands +// somewhere else. +func TestLevenbergMarquardtSingleStepExactQuadratic(t *testing.T) { + // A is 12x3, well conditioned, nothing symmetric in its columns. + A := [][]float64{ + {1, 0.5, -0.25}, {2, -1, 0.75}, {-0.5, 1.5, 2}, {0.25, -0.75, 1}, + {1.5, 2, -1}, {-2, 0.25, 0.5}, {0.75, -2, 1.5}, {1, 1, 1}, + {-1, 0.5, 2}, {0.5, 2, -0.5}, {2, 1, -2}, {-0.25, -1.5, 0.25}, + } + truth := []float64{1.5, -0.75, 2.25} + y := make([]float64, len(A)) + for r := range A { + y[r] = A[r][0]*truth[0] + A[r][1]*truth[1] + A[r][2]*truth[2] + } + residual := func(p *core.Array) (*core.Array, error) { + pv := p.RawFloats() + out := make([]float64, len(A)) + for r := range A { + out[r] = A[r][0]*pv[0] + A[r][1]*pv[1] + A[r][2]*pv[2] - y[r] + } + return core.FromFloats(out, len(A)) + } + p0, err := core.FromFloats([]float64{0, 0, 0}, 3) + if err != nil { + t.Fatal(err) + } + opts := LMOptions{MaxIterations: 8} + pFit, _, err := LevenbergMarquardt(residual, p0, opts) + if err != nil { + t.Fatal(err) + } + got := pFit.RawFloats() + for k := range truth { + if math.Abs(got[k]-truth[k]) > 1e-9 { + t.Fatalf("parameter %d: the fit answered %g, the exact single step is %g", k, got[k], truth[k]) + } + } +} diff --git a/optim/leastsq_test.go b/optim/leastsq_test.go new file mode 100644 index 0000000..00b3a87 --- /dev/null +++ b/optim/leastsq_test.go @@ -0,0 +1,295 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestLevenbergMarquardt fits y = a·e^{−bx} + c to exact data and +// checks that the fitted parameters match the generating values +// within the convergence tolerance of the LM fitter. +func TestLevenbergMarquardt(t *testing.T) { + xData := []float64{0, 1, 2, 3, 4, 5, 6, 7, 8, 9} + yData := make([]float64, len(xData)) + for i, x := range xData { + yData[i] = 3*math.Exp(-0.5*float64(x)) + 0.5 + } + yArr := mustFloats(t, yData, len(yData)) + residual := func(p *core.Array) (*core.Array, error) { + a, b, c := p.FloatAt(0), p.FloatAt(1), p.FloatAt(2) + out := core.New(core.Float, len(xData)) + for i := range len(xData) { + out.RawFloats()[i] = yArr.FloatAt(i) - (a*math.Exp(-b*float64(xData[i])) + c) + } + return out, nil + } + p0 := mustFloats(t, []float64{2, 0.3, 0.1}, 3) + pOpt, chi2, err := LevenbergMarquardt(residual, p0, LMOptions{}) + if err != nil { + t.Fatalf("LevenbergMarquardt: %v", err) + } + if chi2 > 1e-4 { + t.Fatalf("χ² = %v, want < 1e-4 for exact data", chi2) + } + a := pOpt.FloatAt(0) + b := pOpt.FloatAt(1) + if math.Abs(a-3) > 0.01 || math.Abs(b-0.5) > 0.01 { + t.Fatalf("a = %.8f, b = %.8f, want ≈ 3, 0.5", a, b) + } +} + +// TestLevenbergMarquardtErrors pins the error contract. +func TestLevenbergMarquardtErrors(t *testing.T) { + f := func(p *core.Array) (*core.Array, error) { return nil, nil } + cx, _ := core.FromComplexes([]complex128{1}, 1) + if _, _, err := LevenbergMarquardt(f, cx, LMOptions{}); err == nil { + t.Fatal("expected an error for complex parameters") + } + empty, _ := core.FromFloats(nil, 0) + if _, _, err := LevenbergMarquardt(f, empty, LMOptions{}); err == nil { + t.Fatal("expected an error for an empty parameter vector") + } +} + +// TestLevenbergMarquardtAnalyticJacobian fits the same exponential +// model twice, once with the analytic Jacobian and once with the +// central-difference fallback: both must land on the same optimum, +// the analytic one reaching it without the extra evaluations. +func TestLevenbergMarquardtAnalyticJacobian(t *testing.T) { + xData := []float64{0, 1, 2, 3, 4, 5, 6, 7, 8, 9} + yData := make([]float64, len(xData)) + for i, x := range xData { + yData[i] = 3*math.Exp(-0.5*float64(x)) + 0.5 + } + residual := func(p *core.Array) (*core.Array, error) { + a, b, c := p.FloatAt(0), p.FloatAt(1), p.FloatAt(2) + out := core.New(core.Float, len(xData)) + for i := range len(xData) { + out.RawFloats()[i] = yData[i] - (a*math.Exp(-b*float64(xData[i])) + c) + } + return out, nil + } + jacobian := func(p *core.Array) (*core.Array, error) { + a, b := p.FloatAt(0), p.FloatAt(1) + out := core.New(core.Float, len(xData), 3) + for i := range len(xData) { + e := math.Exp(-b * float64(xData[i])) + out.RawFloats()[i*3+0] = -e + out.RawFloats()[i*3+1] = a * float64(xData[i]) * e + out.RawFloats()[i*3+2] = -1 + } + return out, nil + } + p0 := mustFloats(t, []float64{2, 0.3, 0.1}, 3) + pA, chi2A, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: jacobian}) + if err != nil { + t.Fatalf("LevenbergMarquardt analytic: %v", err) + } + pD, chi2D, err := LevenbergMarquardt(residual, p0, LMOptions{}) + if err != nil { + t.Fatalf("LevenbergMarquardt differences: %v", err) + } + for j := range 3 { + if math.Abs(pA.FloatAt(j)-pD.FloatAt(j)) > 1e-6 { + t.Fatalf("parameter %d: analytic %.12g, differences %.12g", j, pA.FloatAt(j), pD.FloatAt(j)) + } + } + if chi2A > 1e-4 { + t.Fatalf("analytic χ² = %v, want < 1e-4", chi2A) + } + if math.Abs(chi2A-chi2D) > 1e-8 { + t.Fatalf("χ² disagree: analytic %v, differences %v", chi2A, chi2D) + } +} + +// TestLevenbergMarquardtAnalyticExact solves a determined linear +// system, where the Gauss-Newton step is exact and the analytic +// Jacobian reaches the solution the equations dictate. +func TestLevenbergMarquardtAnalyticExact(t *testing.T) { + residual := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{ + p.FloatAt(0) + 2*p.FloatAt(1) - 3, + 2*p.FloatAt(0) + p.FloatAt(1) - 4, + }, 2) + } + jacobian := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 2, 2, 1}, 2, 2) + } + p0 := mustFloats(t, []float64{0, 0}, 2) + pOpt, chi2, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: jacobian, Tolerance: 1e-14}) + if err != nil { + t.Fatalf("LevenbergMarquardt: %v", err) + } + if math.Abs(pOpt.FloatAt(0)-5.0/3) > 1e-8 || math.Abs(pOpt.FloatAt(1)-2.0/3) > 1e-8 { + t.Fatalf("p = (%.12g, %.12g), want (5/3, 2/3)", pOpt.FloatAt(0), pOpt.FloatAt(1)) + } + if chi2 > 1e-16 { + t.Fatalf("χ² = %v, want 0", chi2) + } +} + +// TestLevenbergMarquardtJacobianErrors pins the analytic Jacobian's +// own contract: shape mismatches and callback errors surface. +func TestLevenbergMarquardtJacobianErrors(t *testing.T) { + residual := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{p.FloatAt(0) - 1, p.FloatAt(1) - 2}, 2) + } + rank1 := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 1}, 2) + } + if _, _, err := LevenbergMarquardt(residual, mustFloats(t, []float64{0, 0}, 2), + LMOptions{Jacobian: rank1}); err == nil { + t.Fatal("expected an error for a rank-1 Jacobian") + } + wrongDims := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 1, 1}, 3, 1) + } + if _, _, err := LevenbergMarquardt(residual, mustFloats(t, []float64{0, 0}, 2), + LMOptions{Jacobian: wrongDims}); err == nil { + t.Fatal("expected an error for a Jacobian of the wrong dimensions") + } + boom := func(p *core.Array) (*core.Array, error) { + return nil, base.Errf("jacobian failed") + } + if _, _, err := LevenbergMarquardt(residual, mustFloats(t, []float64{0, 0}, 2), + LMOptions{Jacobian: boom}); err == nil { + t.Fatal("expected the Jacobian error to propagate") + } +} + +// TestLevenbergMarquardtDtypes pins the parameter-vector promotion: +// Int and Float32 starts must behave exactly like Float64 ones. The +// old RawFloats() copy silently started those fits from zeros. +func TestLevenbergMarquardtDtypes(t *testing.T) { + residual := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{ + p.FloatAt(0) + 2*p.FloatAt(1) - 3, + 2*p.FloatAt(0) + p.FloatAt(1) - 4, + }, 2) + } + jacobian := func(*core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 2, 2, 1}, 2, 2) + } + ints, err := core.FromInts([]int64{0, 0}, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + thirtyTwo, err := core.FromFloat32s([]float32{0, 0}, 2) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + for name, p0 := range map[string]*core.Array{"int": ints, "float32": thirtyTwo} { + pOpt, chi2, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: jacobian, Tolerance: 1e-14}) + if err != nil { + t.Fatalf("LevenbergMarquardt(%s start): %v", name, err) + } + if math.Abs(pOpt.FloatAt(0)-5.0/3) > 1e-8 || math.Abs(pOpt.FloatAt(1)-2.0/3) > 1e-8 { + t.Fatalf("LevenbergMarquardt(%s start): p = (%.12g, %.12g), want (5/3, 2/3)", + name, pOpt.FloatAt(0), pOpt.FloatAt(1)) + } + if chi2 > 1e-16 { + t.Fatalf("LevenbergMarquardt(%s start): χ² = %v, want 0", name, chi2) + } + } +} + +// TestLevenbergMarquardtResidualLengthChange pins the residual +// contract: a callback whose output length changes mid-fit is an +// error naming the mismatch, not a shape panic. +func TestLevenbergMarquardtResidualLengthChange(t *testing.T) { + calls := 0 + residual := func(*core.Array) (*core.Array, error) { + calls++ + if calls == 1 { + return core.FromFloats([]float64{1, 2, 3}, 3) + } + return core.FromFloats([]float64{1, 2}, 2) + } + _, _, err := LevenbergMarquardt(residual, mustFloats(t, []float64{0, 0}, 2), LMOptions{}) + if err == nil { + t.Fatal("residual length change mid-fit: want an error") + } +} + +// TestLevenbergMarquardtStencilMatchesAnalyticJacobian pins the reused +// difference stencil against the analytic route. The stencil carries +// the offset on one parameter at a time and puts it back as soon as the +// column is differenced, so each column's Jacobian is the central +// difference of that parameter alone and the two routes land on the +// same point. A stencil that leaves its offset behind differences every +// later column at a point that is also displaced earlier: a different +// Jacobian, and a fit that parts company with the analytic route. +func TestLevenbergMarquardtStencilMatchesAnalyticJacobian(t *testing.T) { + // A decay with an offset and a smooth nuisance term, so every + // parameter carries curvature and the columns are coupled. + const n = 40 + xs := make([]float64, n) + ys := make([]float64, n) + for i := range n { + xs[i] = float64(i) * 0.25 + ys[i] = 1.5*math.Exp(-0.8*xs[i]) + 0.7 + 0.05*math.Sin(2*xs[i]) + } + residual := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, n) + v := out.RawFloats() + for i := range n { + v[i] = ys[i] - (p.FloatAt(0)*math.Exp(-p.FloatAt(1)*xs[i]) + + p.FloatAt(2) + p.FloatAt(3)*math.Sin(2*xs[i])) + } + return out, nil + } + jacobian := func(p *core.Array) (*core.Array, error) { + v := make([]float64, n*4) + for i := range n { + e := math.Exp(-p.FloatAt(1) * xs[i]) + v[i*4+0] = -e + v[i*4+1] = p.FloatAt(0) * xs[i] * e + v[i*4+2] = -1 + v[i*4+3] = -math.Sin(2 * xs[i]) + } + return core.FromFloats(v, n, 4) + } + start, err := core.FromFloats([]float64{1.2, 0.9, 0.4, 0.04}, 4) + if err != nil { + t.Fatal(err) + } + fd, fdChi, err := LevenbergMarquardt(residual, start, LMOptions{}) + if err != nil { + t.Fatalf("LevenbergMarquardt with a difference stencil: %v", err) + } + analytic, analyticChi, err := LevenbergMarquardt(residual, start, LMOptions{Jacobian: jacobian}) + if err != nil { + t.Fatalf("LevenbergMarquardt with an analytic Jacobian: %v", err) + } + if fd.Len() != analytic.Len() { + t.Fatalf("the two routes fit %d and %d parameters", fd.Len(), analytic.Len()) + } + for i := range fd.Len() { + if d := math.Abs(fd.FloatAt(i) - analytic.FloatAt(i)); d > 1e-12 { + t.Fatalf("parameter %d: the stencil route gives %.17g, the analytic one %.17g (differ by %g)", + i, fd.FloatAt(i), analytic.FloatAt(i), d) + } + } + if math.Abs(fdChi-analyticChi) > 1e-12 { + t.Fatalf("the two routes report chi2 %.17g and %.17g", fdChi, analyticChi) + } + // The stencil is carried across the fit's iterations, so the same + // start must reproduce the same bits. + again, againChi, err := LevenbergMarquardt(residual, start, LMOptions{}) + if err != nil { + t.Fatalf("the repeated stencil fit: %v", err) + } + for i := range again.Len() { + if math.Float64bits(again.FloatAt(i)) != math.Float64bits(fd.FloatAt(i)) { + t.Fatalf("parameter %d moved between stencil fits: %.17g against %.17g", i, again.FloatAt(i), fd.FloatAt(i)) + } + } + if math.Float64bits(againChi) != math.Float64bits(fdChi) { + t.Fatalf("chi2 moved between stencil fits: %.17g against %.17g", againChi, fdChi) + } +} diff --git a/optim/leastsqfit_test.go b/optim/leastsqfit_test.go new file mode 100644 index 0000000..aa6b212 --- /dev/null +++ b/optim/leastsqfit_test.go @@ -0,0 +1,609 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "math/big" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// linearPair returns the determined linear pair r(p) = (p₀ + 2p₁ − 3, +// 2p₀ + p₁ − 4) with its constant Jacobian, the model whose +// Gauss-Newton step is exact. +func linearPair() (func(*core.Array) (*core.Array, error), func(*core.Array) (*core.Array, error)) { + residual := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{ + p.FloatAt(0) + 2*p.FloatAt(1) - 3, + 2*p.FloatAt(0) + p.FloatAt(1) - 4, + }, 2) + } + jacobian := func(*core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 2, 2, 1}, 2, 2) + } + return residual, jacobian +} + +// TestLevenbergMarquardtFitConverged pins the Fit surface on the model +// the plain entry point already covers: same point, same chi2, and a +// status that says converged. +func TestLevenbergMarquardtFitConverged(t *testing.T) { + residual, jacobian := linearPair() + p0 := mustFloats(t, []float64{0, 0}, 2) + res, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian, Tolerance: 1e-14}) + if err != nil { + t.Fatalf("LevenbergMarquardtFit: %v", err) + } + if res.Status != FitConverged { + t.Fatalf("status = %d, want FitConverged", res.Status) + } + if math.Abs(res.Parameters.FloatAt(0)-5.0/3) > 1e-8 || math.Abs(res.Parameters.FloatAt(1)-2.0/3) > 1e-8 { + t.Fatalf("p = (%.12g, %.12g), want (5/3, 2/3)", res.Parameters.FloatAt(0), res.Parameters.FloatAt(1)) + } + if res.Chi2 > 1e-16 { + t.Fatalf("chi2 = %v, want 0", res.Chi2) + } + if res.Covariance != nil { + t.Fatal("a fit that did not ask for a covariance returned one") + } + legacyP, legacyChi2, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: jacobian, Tolerance: 1e-14}) + if err != nil { + t.Fatalf("LevenbergMarquardt: %v", err) + } + if legacyP.FloatAt(0) != res.Parameters.FloatAt(0) || legacyP.FloatAt(1) != res.Parameters.FloatAt(1) { + t.Fatal("the two entry points returned different parameters") + } + if legacyChi2 != res.Chi2 { + t.Fatalf("the two entry points reported chi2 %v and %v", legacyChi2, res.Chi2) + } +} + +// TestLevenbergMarquardtGradientExit pins the gradient test: with the +// chi2 tolerance set far below anything the fit can reach, the only +// exit that can fire after the first accepted step is ‖Jᵀr‖∞. The +// residual call count proves it: the start and the one trial make +// two, and a third call would mean the run walked into a second +// solve, which the gradient test forbids. +func TestLevenbergMarquardtGradientExit(t *testing.T) { + residual, jacobian := linearPair() + calls := 0 + counted := func(p *core.Array) (*core.Array, error) { + calls++ + return residual(p) + } + p0 := mustFloats(t, []float64{0, 0}, 2) + // Lambda 1e-300 leaves the damping factor 1+λ bit-identically 1, + // so the first step is the exact Gauss-Newton one and the run + // reaches the solution in one acceptance. + res, err := LevenbergMarquardtFit(counted, p0, LMOptions{ + Jacobian: jacobian, + Tolerance: 1e-30, + GradTol: 1e-7, + Lambda: 1e-300, + }) + if err != nil { + t.Fatalf("LevenbergMarquardtFit: %v", err) + } + if res.Status != FitConverged { + t.Fatalf("status = %d, want FitConverged", res.Status) + } + if calls != 2 { + t.Fatalf("the residual was evaluated %d times, want 2 (the gradient test must stop the run before its second solve)", calls) + } + if math.Abs(res.Parameters.FloatAt(0)-5.0/3) > 1e-8 || math.Abs(res.Parameters.FloatAt(1)-2.0/3) > 1e-8 { + t.Fatalf("p = (%.12g, %.12g), want (5/3, 2/3)", res.Parameters.FloatAt(0), res.Parameters.FloatAt(1)) + } +} + +// TestLevenbergMarquardtStepExit pins the step test: a linear +// residual whose first accepted step is negligible against the +// point's own scale converges on that step, with the chi2 tolerance +// far too tight to claim the exit. +func TestLevenbergMarquardtStepExit(t *testing.T) { + residual := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{10 * (p.FloatAt(0) - 10)}, 1) + } + jacobian := func(*core.Array) (*core.Array, error) { + return core.FromFloats([]float64{10}, 1, 1) + } + counted := func(count *int) func(*core.Array) (*core.Array, error) { + return func(p *core.Array) (*core.Array, error) { + *count++ + return residual(p) + } + } + stepCalls := 0 + res, err := LevenbergMarquardtFit(counted(&stepCalls), mustFloats(t, []float64{10.001}, 1), LMOptions{ + Jacobian: jacobian, + Tolerance: 1e-15, + StepTol: 0.5, + }) + if err != nil { + t.Fatalf("LevenbergMarquardtFit: %v", err) + } + if res.Status != FitConverged { + t.Fatalf("status = %d, want FitConverged", res.Status) + } + if stepCalls != 2 { + t.Fatalf("the step-exit run evaluated the residual %d times, want 2", stepCalls) + } + if math.Abs(res.Parameters.FloatAt(0)-10) > 1e-4 { + t.Fatalf("p = %.12g, want 10", res.Parameters.FloatAt(0)) + } + // Without the step test the same run needs the chi2 tolerance, and + // at 1e-6 that fires one acceptance later: three evaluations. + wideCalls := 0 + wide, err := LevenbergMarquardtFit(counted(&wideCalls), mustFloats(t, []float64{10.001}, 1), LMOptions{ + Jacobian: jacobian, + Tolerance: 1e-6, + }) + if err != nil { + t.Fatalf("LevenbergMarquardtFit without StepTol: %v", err) + } + if wide.Status != FitConverged { + t.Fatalf("status = %d, want FitConverged", wide.Status) + } + if wideCalls != 3 { + t.Fatalf("the control run evaluated the residual %d times, want 3", wideCalls) + } +} + +// TestLevenbergMarquardtFitStalled pins the stalled contract: a +// parameter the residual never sees makes the normal equations +// singular on the first iteration, and the fit reports the start +// point back with FitStalled instead of an error. The legacy entry +// point keeps its historical error for the same run. +func TestLevenbergMarquardtFitStalled(t *testing.T) { + residual := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{p.FloatAt(1) - 1, p.FloatAt(1) - 1}, 2) + } + jacobian := func(*core.Array) (*core.Array, error) { + return core.FromFloats([]float64{0, 1, 0, 1}, 2, 2) + } + p0 := mustFloats(t, []float64{5, 5}, 2) + res, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian}) + if err != nil { + t.Fatalf("LevenbergMarquardtFit: %v", err) + } + if res.Status != FitStalled { + t.Fatalf("status = %d, want FitStalled", res.Status) + } + if res.Parameters.FloatAt(0) != 5 || res.Parameters.FloatAt(1) != 5 { + t.Fatalf("p = (%.12g, %.12g), want the untouched start", res.Parameters.FloatAt(0), res.Parameters.FloatAt(1)) + } + if res.Chi2 != 32 { + t.Fatalf("chi2 = %v, want 32", res.Chi2) + } + if _, _, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: jacobian}); err == nil { + t.Fatal("the legacy entry point must keep reporting the singular solve as an error") + } +} + +// TestLevenbergMarquardtFitBudget pins the budget report: a fit given +// two iterations of a problem that needs more returns its last point +// with FitBudget and no error, and the legacy entry point refuses the +// same run unless AllowBudgetExit is set. +func TestLevenbergMarquardtFitBudget(t *testing.T) { + xData := []float64{0, 1, 2, 3, 4, 5, 6, 7, 8, 9} + residual := func(p *core.Array) (*core.Array, error) { + a, b, c := p.FloatAt(0), p.FloatAt(1), p.FloatAt(2) + out := core.New(core.Float, len(xData)) + for i := range len(xData) { + out.RawFloats()[i] = 3*math.Exp(-0.5*float64(xData[i])) + 0.5 - + (a*math.Exp(-b*float64(xData[i])) + c) + } + return out, nil + } + p0 := mustFloats(t, []float64{20, 5, 20}, 3) + res, err := LevenbergMarquardtFit(residual, p0, LMOptions{MaxIterations: 2}) + if err != nil { + t.Fatalf("LevenbergMarquardtFit: %v", err) + } + if res.Status != FitBudget { + t.Fatalf("status = %d, want FitBudget", res.Status) + } + if res.Parameters == nil { + t.Fatal("a budget stop must still report its last point") + } + if _, _, err := LevenbergMarquardt(residual, p0, LMOptions{MaxIterations: 2}); err == nil { + t.Fatal("the legacy entry point must refuse a budget stop without AllowBudgetExit") + } + if _, _, err := LevenbergMarquardt(residual, p0, LMOptions{MaxIterations: 2, AllowBudgetExit: true}); err != nil { + t.Fatalf("LevenbergMarquardt with AllowBudgetExit: %v", err) + } +} + +// weightedModel builds the three-observation linear model r(p) = y − +// Xp with its Jacobian, small enough for the exact rational referent. +func weightedModel() (func(*core.Array) (*core.Array, error), func(*core.Array) (*core.Array, error)) { + x := [][]float64{{1, 0}, {0, 1}, {1, 1}} + y := []float64{1.2, 0.7, 2.4} + residual := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, 3) + for i := range 3 { + out.RawFloats()[i] = y[i] - (x[i][0]*p.FloatAt(0) + x[i][1]*p.FloatAt(1)) + } + return out, nil + } + jacobian := func(*core.Array) (*core.Array, error) { + return core.FromFloats([]float64{-1, 0, 0, -1, -1, -1}, 3, 2) + } + return residual, jacobian +} + +// ratSolve solves a x = b exactly over the rationals by elimination +// with nonzero pivots; a singular system returns nil. The inputs are +// copied: big.Rat values are pointers, and the elimination mutates +// every entry it touches. +func ratSolve(a, b [][]*big.Rat) [][]*big.Rat { + n := len(a) + m := make([][]*big.Rat, n) + for i := range n { + m[i] = make([]*big.Rat, 0, len(a[i])+len(b[i])) + for _, v := range a[i] { + m[i] = append(m[i], new(big.Rat).Set(v)) + } + for _, v := range b[i] { + m[i] = append(m[i], new(big.Rat).Set(v)) + } + } + zero := new(big.Rat) + for col := range n { + piv := -1 + for row := col; row < n; row++ { + if m[row][col].Cmp(zero) != 0 { + piv = row + break + } + } + if piv < 0 { + return nil + } + m[col], m[piv] = m[piv], m[col] + inv := new(big.Rat).Inv(m[col][col]) + for j := col; j < len(m[col]); j++ { + m[col][j].Mul(m[col][j], inv) + } + for row := range n { + if row == col || m[row][col].Cmp(zero) == 0 { + continue + } + f := new(big.Rat).Set(m[row][col]) + for j := col; j < len(m[row]); j++ { + m[row][j].Sub(m[row][j], new(big.Rat).Mul(f, m[col][j])) + } + } + } + out := make([][]*big.Rat, n) + for i := range n { + out[i] = m[i][n:] + } + return out +} + +// ratFroms converts float columns into exact rationals. +func ratFroms(vs []float64) []*big.Rat { + out := make([]*big.Rat, len(vs)) + for i, v := range vs { + out[i] = new(big.Rat).SetFloat64(v) + } + return out +} + +// ratColumn turns a vector into the one-column right-hand side +// ratSolve takes. +func ratColumn(vs []*big.Rat) [][]*big.Rat { + out := make([][]*big.Rat, len(vs)) + for i, v := range vs { + out[i] = []*big.Rat{v} + } + return out +} + +// ratMatrixInvert inverts a small square rational matrix. +func ratMatrixInvert(a [][]*big.Rat) [][]*big.Rat { + n := len(a) + id := make([][]*big.Rat, n) + for i := range n { + row := make([]*big.Rat, n) + for j := range n { + row[j] = new(big.Rat) + if i == j { + row[j].SetInt64(1) + } + } + id[i] = row + } + return ratSolve(a, id) +} + +// TestLevenbergMarquardtWeights pins the weighted fit against an +// exact rational referent: chi2 becomes rᵀC⁻¹r, the parameters +// minimise it, and the diagonal and matrix forms of Sigma agree on +// the same weighting. +func TestLevenbergMarquardtWeights(t *testing.T) { + residual, jacobian := weightedModel() + cov := []float64{ + 2, 1, 0, + 1, 3, 1, + 0, 1, 2, + } + sigma, err := core.FromFloats(cov, 3, 3) + if err != nil { + t.Fatal(err) + } + p0 := mustFloats(t, []float64{0, 0}, 2) + res, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian, Sigma: sigma}) + if err != nil { + t.Fatalf("LevenbergMarquardtFit: %v", err) + } + if res.Status != FitConverged { + t.Fatalf("status = %d, want FitConverged", res.Status) + } + + // The exact referent: r = y − Xp, so the weighted optimum solves + // (XᵀC⁻¹X) p = XᵀC⁻¹y. + xf := [][]float64{{1, 0}, {0, 1}, {1, 1}} + yf := []float64{1.2, 0.7, 2.4} + cr := make([][]*big.Rat, 3) + for i := range 3 { + cr[i] = ratFroms(cov[i*3 : i*3+3]) + } + cinv := ratMatrixInvert(cr) + if cinv == nil { + t.Fatal("the referent covariance is singular") + } + xr := make([][]*big.Rat, 3) + for i := range 3 { + xr[i] = ratFroms(xf[i]) + } + yr := ratFroms(yf) + cinvX := ratSolve(cr, xr) + normal := make([][]*big.Rat, 2) + for i := range 2 { + normal[i] = make([]*big.Rat, 2) + for j := range 2 { + s := new(big.Rat) + for k := range 3 { + s.Add(s, new(big.Rat).Mul(xr[k][i], cinvX[k][j])) + } + normal[i][j] = s + } + } + rhs := make([]*big.Rat, 2) + cinvY := ratSolve(cr, ratColumn(yr)) + if cinvY == nil { + t.Fatal("the referent solve failed") + } + for i := range 2 { + s := new(big.Rat) + for k := range 3 { + s.Add(s, new(big.Rat).Mul(xr[k][i], cinvY[k][0])) + } + rhs[i] = s + } + wantP := ratSolve(normal, ratColumn(rhs)) + if wantP == nil { + t.Fatal("the referent normal equations are singular") + } + for j := range 2 { + got := new(big.Rat).SetFloat64(res.Parameters.FloatAt(j)) + d := new(big.Rat).Sub(got, wantP[j][0]) + d.Abs(d) + if d.Cmp(new(big.Rat).SetFloat64(1e-10)) > 0 { + t.Fatalf("parameter %d = %.12g, want %s", j, res.Parameters.FloatAt(j), wantP[j][0].FloatString(12)) + } + } + // The weighted chi2: rᵀC⁻¹r at the returned point. + rAt := make([]*big.Rat, 3) + for k := range 3 { + rAt[k] = new(big.Rat).SetFloat64(yf[k] - (xf[k][0]*res.Parameters.FloatAt(0) + xf[k][1]*res.Parameters.FloatAt(1))) + } + cinvR := make([]*big.Rat, 3) + for k := range 3 { + s := new(big.Rat) + for j := range 3 { + s.Add(s, new(big.Rat).Mul(cinv[k][j], rAt[j])) + } + cinvR[k] = s + } + chiRef := new(big.Rat) + for k := range 3 { + chiRef.Add(chiRef, new(big.Rat).Mul(rAt[k], cinvR[k])) + } + chiFloat, _ := chiRef.Float64() + if math.Abs(res.Chi2-chiFloat) > 1e-12*math.Max(1, math.Abs(chiFloat)) { + t.Fatalf("chi2 = %.17g, want %.17g", res.Chi2, chiFloat) + } + // The weighted answer must differ from the unweighted one, or the + // test proves nothing about the weighting. + plain, _, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: jacobian}) + if err != nil { + t.Fatalf("the unweighted fit: %v", err) + } + if plain.FloatAt(0) == res.Parameters.FloatAt(0) && plain.FloatAt(1) == res.Parameters.FloatAt(1) { + t.Fatal("the weighted and unweighted fits returned identical parameters") + } + + // The diagonal form weights by the same matrix's diagonal, and the + // legacy entry point carries Sigma too. + variance, err := core.FromFloats([]float64{2, 3, 2}, 3) + if err != nil { + t.Fatal(err) + } + diagRes, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian, Sigma: variance}) + if err != nil { + t.Fatalf("the diagonal Sigma fit: %v", err) + } + if diagRes.Status != FitConverged { + t.Fatalf("the diagonal Sigma status = %d, want FitConverged", diagRes.Status) + } + wp, _, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: jacobian, Sigma: sigma}) + if err != nil { + t.Fatalf("the legacy weighted fit: %v", err) + } + if wp.FloatAt(0) != res.Parameters.FloatAt(0) || wp.FloatAt(1) != res.Parameters.FloatAt(1) { + t.Fatal("the legacy weighted fit returned different parameters") + } +} + +// TestLevenbergMarquardtWeightedCovariance pins the covariance under +// Sigma: (JᵀC⁻¹J)⁻¹ against the exact rational inverse, entry by +// entry. +func TestLevenbergMarquardtWeightedCovariance(t *testing.T) { + residual, jacobian := weightedModel() + cov := []float64{ + 2, 1, 0, + 1, 3, 1, + 0, 1, 2, + } + sigma, err := core.FromFloats(cov, 3, 3) + if err != nil { + t.Fatal(err) + } + xf := [][]float64{{1, 0}, {0, 1}, {1, 1}} + cr := make([][]*big.Rat, 3) + for i := range 3 { + cr[i] = ratFroms(cov[i*3 : i*3+3]) + } + xr := make([][]*big.Rat, 3) + for i := range 3 { + xr[i] = ratFroms(xf[i]) + } + cinvX := ratSolve(cr, xr) + normal := make([][]*big.Rat, 2) + for i := range 2 { + normal[i] = make([]*big.Rat, 2) + for j := range 2 { + s := new(big.Rat) + for k := range 3 { + s.Add(s, new(big.Rat).Mul(xr[k][i], cinvX[k][j])) + } + normal[i][j] = s + } + } + want := ratMatrixInvert(normal) + if want == nil { + t.Fatal("the referent normal matrix is singular") + } + res, err := LevenbergMarquardtFit(residual, mustFloats(t, []float64{0, 0}, 2), LMOptions{ + Jacobian: jacobian, + Sigma: sigma, + RequestCovariance: true, + }) + if err != nil { + t.Fatalf("LevenbergMarquardtFit: %v", err) + } + if res.Covariance == nil { + t.Fatal("the requested covariance is missing") + } + if res.Covariance.NDim() != 2 || res.Covariance.Shape()[0] != 2 || res.Covariance.Shape()[1] != 2 { + t.Fatalf("covariance shape %s, want 2×2", base.ShapeText(res.Covariance.Shape())) + } + for i := range 2 { + for j := range 2 { + got := res.Covariance.FloatAt(i*2 + j) + wantFloat, _ := want[i][j].Float64() + if math.Abs(got-wantFloat) > 1e-9*math.Max(1, math.Abs(wantFloat)) { + t.Fatalf("covariance (%d, %d) = %.17g, want %.17g", i, j, got, wantFloat) + } + mirror := res.Covariance.FloatAt(j*2 + i) + if got != mirror { + t.Fatalf("covariance (%d, %d) = %.17g against (%d, %d) = %.17g", i, j, got, j, i, mirror) + } + } + } +} + +// TestLevenbergMarquardtUnweightedCovariance pins the unweighted +// covariance (JᵀJ)⁻¹ against the exact rational inverse, including on +// the perfect start whose fit never enters the iteration loop. +func TestLevenbergMarquardtUnweightedCovariance(t *testing.T) { + residual, jacobian := weightedModel() + xf := [][]float64{{1, 0}, {0, 1}, {1, 1}} + xr := make([][]*big.Rat, 3) + for i := range 3 { + xr[i] = ratFroms(xf[i]) + } + xtX := make([][]*big.Rat, 2) + for i := range 2 { + xtX[i] = make([]*big.Rat, 2) + for j := range 2 { + s := new(big.Rat) + for k := range 3 { + s.Add(s, new(big.Rat).Mul(xr[k][i], xr[k][j])) + } + xtX[i][j] = s + } + } + want := ratMatrixInvert(xtX) + if want == nil { + t.Fatal("the referent normal matrix is singular") + } + starts := [][]float64{{0, 0}, {1.2, 0.7}} + for _, s := range starts { + res, err := LevenbergMarquardtFit(residual, mustFloats(t, s, 2), LMOptions{ + Jacobian: jacobian, + RequestCovariance: true, + }) + if err != nil { + t.Fatalf("LevenbergMarquardtFit(%v): %v", s, err) + } + if res.Covariance == nil { + t.Fatalf("LevenbergMarquardtFit(%v): the requested covariance is missing", s) + } + for i := range 2 { + for j := range 2 { + got := res.Covariance.FloatAt(i*2 + j) + wantFloat, _ := want[i][j].Float64() + if math.Abs(got-wantFloat) > 1e-9*math.Max(1, math.Abs(wantFloat)) { + t.Fatalf("start %v: covariance (%d, %d) = %.17g, want %.17g", s, i, j, got, wantFloat) + } + } + } + } +} + +// TestLevenbergMarquardtSigmaErrors pins the Sigma contract: shapes, +// positive variances, exact symmetry and positive definiteness are +// all checked before the fit moves. +func TestLevenbergMarquardtSigmaErrors(t *testing.T) { + residual, jacobian := weightedModel() + p0 := mustFloats(t, []float64{0, 0}, 2) + cases := []struct { + name string + sigma *core.Array + want string + }{ + {"wrong length", mustFloats(t, []float64{1, 2}, 2), "one variance per residual"}, + {"wrong shape", mustFloats(t, []float64{1, 0, 0, 1, 0, 0}, 2, 3), "shape (2, 3)"}, + {"zero variance", mustFloats(t, []float64{1, 0, 1}, 3), "positive variances"}, + {"negative variance", mustFloats(t, []float64{1, -2, 1}, 3), "positive variances"}, + } + for _, tc := range cases { + if _, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian, Sigma: tc.sigma}); err == nil { + t.Fatalf("%s: want an error", tc.name) + } else if !strings.Contains(err.Error(), tc.want) { + t.Fatalf("%s: error %q, want it to mention %q", tc.name, err, tc.want) + } + } + asymmetric := mustFloats(t, []float64{2, 1, 0, 1, 3, 1, 2, 1, 2}, 3, 3) + if _, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian, Sigma: asymmetric}); err == nil { + t.Fatal("an asymmetric Sigma: want an error") + } else if !strings.Contains(err.Error(), "symmetric") { + t.Fatalf("asymmetric Sigma: error %q, want it to mention symmetry", err) + } + indefinite := mustFloats(t, []float64{1, 2, 2, 1}, 2, 2) + wrongRows := mustFloats(t, []float64{1, 0, 0, 1}, 2, 2) + if _, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian, Sigma: indefinite}); err == nil { + t.Fatal("an indefinite Sigma: want an error") + } + if _, err := LevenbergMarquardtFit(residual, p0, LMOptions{Jacobian: jacobian, Sigma: wrongRows}); err == nil { + t.Fatal("a Sigma of the wrong row count: want an error") + } +} diff --git a/optim/linear.go b/optim/linear.go new file mode 100644 index 0000000..7c0fef2 --- /dev/null +++ b/optim/linear.go @@ -0,0 +1,308 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/linalg" +) + +// Linear constraints handled by the augmented Lagrangian. Each row of +// A carries a two-sided bound l ≤ A·x ≤ u: a row with l = u is an +// equality constraint, a one-sided row opens the other side with an +// infinity. The outer loop escalates the quadratic penalty and ascends +// the row multipliers until the true violations vanish; the inner +// problem is the box-capable L-BFGS, so the options bounds compose +// with the linear rows. + +// LinearConstraints holds the rows of l ≤ A·x ≤ u. A is an r×n matrix +// over the n variables, Lower and Upper hold one entry per row, and an +// infinite entry opens that side. A row with Lower = Upper is an +// equality constraint. +type LinearConstraints struct { + A *core.Array + Lower []float64 + Upper []float64 +} + +// rowPenalty is one row's augmented-Lagrangian state: the coefficient +// vector, the multipliers (High carries the equality row's signed +// multiplier; the two one-sided multipliers of an inequality row are +// both non-negative) and the violation measured at the point of the +// last objective or gradient call. +type rowPenalty struct { + coeffs []float64 + lower float64 + upper float64 + equal bool + high float64 + low float64 + violate float64 + side float64 // +1 above the upper wall, -1 below the lower one +} + +// The augmented-Lagrangian schedule both constraint entries walk: the +// linear rows of MinimiseConstrained and the nonlinear rows of +// MinimiseNonlinearConstrained share one penalty path, one ceiling and +// one feasibility gate, so the two cannot silently diverge. +const ( + almMuStart = 10.0 + almMuGrowth = 10.0 + almMuCeiling = 1e10 + almOuterRound = 40 + // almFeasibleAt is the absolute row violation the outer loop + // accepts as feasible. Like the other tolerances in the package it + // is absolute, deliberately independent of opts.Tolerance (which + // stays the inner solver's projected-gradient tolerance): tying + // the two made a caller who asked for a tighter inner solve get a + // hard failure instead of a tighter answer. + almFeasibleAt = 1e-10 +) + +// measure evaluates every row at p, records each violation and its +// side for the penalty terms, and returns the worst true violation. +func measure(p *core.Array, rows []rowPenalty) float64 { + worst := 0.0 + n := p.Len() + for i := range rows { + r := &rows[i] + ax := 0.0 + for j := range n { + ax += r.coeffs[j] * p.FloatAt(j) + } + v, side := 0.0, 0.0 + switch { + case ax > r.upper: + v, side = ax-r.upper, 1 + case ax < r.lower: + v, side = r.lower-ax, -1 + } + r.violate, r.side = v, side + if v > worst { + worst = v + } + } + return worst +} + +// slope is d(penalty row)/d(ax) at the measured point: the multiplier +// of the active side plus the penalty's derivative. Zero inside the +// row's bounds. +func (r *rowPenalty) slope(mu float64) float64 { + switch r.side { + case 0: + return 0 + case 1: + // An equality row violates to one side only, so its slope is + // the same whichever sign the violation carries. + return r.high + mu*r.violate + default: + if r.equal { + return r.high - mu*r.violate + } + return -(r.low + mu*r.violate) + } +} + +// term is the row's contribution to the augmented Lagrangian at the +// measured point: the multiplier times the signed violation plus the +// quadratic penalty. An equality row's multiplier is signed and rides +// the side; an inequality row's one-sided multipliers are +// non-negative. +func (r *rowPenalty) term(mu float64) float64 { + if r.side == 0 { + return 0 + } + mult := r.high + if !r.equal && r.side < 0 { + mult = r.low + } + if r.equal { + mult *= r.side + } + return mult*r.violate + 0.5*mu*r.violate*r.violate +} + +// MinimiseConstrained returns the point and value of a local minimum +// of f subject to the box walls carried by opts and the linear rows +// l ≤ A·x ≤ u. It is the only entry point whose matrix is read +// directly, so only here is a complex A refused beside the usual dtype +// check on the starting point. Each row is used in the caller's own +// units: the violation, the multipliers and the penalty inherit the +// row's scale, so a row whose coefficients sit many orders above +// another's is enforced to a correspondingly tighter absolute +// precision. Each outer round minimises f plus the rows' augmented +// Lagrangian terms, then ascends the row multipliers and scales the +// quadratic penalty. The inner solves are MinimiseLBFGS warm-started +// from the previous round's point; the supplied gradient, when given, +// is the gradient of f and the rows' contributions are chained onto it +// analytically, and the callback must return n real values. A start +// outside the feasible set is fine: the outer rounds pull it back, and +// failure to reach feasibility is an error naming the worst remaining +// violation. Feasibility is judged on the caller's rows against a +// fixed absolute threshold of 1e-10, independent of opts.Tolerance: +// the option stays the inner solver's projected-gradient tolerance, +// and asking for a tighter inner solve never turns into a hard failure +// of the outer loop. A badly scaled row (a·x = 1 with a of 1e9 or +// more) can put the step that reduces the augmented Lagrangian beyond +// the inner line search's reach: such a run is refused with the stall +// diagnostic, never returned as a converged answer. +func MinimiseConstrained(f func(*core.Array) (float64, error), grad func(*core.Array) (*core.Array, error), + x0 *core.Array, cons LinearConstraints, opts LBFGSOptions) (*core.Array, float64, error) { + const name = "MinimiseConstrained" + n := x0.Len() + if n == 0 { + return nil, 0, base.Errf("%s: the starting point must have at least one element", name) + } + if x0.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex starting points are not supported", name) + } + if cons.A == nil { + return nil, 0, base.Errf("%s: the constraint matrix is nil", name) + } + if err := requireReal(name, "constraint matrices", cons.A); err != nil { + return nil, 0, err + } + if cons.A.NDim() != 2 || cons.A.Shape()[1] != n { + return nil, 0, base.Errf("%s: the constraint matrix is %s, want r×%d", name, base.ShapeText(cons.A.Shape()), n) + } + rowCount := cons.A.Shape()[0] + if rowCount == 0 { + return nil, 0, base.Errf("%s: the constraint matrix has no rows", name) + } + if len(cons.Lower) != rowCount || len(cons.Upper) != rowCount { + return nil, 0, base.Errf("%s: the bounds hold %d and %d entries for %d rows", + name, len(cons.Lower), len(cons.Upper), rowCount) + } + rows := make([]rowPenalty, rowCount) + for i := range rows { + r := &rows[i] + r.coeffs = make([]float64, n) + for j := range n { + a := cons.A.FloatAt(i*n + j) + if math.IsNaN(a) || math.IsInf(a, 0) { + return nil, 0, base.Errf("%s: row %d carries a non-finite coefficient", name, i+1) + } + r.coeffs[j] = a + } + if math.IsNaN(cons.Lower[i]) || math.IsNaN(cons.Upper[i]) || cons.Lower[i] > cons.Upper[i] { + return nil, 0, base.Errf("%s: row %d has bounds [%g, %g]", name, i+1, cons.Lower[i], cons.Upper[i]) + } + r.lower, r.upper = cons.Lower[i], cons.Upper[i] + r.equal = r.lower == r.upper + } + + x := make([]float64, n) + for i := range n { + x[i] = x0.FloatAt(i) + } + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-8 + } + mu := almMuStart + // The inner solve's start point and the augmented gradient buffer: + // one array each for the whole run. MinimiseLBFGS reads the start + // once at entry and copies every gradient out at once, so both are + // fully rewritten before the next reader sees them. + start := core.New(core.Float, n) + gradBuf := core.New(core.Float, n) + + for range almOuterRound { + objective := func(p *core.Array) (float64, error) { + fv, err := f(p) + if err != nil { + return 0, err + } + measure(p, rows) + total := fv + for i := range rows { + total += rows[i].term(mu) + } + return total, nil + } + objectiveGrad := func(p *core.Array) (*core.Array, error) { + g, err := grad(p) + if err != nil { + return nil, err + } + // The payload is validated here, before anything reads it: + // RawFloats() is nil for a complex or non-float payload and + // as short as the callback's array, so slicing it first + // panics instead of reporting the contract breach. + if err := requireReal(name, "gradients", g); err != nil { + return nil, err + } + if g.Len() != n { + return nil, base.Errf("%s: the gradient callback returned %d elements for %d variables", + name, g.Len(), n) + } + measure(p, rows) + out := gradBuf + for j := range n { + out.RawFloats()[j] = g.FloatAt(j) + } + for i := range rows { + w := rows[i].slope(mu) + if w == 0 { + continue + } + for j := range n { + out.RawFloats()[j] += w * rows[i].coeffs[j] + } + } + return out, nil + } + innerGrad := objectiveGrad + if grad == nil { + innerGrad = nil // the inner solver differences the objective + } + + copy(start.RawFloats(), x) + // The inner solve is inexact by design: the outer loop judges + // the point by the rows' feasibility, so a run that spends its + // iteration budget without reaching the projected-gradient + // tolerance still carries the outer loop forward. + inner := opts + inner.AllowBudgetExit = true + pt, _, err := MinimiseLBFGS(objective, innerGrad, start, inner) + if err != nil { + return nil, 0, base.Errf("%s: %w", name, err) + } + for i := range n { + x[i] = pt.FloatAt(i) + } + + if measure(linalg.ArrayFromFloatsSafe(x, n), rows) <= almFeasibleAt { + // The returned value is f(x) itself, not the augmented + // Lagrangian the inner solver minimised: the penalty terms + // may be small but their multipliers are not, and the + // caller asked for the objective. + fv, ferr := f(linalg.ArrayFromFloatsSafe(x, n)) + if ferr != nil { + return nil, 0, base.Errf("%s: %w", name, ferr) + } + out, _ := packResult(x, fv) + return out, fv, nil + } + for i := range rows { + r := &rows[i] + switch { + case r.equal: + r.high += mu * r.side * r.violate + case r.side > 0: + r.high = max(0, r.high+mu*r.violate) + case r.side < 0: + r.low = max(0, r.low+mu*r.violate) + } + } + if mu < almMuCeiling { + mu *= almMuGrowth + } + } + worst := measure(linalg.ArrayFromFloatsSafe(x, n), rows) + return nil, 0, base.Errf("%s: %d rounds left the worst row violation at %g", name, almOuterRound, worst) +} diff --git a/optim/linear_test.go b/optim/linear_test.go new file mode 100644 index 0000000..b29f897 --- /dev/null +++ b/optim/linear_test.go @@ -0,0 +1,150 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// constrainedBowl returns the separable bowl shifted to the given +// centre, the standard workhorse for verifying constraint handling: +// the unconstrained answer is known and the constrained one is a +// projection of it that admits an analytic check. +func constrainedBowl(cx, cy float64) func(*core.Array) (float64, error) { + return func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-cx, p.FloatAt(1)-cy + return dx*dx + dy*dy, nil + } +} + +// TestConstrainedEquality drives the augmented Lagrangian through an +// equality hyperplane that cuts the bowl's minimum: the constrained +// optimum is the projection of the unconstrained one onto the plane. +func TestConstrainedEquality(t *testing.T) { + // min (x-1)² + (y-1)² subject to x + y = 2: the plane passes + // through the unconstrained minimum, so the answer is (1, 1) with + // value 0 and the multiplier converges to zero. + A, _ := core.FromFloats([]float64{1, 1}, 1, 2) + point, value, err := MinimiseConstrained(constrainedBowl(1, 1), nil, + mustLin(t, []float64{5, -3}), LinearConstraints{A: A, Lower: []float64{2}, Upper: []float64{2}}, + LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatalf("MinimiseConstrained: %v", err) + } + if math.Abs(point.FloatAt(0)-1) > 1e-4 || math.Abs(point.FloatAt(1)-1) > 1e-4 { + t.Fatalf("point = (%.8g, %.8g), want (1, 1)", point.FloatAt(0), point.FloatAt(1)) + } + if value > 1e-6 { + t.Fatalf("value = %.8g, want 0", value) + } +} + +// TestConstrainedInequality forces an active inequality: the bowl's +// minimum at (2, -1) lies beyond x + y = 0, so the constrained +// optimum sits on the wall at (1.5, -1.5) with value 0.5. +func TestConstrainedInequality(t *testing.T) { + A, _ := core.FromFloats([]float64{1, 1}, 1, 2) + gradFn := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, 2) + out.RawFloats()[0] = 2 * (p.FloatAt(0) - 2) + out.RawFloats()[1] = 2 * (p.FloatAt(1) + 1) + return out, nil + } + for _, grad := range []func(*core.Array) (*core.Array, error){nil, gradFn} { + point, value, err := MinimiseConstrained(constrainedBowl(2, -1), grad, + mustLin(t, []float64{4, 4}), LinearConstraints{A: A, Lower: []float64{math.Inf(-1)}, Upper: []float64{0}}, + LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatalf("MinimiseConstrained: %v", err) + } + if math.Abs(point.FloatAt(0)-1.5) > 1e-4 || math.Abs(point.FloatAt(1)+1.5) > 1e-4 { + t.Fatalf("point = (%.8g, %.8g), want (1.5, -1.5)", point.FloatAt(0), point.FloatAt(1)) + } + if math.Abs(value-0.5) > 1e-6 { + t.Fatalf("value = %.8g, want 0.5", value) + } + } +} + +// TestConstrainedTwoSided clamps the bowl between two parallel walls: +// 1 ≤ x ≤ 2 drags the first coordinate to the upper wall and leaves +// the second coordinate free. +func TestConstrainedTwoSided(t *testing.T) { + A, _ := core.FromFloats([]float64{1, 0}, 1, 2) + point, _, err := MinimiseConstrained(constrainedBowl(3, 1), nil, + mustLin(t, []float64{0, 0}), LinearConstraints{A: A, Lower: []float64{1}, Upper: []float64{2}}, + LBFGSOptions{}) + if err != nil { + t.Fatalf("MinimiseConstrained: %v", err) + } + if math.Abs(point.FloatAt(0)-2) > 1e-4 { + t.Fatalf("first coordinate = %.8g, want 2 on the upper wall", point.FloatAt(0)) + } + if math.Abs(point.FloatAt(1)-1) > 1e-4 { + t.Fatalf("second coordinate = %.8g, want 1", point.FloatAt(1)) + } +} + +// TestConstrainedWithBox composes linear rows with the box walls from +// the options: x + y = 0 pins the pair to a line and y ≥ -0.5 cuts +// the line at x = 0.5. The bowl around (1, 0) then bottoms out at +// (0.5, -0.5) with value 0.5, both terms contributing equally. +func TestConstrainedWithBox(t *testing.T) { + A, _ := core.FromFloats([]float64{1, 1}, 1, 2) + point, value, err := MinimiseConstrained(constrainedBowl(1, 0), nil, + mustLin(t, []float64{2, 2}), + LinearConstraints{A: A, Lower: []float64{0}, Upper: []float64{0}}, + LBFGSOptions{Tolerance: 1e-10, Lower: []float64{math.Inf(-1), -0.5}}) + if err != nil { + t.Fatalf("MinimiseConstrained: %v", err) + } + if math.Abs(point.FloatAt(0)-0.5) > 1e-4 || math.Abs(point.FloatAt(1)+0.5) > 1e-4 { + t.Fatalf("point = (%.8g, %.8g), want (0.5, -0.5)", point.FloatAt(0), point.FloatAt(1)) + } + if math.Abs(value-0.5) > 1e-5 { + t.Fatalf("value = %.8g, want 0.5", value) + } +} + +// TestConstrainedRefusals checks the loud rejections: nil matrix, +// wrong shape, bound or coefficient nonsense. +func TestConstrainedRefusals(t *testing.T) { + start := mustLin(t, []float64{0, 0}) + A, _ := core.FromFloats([]float64{1, 1}, 1, 2) + cases := []struct { + name string + cons LinearConstraints + }{ + {"nil matrix", LinearConstraints{Lower: []float64{0}, Upper: []float64{1}}}, + {"wrong shape", LinearConstraints{A: mustLin(t, []float64{1, 1}), Lower: []float64{0}, Upper: []float64{1}}}, + {"short bounds", LinearConstraints{A: A, Lower: []float64{0}, Upper: []float64{}}}, + {"crossed bounds", LinearConstraints{A: A, Lower: []float64{1}, Upper: []float64{0}}}, + } + for _, c := range cases { + if _, _, err := MinimiseConstrained(constrainedBowl(0, 0), nil, start, c.cons, LBFGSOptions{}); err == nil { + t.Fatalf("%s accepted", c.name) + } + } + nan, _ := core.FromFloats([]float64{math.NaN(), 1}, 1, 2) + if _, _, err := MinimiseConstrained(constrainedBowl(0, 0), nil, start, + LinearConstraints{A: nan, Lower: []float64{0}, Upper: []float64{1}}, LBFGSOptions{}); err == nil { + t.Fatal("NaN coefficient accepted") + } +} + +// mustLin builds a float vector for the constraint tests. +func mustLin(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 +} diff --git a/optim/linesearch_pins_test.go b/optim/linesearch_pins_test.go new file mode 100644 index 0000000..27c8f88 --- /dev/null +++ b/optim/linesearch_pins_test.go @@ -0,0 +1,650 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Line-search and convergence-exit pins for the optimisers. Each test +// names the defect it pins; the numeric fixtures are hand-derived +// optima. + +// stiffQuadratic is f = k·(x−3)², whose minimum is x = 3 with value 0. +func stiffQuadratic(k float64) func(*core.Array) (float64, error) { + return func(p *core.Array) (float64, error) { + d := p.FloatAt(0) - 3 + return k * d * d, nil + } +} + +// TestLBFGSLineSearchStiffQuadraticConverges pins the line-search +// budget: the step that reduces a quadratic is ≈ 1/L for a curvature +// L, so with 20 halvings from a unit step the first trial is never +// acceptable once L ≳ 1e6 and L-BFGS used to return the start point as +// a converged answer (x = 0, value 9k). +func TestLBFGSLineSearchStiffQuadraticConverges(t *testing.T) { + for _, k := range []float64{1e6, 1e8, 1e12} { + point, value, err := MinimiseLBFGS(stiffQuadratic(k), nil, mustFloats(t, []float64{0}), LBFGSOptions{}) + if err != nil { + t.Errorf("k=%g: MinimiseLBFGS: %v", k, err) + continue + } + if got := point.FloatAt(0); math.Abs(got-3) > 1e-6 { + t.Errorf("k=%g: x = %.12g, want 3", k, got) + } + if value > 1e-6 { + t.Errorf("k=%g: value = %.12g, want 0", k, value) + } + } +} + +// TestLBFGSLineSearchMixedUnitsConverges is the same failure through an +// ordinary two-parameter fit: the y coordinate's curvature sets the +// first step and takes x down with it when the step is capped at a +// unit. The optimum of (x−3)² + K·(y−5)² is (3, 5) with value 0. +func TestLBFGSLineSearchMixedUnitsConverges(t *testing.T) { + for _, k := range []float64{1e6, 1e8} { + f := func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-3, p.FloatAt(1)-5 + return dx*dx + k*dy*dy, nil + } + grad := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{2 * (p.FloatAt(0) - 3), 2 * k * (p.FloatAt(1) - 5)}, 2) + } + for _, g := range []struct { + name string + fn func(*core.Array) (*core.Array, error) + }{{"finite differences", nil}, {"analytic gradient", grad}} { + point, value, err := MinimiseLBFGS(f, g.fn, mustFloats(t, []float64{0, 0}, 2), LBFGSOptions{}) + if err != nil { + t.Errorf("K=%g (%s): %v", k, g.name, err) + continue + } + if math.Abs(point.FloatAt(0)-3) > 1e-6 || math.Abs(point.FloatAt(1)-5) > 1e-6 { + t.Errorf("K=%g (%s): point = (%.12g, %.12g), want (3, 5)", + k, g.name, point.FloatAt(0), point.FloatAt(1)) + } + if value > 1e-6 { + t.Errorf("K=%g (%s): value = %.12g, want 0", k, g.name, value) + } + } + } +} + +// TestLBFGSLineSearchStallIsAnError pins the silence: an exhausted +// line search used to break out of the iteration and hand the start +// point back as converged. The objective here is flat below zero and +// jumps at it, so the finite-difference stencil just below the jump +// reports a large gradient while no step along it can reduce the +// value: the search stalls and must say so. +func TestLBFGSLineSearchStallIsAnError(t *testing.T) { + f := func(p *core.Array) (float64, error) { + if p.FloatAt(0) >= 0 { + return 1, nil + } + return 0, nil + } + start := mustFloats(t, []float64{-1e-9}) + point, value, err := MinimiseLBFGS(f, nil, start, LBFGSOptions{}) + if err == nil { + t.Fatalf("a stalled line search was reported as convergence: point = %v, value = %g", + floatsOf(point), value) + } + if !strings.Contains(err.Error(), "line search") { + t.Fatalf("error = %v, want a line-search refusal", err) + } +} + +// TestLBFGSBoxOptimaFixtures checks the box-constrained answers against +// hand-derived optima: (x−3)²+(y+2)² on −1 ≤ x ≤ 1, y free reaches the +// upper wall at (1, −2) with value 4; (x+5)²+(y−5)² on 0 ≤ x, y ≤ 2 +// pins both coordinates at (0, 2) with value 34; (x+y−3)²+x² on +// x, y ≥ 1 bottoms out at the corner (1, 2) with value 1. +func TestLBFGSBoxOptimaFixtures(t *testing.T) { + inf := math.Inf(1) + cases := []struct { + name string + f func(*core.Array) (float64, error) + lower []float64 + upper []float64 + wantX []float64 + wantValue float64 + }{ + { + name: "upper wall, y free", + f: func(p *core.Array) (float64, error) { + return (p.FloatAt(0)-3)*(p.FloatAt(0)-3) + (p.FloatAt(1)+2)*(p.FloatAt(1)+2), nil + }, + lower: []float64{-1, -inf}, upper: []float64{1, inf}, + wantX: []float64{1, -2}, wantValue: 4, + }, + { + name: "both coordinates pinned", + f: func(p *core.Array) (float64, error) { + return (p.FloatAt(0)+5)*(p.FloatAt(0)+5) + (p.FloatAt(1)-5)*(p.FloatAt(1)-5), nil + }, + lower: []float64{0, 0}, upper: []float64{inf, 2}, + wantX: []float64{0, 2}, wantValue: 34, + }, + { + name: "coupled bowl at the corner", + f: func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return (x+y-3)*(x+y-3) + x*x, nil + }, + lower: []float64{1, 1}, upper: []float64{inf, inf}, + wantX: []float64{1, 2}, wantValue: 1, + }, + } + for _, c := range cases { + point, value, err := MinimiseLBFGS(c.f, nil, mustFloats(t, []float64{0, 0}, 2), + LBFGSOptions{Lower: c.lower, Upper: c.upper}) + if err != nil { + t.Errorf("%s: %v", c.name, err) + continue + } + for i, want := range c.wantX { + if math.Abs(point.FloatAt(i)-want) > 1e-4 { + t.Errorf("%s: x[%d] = %.12g, want %.12g", c.name, i, point.FloatAt(i), want) + } + } + if math.Abs(value-c.wantValue) > 1e-6 { + t.Errorf("%s: value = %.12g, want %.12g", c.name, value, c.wantValue) + } + } +} + +// TestLBFGSProjectionOntoBindingWall pins the box projection with a +// box that actually binds: (x−3)² on x ≤ 1 has its constrained minimum +// on the wall at x = 1 with value 4, and a start above the wall is +// projected onto it rather than refused. +func TestLBFGSProjectionOntoBindingWall(t *testing.T) { + f := func(p *core.Array) (float64, error) { + d := p.FloatAt(0) - 3 + return d * d, nil + } + for _, start := range []float64{0, 5, -10} { + point, value, err := MinimiseLBFGS(f, nil, mustFloats(t, []float64{start}), + LBFGSOptions{Upper: []float64{1}}) + if err != nil { + t.Fatalf("start %g: %v", start, err) + } + if math.Abs(point.FloatAt(0)-1) > 1e-6 { + t.Errorf("start %g: x = %.12g, want 1 on the wall", start, point.FloatAt(0)) + } + if math.Abs(value-4) > 1e-6 { + t.Errorf("start %g: value = %.12g, want 4", start, value) + } + } +} + +// TestMinimiseLevelSetStallIsNotConvergence pins the level-set trap: the +// value spread was the only convergence test, so a simplex whose +// vertices happened to lie on one level set stopped while spanning the +// space. +// For (x−1)² + (y−2)² from (0, 0) all three vertices of the stalled +// simplex sit on the circle of radius √0.5 about (1, 2), so the +// spread is zero and the answer used to be (1.5, 1.5) with value 0.5. +func TestMinimiseLevelSetStallIsNotConvergence(t *testing.T) { + f := func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-1, p.FloatAt(1)-2 + return dx*dx + dy*dy, nil + } + start := mustFloats(t, []float64{0, 0}, 2) + point, value, err := Minimise(f, start, MinimiseOptions{}) + if err != nil { + t.Fatalf("Minimise: %v", err) + } + if math.Abs(point.FloatAt(0)-1) > 1e-4 || math.Abs(point.FloatAt(1)-2) > 1e-4 { + t.Fatalf("point = (%v), want (1, 2)", floatsOf(point)) + } + if value > 1e-8 { + t.Fatalf("value = %g, want 0", value) + } + // Neither a larger budget nor a tighter tolerance may rescue the + // stall: the loop exits at the top-of-loop test either way. + for _, opts := range []MinimiseOptions{ + {MaxIterations: 20000}, + {Tolerance: 1e-20}, + } { + point, value, err := Minimise(f, start, opts) + if err != nil { + t.Fatalf("%+v: %v", opts, err) + } + if value > 1e-8 { + t.Fatalf("%+v: value = %g at (%v), want 0 at (1, 2)", opts, value, floatsOf(point)) + } + } +} + +// TestMinimiseHitRateOnShiftedBowls sweeps starts over a grid: every +// one of them must find the minimum of the axis-aligned bowl, the +// rotated bowl and the 3-D sphere. The level-set stall used to return +// a non-minimal point for 5 of 49 axis-bowl starts and 1 of 49 rotated +// starts with the default options. +func TestMinimiseHitRateOnShiftedBowls(t *testing.T) { + bowls := []struct { + name string + f func(*core.Array) (float64, error) + dim int + }{ + {"axis bowl", func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-1, p.FloatAt(1)-2 + return dx*dx + dy*dy, nil + }, 2}, + {"rotated bowl", func(p *core.Array) (float64, error) { + u := (p.FloatAt(0) - 1) + (p.FloatAt(1) - 2) + v := (p.FloatAt(0) - 1) - (p.FloatAt(1) - 2) + return u*u + 3*v*v, nil + }, 2}, + {"3-D sphere", func(p *core.Array) (float64, error) { + s := 0.0 + for i, c := range []float64{1, 2, 3} { + d := p.FloatAt(i) - c + s += d * d + } + return s, nil + }, 3}, + } + for _, b := range bowls { + for x := -3.0; x <= 3; x++ { + for y := -3.0; y <= 3; y++ { + start := []float64{x, y} + if b.dim == 3 { + start = []float64{x, y, x - y} + } + point, value, err := Minimise(b.f, mustFloats(t, start, b.dim), MinimiseOptions{}) + if err != nil { + t.Fatalf("%s from %v: %v", b.name, start, err) + } + if value > 1e-6 { + t.Errorf("%s from %v: value = %g at (%v), want 0", b.name, start, value, floatsOf(point)) + } + } + } + } +} + +// TestMinimiseToleranceScaleIsDocumented pins the documented absolute +// tolerance: an objective whose values are ~1e-14 already counts as +// flat (the default spread test is 1e-10·max(1, |f|)), so Minimise +// reports the start point and a nil error, and rescaling the objective +// to O(1), the documented remedy, resolves the minimum (1, 2). +func TestMinimiseToleranceScaleIsDocumented(t *testing.T) { + const scale = 1e-14 + tiny := func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-1, p.FloatAt(1)-2 + return scale * (dx*dx + dy*dy), nil + } + start := mustFloats(t, []float64{0, 0}, 2) + point, value, err := Minimise(tiny, start, MinimiseOptions{}) + if err != nil { + t.Fatalf("Minimise: %v", err) + } + // The start's own value is 5e-14; the run stops on the value spread + // long before the minimum is reached, which the documentation now + // warns about. + if math.Abs(point.FloatAt(0)-1) <= 0.1 || math.Abs(point.FloatAt(1)-2) <= 0.1 { + t.Fatalf("tiny objective: point = (%v), the documented absolute tolerance claims (1, 2) is left unfound", + floatsOf(point)) + } + if value > 5e-14 { + t.Fatalf("tiny objective: value = %g, want a value no larger than the start's 5e-14", value) + } + unit := func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-1, p.FloatAt(1)-2 + return dx*dx + dy*dy, nil + } + point, value, err = Minimise(unit, start, MinimiseOptions{}) + if err != nil { + t.Fatalf("Minimise rescaled: %v", err) + } + if math.Abs(point.FloatAt(0)-1) > 1e-4 || math.Abs(point.FloatAt(1)-2) > 1e-4 || value > 1e-8 { + t.Fatalf("rescaled objective: point = (%v), value = %g, want (1, 2) and 0", floatsOf(point), value) + } +} + +// TestMinimiseConstrainedComplexMatrixRefused pins the dtype guard on +// the constraint matrix: a complex A used to reach FloatAt and panic +// with an indexing error instead of the family's dtype refusal. +func TestMinimiseConstrainedComplexMatrixRefused(t *testing.T) { + A, err := core.FromComplexes([]complex128{1, 1}, 1, 2) + if err != nil { + t.Fatal(err) + } + start := mustFloats(t, []float64{5, -3}, 2) + if _, _, err := MinimiseConstrained(bowlAt(1, 2), nil, start, + LinearConstraints{A: A, Lower: []float64{3}, Upper: []float64{3}}, LBFGSOptions{}); err == nil { + t.Fatal("expected a complex constraint matrix to be refused") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-input refusal", err) + } +} + +// TestMinimiseConstrainedComplexGradientRefused pins the callback guard +// in the constrained wrapper: a complex gradient used to be sliced with +// RawFloats (nil for a complex payload) and panic. +func TestMinimiseConstrainedComplexGradientRefused(t *testing.T) { + grad := func(*core.Array) (*core.Array, error) { + g, err := core.FromComplexes([]complex128{1, 1}, 2) + return g, err + } + A, err := core.FromFloats([]float64{1, 1}, 1, 2) + if err != nil { + t.Fatal(err) + } + start := mustFloats(t, []float64{5, -3}, 2) + if _, _, err := MinimiseConstrained(bowlAt(1, 2), grad, start, + LinearConstraints{A: A, Lower: []float64{3}, Upper: []float64{3}}, LBFGSOptions{}); err == nil { + t.Fatal("expected a complex gradient to be refused") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-input refusal", err) + } +} + +// TestMinimiseConstrainedShortGradientRefused pins the length contract +// in the constrained wrapper: a gradient one element short used to +// panic in the slice before MinimiseLBFGS could report it. +func TestMinimiseConstrainedShortGradientRefused(t *testing.T) { + grad := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{1}), nil } + A, err := core.FromFloats([]float64{1, 1}, 1, 2) + if err != nil { + t.Fatal(err) + } + start := mustFloats(t, []float64{5, -3}, 2) + if _, _, err := MinimiseConstrained(bowlAt(1, 2), grad, start, + LinearConstraints{A: A, Lower: []float64{3}, Upper: []float64{3}}, LBFGSOptions{}); err == nil { + t.Fatal("expected a short gradient to be refused") + } else if !strings.Contains(err.Error(), "gradient callback") { + t.Fatalf("error = %v, want a callback-length refusal", err) + } +} + +// TestLBFGSComplexGradientRefused pins the dtype guard on the L-BFGS +// gradient callback, which used to dereference the nil int payload of a +// complex array. +func TestLBFGSComplexGradientRefused(t *testing.T) { + grad := func(*core.Array) (*core.Array, error) { + g, err := core.FromComplexes([]complex128{1, 1}, 2) + return g, err + } + if _, _, err := MinimiseLBFGS(bowlAt(1, 2), grad, mustFloats(t, []float64{5, -3}, 2), LBFGSOptions{}); err == nil { + t.Fatal("expected a complex gradient to be refused") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-input refusal", err) + } +} + +// TestLevenbergMarquardtComplexPayloadRefused pins both callback guards +// of the LM fitter: a complex residual and a complex analytic Jacobian +// used to panic in FloatAt. +func TestLevenbergMarquardtComplexPayloadRefused(t *testing.T) { + p0 := mustFloats(t, []float64{0, 0}, 2) + complexResidual := func(*core.Array) (*core.Array, error) { + r, err := core.FromComplexes([]complex128{1, 2, 3, 4}, 4) + return r, err + } + if _, _, err := LevenbergMarquardt(complexResidual, p0, LMOptions{}); err == nil { + t.Fatal("expected a complex residual to be refused") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("residual: error = %v, want a complex-input refusal", err) + } + residual := func(p *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{p.FloatAt(0) - 1, p.FloatAt(1) - 2}, 2) + } + complexJacobian := func(*core.Array) (*core.Array, error) { + j, err := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + return j, err + } + if _, _, err := LevenbergMarquardt(residual, p0, LMOptions{Jacobian: complexJacobian}); err == nil { + t.Fatal("expected a complex Jacobian to be refused") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("Jacobian: error = %v, want a complex-input refusal", err) + } +} + +// TestFindRootSystemComplexResidualRefused pins the dtype guard on the +// root-system residual, which used to dereference the nil int payload +// of a complex array. +func TestFindRootSystemComplexResidualRefused(t *testing.T) { + f := func(*core.Array) (*core.Array, error) { + r, err := core.FromComplexes([]complex128{1, 2}, 2) + return r, err + } + if _, _, err := FindRootSystem(f, mustFloats(t, []float64{1, 1}, 2), RootSystemOptions{}); err == nil { + t.Fatal("expected a complex residual to be refused") + } else if !strings.Contains(err.Error(), "complex") { + t.Fatalf("error = %v, want a complex-input refusal", err) + } +} + +// TestLevenbergMarquardtChi2MatchesReturnedPoint pins the reported fit +// quality: the relative-improvement break published the new, lower χ² +// while the parameters were still the old ones, so the answer looked +// 99.9 % better than the point that came back. +func TestLevenbergMarquardtChi2MatchesReturnedPoint(t *testing.T) { + xs := []float64{0, 1, 2, 3, 4, 5, 6, 7, 8, 9} + ys := make([]float64, len(xs)) + for i, x := range xs { + ys[i] = 3*math.Exp(-0.5*x) + 0.5 + 0.001*math.Sin(7*x) + } + residual := func(p *core.Array) (*core.Array, error) { + a, b, c := p.FloatAt(0), p.FloatAt(1), p.FloatAt(2) + out := core.New(core.Float, len(xs)) + for i := range xs { + out.RawFloats()[i] = ys[i] - (a*math.Exp(-b*xs[i]) + c) + } + return out, nil + } + point, chi2, err := LevenbergMarquardt(residual, mustFloats(t, []float64{2, 0.3, 0.1}, 3), + LMOptions{Tolerance: 0.05}) + if err != nil { + t.Fatalf("LevenbergMarquardt: %v", err) + } + r, err := residual(point) + if err != nil { + t.Fatal(err) + } + actual := 0.0 + for i := range r.Len() { + actual += r.FloatAt(i) * r.FloatAt(i) + } + if math.Abs(chi2-actual) > 1e-9*math.Max(1, actual) { + t.Fatalf("reported χ² = %.14g, χ² at the returned point = %.14g", chi2, actual) + } +} + +// TestMinimiseConstrainedExactFixtures checks the augmented Lagrangian +// against hand-derived optima: min (x−1)² + (y−2)² subject to x+y = 5 +// projects to (2, 3) with value 2; min (x−10)² + (y−10)² subject to +// x+y = 1 and x ≤ −4 bottoms out at (−4, 5) with value 221; and the +// degenerate row min x+3 subject to x = 1 reaches 4. +func TestMinimiseConstrainedExactFixtures(t *testing.T) { + bowlPeak, err := core.FromFloats([]float64{1, 1}, 1, 2) + if err != nil { + t.Fatal(err) + } + boxCut, err := core.FromFloats([]float64{1, 1}, 1, 2) + if err != nil { + t.Fatal(err) + } + cases := []struct { + name string + f func(*core.Array) (float64, error) + A *core.Array + lower []float64 + upper []float64 + boxUpper []float64 + start []float64 + wantX []float64 + wantValue float64 + }{ + { + name: "equality cuts the unconstrained minimum", + f: func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-1, p.FloatAt(1)-2 + return dx*dx + dy*dy, nil + }, + A: bowlPeak, + lower: []float64{5}, upper: []float64{5}, + start: []float64{0, 0}, wantX: []float64{2, 3}, wantValue: 2, + }, + { + name: "row cut by a box wall", + f: func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-10, p.FloatAt(1)-10 + return dx*dx + dy*dy, nil + }, + A: boxCut, + lower: []float64{1}, upper: []float64{1}, + boxUpper: []float64{-4, math.Inf(1)}, + start: []float64{0, 0}, wantX: []float64{-4, 5}, wantValue: 221, + }, + } + for _, c := range cases { + point, value, err := MinimiseConstrained(c.f, nil, mustFloats(t, c.start, len(c.start)), + LinearConstraints{A: c.A, Lower: c.lower, Upper: c.upper}, + LBFGSOptions{Tolerance: 1e-10, Upper: c.boxUpper}) + if err != nil { + t.Errorf("%s: %v", c.name, err) + continue + } + for i, want := range c.wantX { + if math.Abs(point.FloatAt(i)-want) > 1e-4 { + t.Errorf("%s: x[%d] = %.12g, want %.12g", c.name, i, point.FloatAt(i), want) + } + } + if math.Abs(value-c.wantValue) > 1e-4*math.Max(1, c.wantValue) { + t.Errorf("%s: value = %.12g, want %.12g", c.name, value, c.wantValue) + } + } +} + +// TestMinimiseConstrainedBadlyScaledRow pins the badly scaled row: min +// x² + y² subject to a·x + y = 1 has the closed form x = a/(a²+1), +// y = 1/(a²+1) with value 1/(a²+1). The row keeps the caller's scale, +// so the penalty's gradient at the start is mu·a and its curvature +// mu·a², which put the step that reduces the augmented Lagrangian +// below the old 20-halving line search for every a ≥ 1000: the inner +// solve took the silent stall exit and the outer loop reported "40 +// rounds left the worst row violation at 1" without moving. The rows +// the long line search can reach must now reach the closed form; rows +// beyond its reach must be refused with the stall diagnostic, never +// returned as a converged answer. +func TestMinimiseConstrainedBadlyScaledRow(t *testing.T) { + f := func(p *core.Array) (float64, error) { + return p.FloatAt(0)*p.FloatAt(0) + p.FloatAt(1)*p.FloatAt(1), nil + } + rowResidual := func(a float64, point *core.Array) float64 { + return math.Abs(a*point.FloatAt(0) + point.FloatAt(1) - 1) + } + for _, a := range []float64{1e3, 1e4, 1e6, 1e8} { + A, err := core.FromFloats([]float64{a, 1}, 1, 2) + if err != nil { + t.Fatal(err) + } + point, value, err := MinimiseConstrained(f, nil, mustFloats(t, []float64{0, 0}, 2), + LinearConstraints{A: A, Lower: []float64{1}, Upper: []float64{1}}, LBFGSOptions{}) + if err != nil { + t.Errorf("a=%g: %v", a, err) + continue + } + wantX, wantY, wantValue := a/(a*a+1), 1/(a*a+1), 1/(a*a+1) + if got := point.FloatAt(0); math.Abs(got-wantX) > 1e-6*wantX { + t.Errorf("a=%g: x = %.12g, want %.12g", a, got, wantX) + } + if got := point.FloatAt(1); math.Abs(got-wantY) > 1e-6*wantY { + t.Errorf("a=%g: y = %.12g, want %.12g", a, got, wantY) + } + if math.Abs(value-wantValue) > 1e-6*wantValue { + t.Errorf("a=%g: value = %.12g, want %.12g", a, value, wantValue) + } + if res := rowResidual(a, point); res > 1e-6 { + t.Errorf("a=%g: row residual = %g, want ≤ 1e-6", a, res) + } + } + // Past the line search's reach the answer is an error, not the start + // point dressed as convergence: a returned point must satisfy the + // row it claims to solve. + for _, a := range []float64{1e9, 1e12} { + A, err := core.FromFloats([]float64{a, 1}, 1, 2) + if err != nil { + t.Fatal(err) + } + point, _, err := MinimiseConstrained(f, nil, mustFloats(t, []float64{0, 0}, 2), + LinearConstraints{A: A, Lower: []float64{1}, Upper: []float64{1}}, LBFGSOptions{}) + if err != nil { + if !strings.Contains(err.Error(), "line search") { + t.Errorf("a=%g: error = %v, want the inner line search's stall diagnostic", a, err) + } + continue + } + if res := rowResidual(a, point); res > 1e-6 { + t.Errorf("a=%g: returned (%v) as converged with a row residual of %g, want the stall reported", + a, floatsOf(point), res) + } + } +} + +// TestMinimiseConstrainedToleranceIsNotTheFeasibilityThreshold pins the +// decoupling: opts.Tolerance is the inner solver's projected-gradient +// tolerance and used to double as the absolute row-feasibility +// threshold (feasibleAt = max(Tolerance, 1e-10)), so a looser inner +// solve bought a looser row. For min (x−3)² + y² subject to 1 ≤ x ≤ 2 +// (answer (2, 0), value 1) the old coupling returned (2.0033, 0) at +// Tolerance = 1e-2, a row violation of 3.3e-3, and a tolerance below +// 1e-10 hard-failed the solve with "40 rounds left the worst row +// violation at 1.3e-9". The row is now judged against the fixed 1e-10 +// at every inner tolerance. +func TestMinimiseConstrainedToleranceIsNotTheFeasibilityThreshold(t *testing.T) { + f := func(p *core.Array) (float64, error) { + dx := p.FloatAt(0) - 3 + return dx*dx + p.FloatAt(1)*p.FloatAt(1), nil + } + A, err := core.FromFloats([]float64{1, 0}, 1, 2) + if err != nil { + t.Fatal(err) + } + for _, tol := range []float64{0, 1e-8, 1e-10, 1e-12, 1e-4, 1e-2} { + point, value, err := MinimiseConstrained(f, nil, mustFloats(t, []float64{0, 0}, 2), + LinearConstraints{A: A, Lower: []float64{1}, Upper: []float64{2}}, + LBFGSOptions{Tolerance: tol}) + if err != nil { + t.Errorf("Tolerance=%g: %v", tol, err) + continue + } + if math.Abs(point.FloatAt(0)-2) > 1e-6 || math.Abs(point.FloatAt(1)) > 1e-6 { + t.Errorf("Tolerance=%g: point = (%v), want (2, 0)", tol, floatsOf(point)) + } + if violation := math.Abs(point.FloatAt(0) - 2); violation > 1e-8 { + t.Errorf("Tolerance=%g: row violation = %g, want ≤ 1e-8", tol, violation) + } + if math.Abs(value-1) > 1e-6 { + t.Errorf("Tolerance=%g: value = %.12g, want 1", tol, value) + } + } +} + +// bowlAt returns the separable bowl centred at (cx, cy), the fixture +// the constrained tests share. +func bowlAt(cx, cy float64) func(*core.Array) (float64, error) { + return func(p *core.Array) (float64, error) { + dx, dy := p.FloatAt(0)-cx, p.FloatAt(1)-cy + return dx*dx + dy*dy, nil + } +} + +// floatsOf copies an array's values out for readable failure messages. +func floatsOf(a *core.Array) []float64 { + out := make([]float64, a.Len()) + for i := range out { + out[i] = a.FloatAt(i) + } + return out +} diff --git a/optim/nonfinite_exit_pins_test.go b/optim/nonfinite_exit_pins_test.go new file mode 100644 index 0000000..788c88d --- /dev/null +++ b/optim/nonfinite_exit_pins_test.go @@ -0,0 +1,157 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestLBFGSRejectsNonFiniteObjectiveAndGradient pins that a +// non-finite objective value on the finite-difference path is an +// error: the NaN used to slip past the projected-gradient max and read +// as converged. +func TestLBFGSRejectsNonFiniteObjectiveAndGradient(t *testing.T) { + f := func(x *core.Array) (float64, error) { + if x.FloatAt(0) < 0 { + return math.NaN(), nil + } + return (x.FloatAt(0) - 3) * (x.FloatAt(0) - 3), nil + } + x0 := mustFloats(t, []float64{0}, 1) + if _, _, err := MinimiseLBFGS(f, nil, x0, LBFGSOptions{}); err == nil { + t.Fatal("expected an error for a NaN objective on the FD path") + } + // A NaN objective at the start is refused before any step. + f2 := func(*core.Array) (float64, error) { return math.Inf(1), nil } + if _, _, err := MinimiseLBFGS(f2, nil, x0, LBFGSOptions{}); err == nil { + t.Fatal("expected an error for an Inf objective at the start") + } +} + +// TestLBFGSConvergesOnTheLastIteration pins that a tolerance +// met exactly on the final permitted iteration is a success, not a +// budget error. +func TestLBFGSConvergesOnTheLastIteration(t *testing.T) { + f := func(x *core.Array) (float64, error) { + return (x.FloatAt(0) - 1.5) * (x.FloatAt(0) - 1.5), nil + } + g := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{2 * (x.FloatAt(0) - 1.5)}, 1) + } + x0 := mustFloats(t, []float64{0}, 1) + out, _, err := MinimiseLBFGS(f, g, x0, LBFGSOptions{MaxIterations: 1, Tolerance: 1e-10}) + if err != nil { + t.Fatalf("MinimiseLBFGS 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)) + } +} + +// TestFindRootSystemRejectsNonFiniteResidual pins that a NaN +// residual at an adopted state is an error: the infinity norm used to +// read it as a converged root with residual zero. +func TestFindRootSystemRejectsNonFiniteResidual(t *testing.T) { + f := func(*core.Array) (*core.Array, error) { + return mustFloats(t, []float64{math.NaN(), math.NaN()}), nil + } + x0 := mustFloats(t, []float64{1, 1}, 2) + _, _, err := FindRootSystem(f, x0, RootSystemOptions{}) + if err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("err = %v, want the non-finite refusal", err) + } + // One prefix only: the message carries the entry point once. + if strings.Count(err.Error(), "FindRootSystem") != 1 { + t.Fatalf("err = %v, want exactly one entry-point prefix", err) + } +} + +// TestFindRootSystemSurvivesOverflowingTrials pins the trial +// softening: log(x) is NaN on the negative half-line, the first +// Newton step from afar overshoots straight into it, and the damping +// halves past the rejected trial instead of dying. The start itself +// stays finite, so only the trials ever see NaN. +func TestFindRootSystemSurvivesOverflowingTrials(t *testing.T) { + f := func(x *core.Array) (*core.Array, error) { + return mustFloats(t, []float64{math.Log(x.FloatAt(0)) - 1}), nil + } + x0 := mustFloats(t, []float64{1000}, 1) + sol, _, err := FindRootSystem(f, x0, RootSystemOptions{}) + if err != nil { + t.Fatalf("FindRootSystem through NaN trials: %v", err) + } + if math.Abs(sol.FloatAt(0)-math.E) > 1e-8 { + t.Fatalf("root = %.12g, want e", sol.FloatAt(0)) + } +} + +// TestFindRootSystemVanishedStepIsNotARoot pins that a step +// collapsing to zero under a numerically zero Jacobian is a stall the +// iteration reports, never a root published beside a live residual. +func TestFindRootSystemVanishedStepIsNotARoot(t *testing.T) { + f := func(x *core.Array) (*core.Array, error) { + return mustFloats(t, []float64{1e-320*x.FloatAt(0)*x.FloatAt(0) - 10}), nil + } + x0 := mustFloats(t, []float64{0}, 1) + sol, res, err := FindRootSystem(f, x0, RootSystemOptions{}) + if err == nil { + t.Fatalf("a vanished step published root %g with residual %g", sol.FloatAt(0), res) + } +} + +// TestFindRootSystemConvergesOnTheLastIteration pins the +// exact-final-step convergence of the root solver. +func TestFindRootSystemConvergesOnTheLastIteration(t *testing.T) { + f := func(x *core.Array) (*core.Array, error) { + return mustFloats(t, []float64{x.FloatAt(0) - 2, x.FloatAt(1) + 1}), nil + } + x0 := mustFloats(t, []float64{0, 0}, 2) + out, _, err := FindRootSystem(f, x0, RootSystemOptions{MaxIterations: 1, Tolerance: 1e-10}) + if err != nil { + t.Fatalf("FindRootSystem on a linear system with one exact step: %v", err) + } + if out.FloatAt(0) != 2 || out.FloatAt(1) != -1 { + t.Fatalf("root = (%g, %g), want (2, -1)", out.FloatAt(0), out.FloatAt(1)) + } +} + +// TestLevenbergMarquardtPerfectStart pins that a start whose +// residual cancels exactly is the answer, not the damping collapse the +// strict improvement test used to die in. +func TestLevenbergMarquardtPerfectStart(t *testing.T) { + f := func(p *core.Array) (*core.Array, error) { + return mustFloats(t, []float64{p.FloatAt(0) - 2}), nil + } + p0 := mustFloats(t, []float64{2}, 1) + out, chi2, err := LevenbergMarquardt(f, p0, LMOptions{}) + if err != nil { + t.Fatalf("LevenbergMarquardt at the exact solution: %v", err) + } + if chi2 != 0 || out.FloatAt(0) != 2 { + t.Fatalf("fit = (%g, %g), want (2, 0)", out.FloatAt(0), chi2) + } +} + +// TestMinimiseConvergesOnTheLastIteration exercises the simplex +// re-test after the final move; the convergence here lands well inside +// the budget, so the case is a smoke of the returned pair, not a pin of +// the exact final move (a deterministic last-move budget is not stable +// across the portable and vector builds). +func TestMinimiseConvergesOnTheLastIteration(t *testing.T) { + f := func(x *core.Array) (float64, error) { + return math.Abs(x.FloatAt(0)-1) + 0.5*math.Abs(x.FloatAt(1)+1), nil + } + x0 := mustFloats(t, []float64{0, 0}, 2) + out, _, err := Minimise(f, x0, MinimiseOptions{MaxIterations: 200, Tolerance: 1e-8}) + if err != nil { + t.Fatalf("Minimise: %v", err) + } + if math.Abs(out.FloatAt(0)-1) > 1e-4 || math.Abs(out.FloatAt(1)+1) > 1e-4 { + t.Fatalf("minimum = (%g, %g), want (1, -1)", out.FloatAt(0), out.FloatAt(1)) + } +} diff --git a/optim/nonlinearconstr.go b/optim/nonlinearconstr.go new file mode 100644 index 0000000..f3a18dc --- /dev/null +++ b/optim/nonlinearconstr.go @@ -0,0 +1,313 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/linalg" +) + +// Nonlinear equality and inequality constraints through the augmented +// Lagrangian machinery the linear rows of linear.go already run on. +// +// How the nonlinear rows enter: each row carries the same rowPenalty +// state the linear rows carry (the multipliers, the violation and its +// side, and the equality flag), so the penalty term and the slope that +// chains onto the gradient are the shared functions +// rowPenalty.term and rowPenalty.slope unchanged. What differs is only +// where the row's value and gradient come from: a linear row reads +// A·x and a fixed coefficient vector, a nonlinear row calls the +// caller's function at the measured point and takes central +// differences of it for the chain rule. The outer loop itself is the +// same schedule as MinimiseConstrained's: the multipliers ascend by +// the violation, the quadratic penalty grows by the same factors, and +// feasibility is judged on the true rows against the same fixed +// 1e-10 threshold, independent of opts.Tolerance, which stays the +// inner solver's projected-gradient tolerance. The inner solves stay +// inexact by design (MinimiseLBFGS with AllowBudgetExit), judged by +// that outer feasibility check, exactly as the linear entry's document +// says. Mixed problems need nothing new: a linear row a·x ≤ u is a +// legal nonlinear row h(x) = a·x − u. + +// NonlinearConstraints carries functional rows: each equality is +// g(x) = 0 and each inequality h(x) ≤ 0, both as callbacks receiving +// the candidate point as a rank-1 array. The functions must return +// finite values; an error they return is fatal for the run. +type NonlinearConstraints struct { + Equalities []func(*core.Array) (float64, error) + Inequalities []func(*core.Array) (float64, error) +} + +// nonlinearRow is one functional row: the shared rowPenalty state (the +// equality rows sit at lower = upper = 0, the inequality rows at +// lower = −∞, upper = 0), the caller's function, the scratch for the +// function's gradient at the last measured point, and the function's +// own signed value there, which the inequality ascent rides so a slack +// row's multiplier decays instead of freezing. +type nonlinearRow struct { + rowPenalty + fn func(*core.Array) (float64, error) + grad []float64 + value float64 +} + +// MinimiseNonlinearConstrained returns the point, the value of f and +// the row multipliers of a local minimum of f subject to the nonlinear +// rows of cons. The multipliers hold the equality rows' signed +// estimates first, then the inequality rows' non-negative ones, in +// declaration order; they are the outer loop's final estimates, which +// converge to the KKT multipliers of the rows active at the answer, +// while an inactive inequality row's estimate decays to its KKT zero. +// A row with both slices empty is the unconstrained problem and +// delegates to MinimiseLBFGS with nil multipliers. +// +// The constants of the schedule (the initial penalty of 10, its growth +// of 10 per round, 40 outer rounds) and the fixed feasibility +// threshold of 1e-10 are the linear entry's; a run that spends its 40 +// rounds without reaching feasibility is refused with the worst +// remaining violation, and a badly scaled row can put the inner step +// beyond the line search's reach, which is refused with the stall +// diagnostic, never returned as a converged answer. +func MinimiseNonlinearConstrained(f func(*core.Array) (float64, error), + grad func(*core.Array) (*core.Array, error), + x0 *core.Array, cons NonlinearConstraints, opts LBFGSOptions) (*core.Array, float64, []float64, error) { + const name = "MinimiseNonlinearConstrained" + n := x0.Len() + if n == 0 { + return nil, 0, nil, base.Errf("%s: the starting point must have at least one element", name) + } + if x0.Dtype() == core.Complex { + return nil, 0, nil, base.Errf("%s: complex starting points are not supported", name) + } + for i, fn := range cons.Equalities { + if fn == nil { + return nil, 0, nil, base.Errf("%s: equality row %d is a nil function", name, i+1) + } + } + for i, fn := range cons.Inequalities { + if fn == nil { + return nil, 0, nil, base.Errf("%s: inequality row %d is a nil function", name, i+1) + } + } + if len(cons.Equalities) == 0 && len(cons.Inequalities) == 0 { + pt, fv, err := MinimiseLBFGS(f, grad, x0, opts) + return pt, fv, nil, err + } + rows := make([]nonlinearRow, 0, len(cons.Equalities)+len(cons.Inequalities)) + for _, fn := range cons.Equalities { + rows = append(rows, nonlinearRow{ + equal: true, + fn: fn, + grad: make([]float64, n), + }) + } + for _, fn := range cons.Inequalities { + rows = append(rows, nonlinearRow{ + lower: math.Inf(-1), + fn: fn, + grad: make([]float64, n), + }) + } + + x := make([]float64, n) + for i := range n { + x[i] = x0.FloatAt(i) + } + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-8 + } + mu := almMuStart + worst := 0.0 + // The chain-rule stencils, the point the caller's rows are measured + // at, and the working copy of a point handed to them: allocated once + // for the whole run. The stencil carries the offset on one + // coordinate at a time, restored as soon as that coordinate is + // differenced, so nothing is copied per coordinate and no slice per + // variable per row reaches the heap. + point := make([]float64, n) + xp := make([]float64, n) + xm := make([]float64, n) + // The inner solve's start point and the augmented gradient buffer: + // one array each for the whole run. MinimiseLBFGS reads the start + // once at entry and copies every gradient out at once, so both are + // fully rewritten before the next reader sees them. + start := core.New(core.Float, n) + gradBuf := core.New(core.Float, n) + measure := func(p []float64, wantGrad bool) (float64, error) { + worst := 0.0 + arr := linalg.ArrayFromFloatsSafe(p, n) + if wantGrad { + copy(xp, p) + copy(xm, p) + } + for i := range rows { + r := &rows[i] + v, err := r.fn(arr) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, base.Errf("%s: a row function returned the non-finite value %g", name, v) + } + r.value = v + vio, side := 0.0, 0.0 + switch { + case v > r.upper: + vio, side = v-r.upper, 1 + case v < r.lower: + vio, side = r.lower-v, -1 + } + r.violate, r.side = vio, side + worst = math.Max(worst, vio) + if wantGrad { + // Central differences of the row function: the chain + // rule needs the row's gradient at the measured point, + // the functional counterpart of a linear row's fixed + // coefficients. + for j := range n { + eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(p[j])) + xp[j] += eps + xm[j] -= eps + fp, err := r.fn(linalg.ArrayFromFloatsSafe(xp, n)) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + fm, err := r.fn(linalg.ArrayFromFloatsSafe(xm, n)) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + r.grad[j] = (fp - fm) / (2 * eps) + xp[j], xm[j] = p[j], p[j] + } + } + } + return worst, nil + } + + // readPoint copies a point into the shared working buffer: the rows + // are measured through it and it is fully overwritten every call. + readPoint := func(p *core.Array) []float64 { + for i := range n { + point[i] = p.FloatAt(i) + } + return point + } + + for range almOuterRound { + objective := func(p *core.Array) (float64, error) { + fv, err := f(p) + if err != nil { + return 0, err + } + if _, err := measure(readPoint(p), false); err != nil { + return 0, err + } + total := fv + for i := range rows { + total += rows[i].term(mu) + } + return total, nil + } + objectiveGrad := func(p *core.Array) (*core.Array, error) { + g, err := grad(p) + if err != nil { + return nil, err + } + if err := requireReal(name, "gradients", g); err != nil { + return nil, err + } + if g.Len() != n { + return nil, base.Errf("%s: the gradient callback returned %d elements for %d variables", + name, g.Len(), n) + } + if _, err := measure(readPoint(p), true); err != nil { + return nil, err + } + out := gradBuf + for j := range n { + out.RawFloats()[j] = g.FloatAt(j) + } + for i := range rows { + w := rows[i].slope(mu) + if w == 0 { + continue + } + for j := range n { + out.RawFloats()[j] += w * rows[i].grad[j] + } + } + return out, nil + } + innerGrad := objectiveGrad + if grad == nil { + innerGrad = nil // the inner solver differences the objective + } + + copy(start.RawFloats(), x) + // The inner solve is inexact by design, as in the linear + // entry: the outer feasibility check is what judges it. + inner := opts + inner.AllowBudgetExit = true + pt, _, err := MinimiseLBFGS(objective, innerGrad, start, inner) + if err != nil { + return nil, 0, nil, base.Errf("%s: %w", name, err) + } + for i := range n { + x[i] = pt.FloatAt(i) + } + + // Assigned, not declared: the function-level worst carries the + // figure the budget refusal reports. + var merr error + worst, merr = measure(x, false) + if merr != nil { + return nil, 0, nil, base.Errf("%s: %w", name, merr) + } + if worst <= almFeasibleAt { + // The answer carries f's own value, not the augmented + // Lagrangian's. The multipliers take the estimates the + // rounds converged to, with complementary slackness pinned + // at the answer itself: an inequality row that sits strictly + // slack has KKT multiplier exactly zero, while an active or + // equality row keeps its estimate. + for i := range rows { + r := &rows[i] + if !r.equal && r.value < -almFeasibleAt { + r.high = 0 + } + } + fv, ferr := f(linalg.ArrayFromFloatsSafe(x, n)) + if ferr != nil { + return nil, 0, nil, base.Errf("%s: %w", name, ferr) + } + multipliers := make([]float64, len(rows)) + for i := range rows { + multipliers[i] = rows[i].high + } + out, fv := packResult(x, fv) + return out, fv, multipliers, nil + } + for i := range rows { + r := &rows[i] + // The ascent mirrors the linear entry: an equality row's + // signed multiplier rides the side, an inequality row's + // non-negative one rides the row's own signed value, so a + // violated row lifts it exactly as before while a row that + // turned slack pulls it back toward its KKT zero instead + // of freezing at what the violated rounds left. + if r.equal { + r.high += mu * r.side * r.violate + } else { + r.high = max(0, r.high+mu*r.value) + } + } + if mu < almMuCeiling { + mu *= almMuGrowth + } + } + return nil, 0, nil, base.Errf("%s: %d rounds left the worst row violation at %g", name, almOuterRound, worst) +} diff --git a/optim/nonlinearconstr_test.go b/optim/nonlinearconstr_test.go new file mode 100644 index 0000000..26e4f9f --- /dev/null +++ b/optim/nonlinearconstr_test.go @@ -0,0 +1,306 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestNonlinearEqualityCircle pins the hand-solved circle case: min x +// subject to x² + y² = 1. The constrained minimum is (−1, 0) with +// value −1, and the KKT stationarity ∇f + λ∇g = 0 there reads +// (1, 0) + λ(−2, 0) = 0, so the multiplier converges to the analytic +// 0.5. The row's own gradient is differentiated, so the tolerance on +// the multiplier is the finite-difference one. +func TestNonlinearEqualityCircle(t *testing.T) { + cons := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return x*x + y*y - 1, nil + }, + }, + } + f := func(p *core.Array) (float64, error) { return p.FloatAt(0), nil } + x, value, multipliers, err := MinimiseNonlinearConstrained(f, nil, mustFloats(t, []float64{-2, 0.5}), cons, + LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatalf("MinimiseNonlinearConstrained: %v", err) + } + if math.Abs(x.FloatAt(0)+1) > 1e-3 || math.Abs(x.FloatAt(1)) > 1e-3 { + t.Fatalf("point = (%.10g, %.10g), want (−1, 0)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(value+1) > 1e-4 { + t.Fatalf("value = %.10g, want −1", value) + } + if len(multipliers) != 1 { + t.Fatalf("multipliers = %v, want one entry for the equality row", multipliers) + } + if math.Abs(multipliers[0]-0.5) > 1e-3 { + t.Fatalf("multiplier = %.10g, want the analytic 0.5", multipliers[0]) + } +} + +// TestNonlinearInequalityCircle pins the active inequality: min +// −(x + y) subject to x² + y² ≤ 1. The unconstrained minimum runs to +// infinity, the constrained one sits on the circle at (1/√2, 1/√2) +// with value −√2, and the stationarity (−1, −1) + μ(√2, √2) = 0 fixes +// the multiplier at 1/√2. +func TestNonlinearInequalityCircle(t *testing.T) { + cons := NonlinearConstraints{ + Inequalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return x*x + y*y - 1, nil + }, + }, + } + f := func(p *core.Array) (float64, error) { return -(p.FloatAt(0) + p.FloatAt(1)), nil } + x, value, multipliers, err := MinimiseNonlinearConstrained(f, nil, mustFloats(t, []float64{0.5, 0.5}), cons, + LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatalf("MinimiseNonlinearConstrained: %v", err) + } + root := 1 / math.Sqrt2 + if math.Abs(x.FloatAt(0)-root) > 1e-3 || math.Abs(x.FloatAt(1)-root) > 1e-3 { + t.Fatalf("point = (%.10g, %.10g), want (%g, %g)", x.FloatAt(0), x.FloatAt(1), root, root) + } + if math.Abs(value+math.Sqrt2) > 1e-4 { + t.Fatalf("value = %.10g, want −√2", value) + } + if len(multipliers) != 1 || math.Abs(multipliers[0]-math.Sqrt2/2) > 5e-3 { + t.Fatalf("multipliers = %v, want [1/√2] within the finite-difference tolerance", multipliers) + } + if multipliers[0] < 0 { + t.Fatalf("multiplier = %.10g, non-negative on an active inequality", multipliers[0]) + } +} + +// TestNonlinearInequalitySlackMultiplier pins the complementary +// slackness the returned slice carries: a row that sits strictly slack +// at the answer has KKT multiplier exactly zero, not whatever the +// rounds its row was violated in left behind. The quartic's minimum +// over x ≤ 0.5 from x₀ = 3 is the unconstrained x = −2, the row deep +// in its slack. +func TestNonlinearInequalitySlackMultiplier(t *testing.T) { + f := func(p *core.Array) (float64, error) { + v := p.FloatAt(0) + return 100 * (v + 2) * (v + 2) * (v - 3) * (v - 3), nil + } + cons := NonlinearConstraints{ + Inequalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { return p.FloatAt(0) - 0.5, nil }, + }, + } + x, _, multipliers, err := MinimiseNonlinearConstrained(f, nil, mustFloats(t, []float64{3}), cons, + LBFGSOptions{}) + if err != nil { + t.Fatalf("MinimiseNonlinearConstrained: %v", err) + } + if math.Abs(x.FloatAt(0)+2) > 1e-3 { + t.Fatalf("x = %.10g, want the unconstrained −2", x.FloatAt(0)) + } + if len(multipliers) != 1 || multipliers[0] != 0 { + t.Fatalf("multipliers = %v, want the slack row's KKT zero", multipliers) + } +} + +// TestNonlinearLinearRowComposition composes an affine row as a +// function: min (x−1)² + (y−1)² subject to x + y − 2 = 0. The plane +// passes through the unconstrained minimum, so the answer is (1, 1) +// with value 0 and the multiplier converging to zero, the linear +// entry's TestConstrainedEquality through the nonlinear door. +func TestNonlinearLinearRowComposition(t *testing.T) { + cons := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { + return p.FloatAt(0) + p.FloatAt(1) - 2, nil + }, + }, + } + f := constrainedBowl(1, 1) + // The bowl's gradient, supplied in one of the two runs so both the + // chained-gradient path and the slope-zero skip are exercised. + grad := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, 2) + out.RawFloats()[0] = 2 * (p.FloatAt(0) - 1) + out.RawFloats()[1] = 2 * (p.FloatAt(1) - 1) + return out, nil + } + for _, grad := range []func(*core.Array) (*core.Array, error){nil, grad} { + x, value, multipliers, err := MinimiseNonlinearConstrained(f, grad, mustFloats(t, []float64{5, -3}), cons, + LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatalf("MinimiseNonlinearConstrained: %v", err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-4 || math.Abs(x.FloatAt(1)-1) > 1e-4 { + t.Fatalf("point = (%.8g, %.8g), want (1, 1)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(value) > 1e-6 { + t.Fatalf("value = %.8g, want 0", value) + } + if len(multipliers) != 1 || math.Abs(multipliers[0]) > 1e-3 { + t.Fatalf("multipliers = %v, want [≈0]", multipliers) + } + } +} + +// TestNonlinearStalledInnerRefused pins the honest refusal of a +// stalled inner solve: min 10⁶x² + y² subject to xy = 2 with one +// L-BFGS iteration per round crawls toward the hyperbola, the outer +// feasibility check stays unsatisfied through the 40 rounds, and the +// run is refused with the remaining violation, never returned as a +// solution. +func TestNonlinearStalledInnerRefused(t *testing.T) { + cons := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { + return p.FloatAt(0)*p.FloatAt(1) - 2, nil + }, + }, + } + f := func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return 1e6*x*x + y*y, nil + } + _, _, _, err := MinimiseNonlinearConstrained(f, nil, mustFloats(t, []float64{5, 5}), cons, + LBFGSOptions{MaxIterations: 1}) + if err == nil { + t.Fatal("a stalled inner solve was returned as a solution") + } + if !strings.Contains(err.Error(), "violation") { + t.Fatalf("error = %v, want the outer feasibility refusal", err) + } +} + +// TestNonlinearDelegatesUnconstrained pins the empty-set composition: +// no rows at all is MinimiseLBFGS with nil multipliers. +func TestNonlinearDelegatesUnconstrained(t *testing.T) { + x, value, multipliers, err := MinimiseNonlinearConstrained(constrainedBowl(2, -1), nil, + mustFloats(t, []float64{0, 0}), NonlinearConstraints{}, LBFGSOptions{}) + if err != nil { + t.Fatalf("MinimiseNonlinearConstrained: %v", err) + } + if math.Abs(x.FloatAt(0)-2) > 1e-4 || math.Abs(x.FloatAt(1)+1) > 1e-4 { + t.Fatalf("point = (%.8g, %.8g), want (2, −1)", x.FloatAt(0), x.FloatAt(1)) + } + if value > 1e-8 { + t.Fatalf("value = %.8g, want ≈ 0", value) + } + if multipliers != nil { + t.Fatalf("multipliers = %v, want nil", multipliers) + } +} + +// TestNonlinearWithAnalyticGradient runs the same hand-solved circle +// case with f's gradient supplied: the row terms are then chained onto +// it analytically, the row's own gradient differentiated at each +// measured point, and the answer must match the finite-difference run. +func TestNonlinearWithAnalyticGradient(t *testing.T) { + cons := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return x*x + y*y - 1, nil + }, + }, + } + f := func(p *core.Array) (float64, error) { return p.FloatAt(0), nil } + grad := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, 2) + out.RawFloats()[0] = 1 + return out, nil + } + x, value, multipliers, err := MinimiseNonlinearConstrained(f, grad, mustFloats(t, []float64{-2, 0.5}), cons, + LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatalf("MinimiseNonlinearConstrained: %v", err) + } + if math.Abs(x.FloatAt(0)+1) > 1e-3 || math.Abs(x.FloatAt(1)) > 1e-3 { + t.Fatalf("point = (%.10g, %.10g), want (−1, 0)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(value+1) > 1e-4 { + t.Fatalf("value = %.10g, want −1", value) + } + if len(multipliers) != 1 || math.Abs(multipliers[0]-0.5) > 1e-3 { + t.Fatalf("multipliers = %v, want [0.5]", multipliers) + } + // A gradient of the wrong length is refused. + if _, _, _, err := MinimiseNonlinearConstrained(f, func(*core.Array) (*core.Array, error) { + return core.New(core.Float, 3), nil + }, mustFloats(t, []float64{-2, 0.5}), cons, LBFGSOptions{}); err == nil { + t.Fatal("a wrong-length gradient was accepted") + } + // A complex gradient payload is refused. + if _, _, _, err := MinimiseNonlinearConstrained(f, func(*core.Array) (*core.Array, error) { + return mustComplexPoint(t), nil + }, mustFloats(t, []float64{-2, 0.5}), cons, LBFGSOptions{}); err == nil { + t.Fatal("a complex gradient was accepted") + } + // A gradient callback's own error propagates. + if _, _, _, err := MinimiseNonlinearConstrained(f, func(*core.Array) (*core.Array, error) { + return nil, base.Errf("the gradient exploded") + }, mustFloats(t, []float64{-2, 0.5}), cons, LBFGSOptions{}); err == nil { + t.Fatal("the gradient's error did not propagate") + } + // The objective's own error propagates out of the inner solve. + if _, _, _, err := MinimiseNonlinearConstrained(func(*core.Array) (float64, error) { + return 0, base.Errf("the objective exploded") + }, nil, mustFloats(t, []float64{-2, 0.5}), cons, LBFGSOptions{}); err == nil { + t.Fatal("the objective's error did not propagate") + } +} + +// TestNonlinearRefusals checks the loud rejections: nil functions, +// empty and complex starts, a non-finite row value and an objective +// error. +func TestNonlinearRefusals(t *testing.T) { + good := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){ + func(p *core.Array) (float64, error) { return p.FloatAt(0) - 1, nil }, + }, + } + f := constrainedBowl(0, 0) + start := mustFloats(t, []float64{0, 0}) + // Nil equality function. + nilEq := NonlinearConstraints{Equalities: []func(*core.Array) (float64, error){nil}} + if _, _, _, err := MinimiseNonlinearConstrained(f, nil, start, nilEq, LBFGSOptions{}); err == nil { + t.Fatal("a nil equality function was accepted") + } + nilIn := NonlinearConstraints{Inequalities: []func(*core.Array) (float64, error){nil}} + if _, _, _, err := MinimiseNonlinearConstrained(f, nil, start, nilIn, LBFGSOptions{}); err == nil { + t.Fatal("a nil inequality function was accepted") + } + // Empty start. + if _, _, _, err := MinimiseNonlinearConstrained(f, nil, core.New(core.Float, 0), good, LBFGSOptions{}); err == nil { + t.Fatal("an empty starting point was accepted") + } + // Complex start. + if _, _, _, err := MinimiseNonlinearConstrained(f, nil, mustComplexPoint(t), good, LBFGSOptions{}); err == nil { + t.Fatal("a complex starting point was accepted") + } + // A row function that returns NaN is fatal. + nan := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){ + func(*core.Array) (float64, error) { return math.NaN(), nil }, + }, + } + if _, _, _, err := MinimiseNonlinearConstrained(f, nil, start, nan, LBFGSOptions{}); err == nil { + t.Fatal("a NaN row value was accepted") + } + // A row function's own error propagates. + failing := NonlinearConstraints{ + Inequalities: []func(*core.Array) (float64, error){ + func(*core.Array) (float64, error) { return 0, base.Errf("the row exploded") }, + }, + } + if _, _, _, err := MinimiseNonlinearConstrained(f, nil, start, failing, LBFGSOptions{}); err == nil { + t.Fatal("the row function's error did not propagate") + } +} diff --git a/optim/optimise.go b/optim/optimise.go new file mode 100644 index 0000000..87968db --- /dev/null +++ b/optim/optimise.go @@ -0,0 +1,412 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "fmt" + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Scalar root finding and multivariate minimisation. Two root finders +// cover the common cases: FindRoot needs only a sign-changing bracket +// and converges unconditionally; FindRootNewton needs the derivative +// and converges quadratically when a good starting guess and a smooth +// derivative are available. Minimise is the derivative-free simplex +// method, the standard choice for objectives that are noisy, opaque or +// expensive to differentiate. + +// FindRoot returns a root of f in the bracket [a, b] by Brent's +// method, which combines inverse quadratic interpolation, the secant +// step and bisection. f(a) and f(b) must be finite with opposite +// signs, so a root is guaranteed inside. A tol ≤ 0 defaults to 1e-12. +func FindRoot(f func(float64) float64, a, b, tol float64) (float64, error) { + if tol <= 0 { + tol = 1e-12 + } + fa, fb := f(a), f(b) + if math.IsNaN(fa) || math.IsNaN(fb) || math.IsInf(fa, 0) || math.IsInf(fb, 0) { + return 0, base.Errf("FindRoot: the bracket must evaluate to finite values, got f(%g)=%g, f(%g)=%g", a, fa, b, fb) + } + if fa == 0 { + return a, nil + } + if fb == 0 { + return b, nil + } + if fa*fb > 0 { + return 0, base.Errf("FindRoot: the bracket [%g, %g] does not change sign (f(a)=%g, f(b)=%g)", a, b, fa, fb) + } + // Brent's iteration (fbrent): c tracks the opposite-sign end, e + // the previous step width; interpolation is tried first and + // bisection keeps the step honest. + c, fc := a, fa + d, e := b-a, b-a + for range 200 { + if fb*fc > 0 { + c, fc = a, fa + d, e = b-a, b-a + } + if math.Abs(fc) < math.Abs(fb) { + a, b, c = b, c, b + fa, fb, fc = fb, fc, fb + } + tol1 := 2*base.EpsF*math.Abs(b) + 0.5*tol + xm := 0.5 * (c - b) + if math.Abs(xm) <= tol1 || fb == 0 { + return b, nil + } + if math.Abs(e) >= tol1 && math.Abs(fa) > math.Abs(fb) { + s := fb / fa + var p, q float64 + if a == c { + // Secant. + p = 2 * xm * s + q = 1 - s + } else { + // Inverse quadratic interpolation. + q = fa / fc + r := fb / fc + p = s * (2*xm*q*(q-r) - (b-a)*(r-1)) + q = (q - 1) * (r - 1) * (s - 1) + } + if p > 0 { + q = -q + } + p = math.Abs(p) + if 2*p < min(3*xm*q-math.Abs(tol1*q), math.Abs(e*q)) { + e, d = d, p/q + } else { + d, e = xm, xm + } + } else { + d, e = xm, xm + } + a, fa = b, fb + if math.Abs(d) > tol1 { + b += d + } else { + b += tol1 * signOf(xm) + } + fb = f(b) + if fb == 0 { + return b, nil + } + } + return 0, base.Errf("FindRoot: no convergence in 200 iterations") +} + +// FindRootNewton returns a root of f near x0 by Newton's iteration +// with the supplied derivative df. A tol ≤ 0 defaults to 1e-12 and +// maxIter ≤ 0 to 100. A vanishing derivative or an exhausted budget +// is an error, not a silent guess. +func FindRootNewton(f, df func(float64) float64, x0, tol float64, maxIter int) (float64, error) { + if tol <= 0 { + tol = 1e-12 + } + if maxIter <= 0 { + maxIter = 100 + } + x := x0 + for range maxIter { + fx := f(x) + if math.IsNaN(fx) || math.IsInf(fx, 0) { + return 0, base.Errf("FindRootNewton: the objective left the real numbers at x=%g", x) + } + d := df(x) + if d == 0 { + return 0, base.Errf("FindRootNewton: the derivative vanishes at x=%g", x) + } + step := fx / d + x -= step + if math.Abs(step) <= tol*(1+math.Abs(x)) { + return x, nil + } + } + return 0, base.Errf("FindRootNewton: no convergence in %d steps from x0=%g", maxIter, x0) +} + +// MinimiseOptions tunes the simplex minimisation. MaxIterations ≤ 0 +// means 2000, Tolerance ≤ 0 means 1e-10, InitialStep ≤ 0 means 1. +// +// Both convergence tests are absolute in the objective's own scale: +// the value spread is compared against Tolerance·max(1, |f|) and the +// simplex diameter against Tolerance·max(1, |x|). An objective whose +// values sit many orders of magnitude below one therefore counts as +// flat from the start, and its start point is returned with a nil +// error; rescale the objective (and the variables) to O(1) before +// calling Minimise when the natural units are not of that size. +type MinimiseOptions struct { + MaxIterations int + Tolerance float64 + InitialStep float64 + // AllowBudgetExit makes a run that exhausts MaxIterations report + // its best point instead of an error. The default is false, so a + // budget stop is never mistaken for a converged answer; the flag + // mirrors LBFGSOptions.AllowBudgetExit. + AllowBudgetExit bool +} + +// Minimise returns the point and value of a local minimum of f near +// x0 by the Nelder-Mead simplex method, which needs no derivatives. +// The answer is a local minimum: multistart from several x0 when the +// objective may have several basins. +func Minimise(f func(*core.Array) (float64, error), x0 *core.Array, opts MinimiseOptions) (*core.Array, float64, error) { + if x0.Dtype() == core.Complex { + return nil, 0, base.Errf("Minimise: complex starting points are not supported") + } + n := x0.Len() + if n == 0 { + return nil, 0, base.Errf("Minimise: the starting point must have at least one element") + } + if opts.MaxIterations <= 0 { + opts.MaxIterations = 2000 + } + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-10 + } + if opts.InitialStep <= 0 { + opts.InitialStep = 1 + } + + eval := func(v []float64) (float64, error) { + a, err := core.FromFloats(v, n) + if err != nil { + return 0, err + } + fv, ferr := f(a) + if ferr != nil { + return 0, ferr + } + // A non-finite objective is an error naming the point: the + // spread test compares false against NaN, so a NaN vertex + // would burn the whole budget and be reported as a budget + // problem, and with AllowBudgetExit set it can sit at index 0 + // and come back as the answer. The error stays bare: every + // call site of eval wraps it once with "Minimise: %w" through + // base.Errf, which adds the entry-point name and the library + // tag, so the surfaced message carries exactly one prefix + // instead of the doubled "tensor: Minimise: tensor: + // Minimise:" a prefixed inner error produced. + if math.IsNaN(fv) || math.IsInf(fv, 0) { + return 0, fmt.Errorf("f returned the non-finite value %g at %v", fv, v) + } + return fv, nil + } + + // The simplex: n+1 vertices, x0 plus one offset per coordinate. + simplex := make([][]float64, n+1) + values := make([]float64, n+1) + simplex[0] = cloneDense(x0) + for i := range n { + v := cloneDense(x0) + step := opts.InitialStep * math.Max(1, math.Abs(v[i])) + v[i] += step + simplex[i+1] = v + } + for i, v := range simplex { + fv, err := eval(v) + if err != nil { + return nil, 0, base.Errf("Minimise: %w", err) + } + values[i] = fv + } + + // Scratch reused across iterations: the centroid and the reflected, + // expanded and contracted candidates. A candidate that wins is + // copied into the worst vertex's slot, whose slice then keeps its + // identity; the scratch is fully rewritten before every read. + centre := make([]float64, n) + reflected := make([]float64, n) + expanded := make([]float64, n) + contracted := make([]float64, n) + + converged := false + for range opts.MaxIterations { + // Order so vertex 0 is the best and vertex n the worst. + orderSimplex(simplex, values) + spread := values[n] - values[0] + if spread <= opts.Tolerance*math.Max(1, math.Abs(values[0])) { + // A small value spread on its own is not convergence: the + // vertices can agree on the value while spanning the + // parameter space, because all of them sit on one level set + // of the objective. The spatial test catches exactly that + // simplex, whose best vertex is no minimum at all. + if simplexDiameter(simplex) <= opts.Tolerance*math.Max(1, maxAbs(simplex[0])) { + converged = true + break + } + // A stalled level set is collapsed onto its best vertex: + // each shrink halves the diameter, so the spatial test is + // reached in a bounded number of rounds and the returned + // point is the best one the simplex found, never a random + // vertex of an equally-valued set. + if err := shrinkSimplex(simplex, values, eval); err != nil { + return nil, 0, base.Errf("Minimise: %w", err) + } + continue + } + worst := simplex[n] + clear(centre) + for i := range n { + for j := range n { + centre[j] += simplex[i][j] / float64(n) + } + } + + // Reflect the worst vertex through the centroid. + for j := range n { + reflected[j] = centre[j] + (centre[j] - worst[j]) + } + fr, err := eval(reflected) + if err != nil { + return nil, 0, base.Errf("Minimise: %w", err) + } + switch { + case fr < values[0]: + // Reflected better than the best: try expanding further. + for j := range n { + expanded[j] = centre[j] + 2*(centre[j]-worst[j]) + } + fe, err := eval(expanded) + if err != nil { + return nil, 0, base.Errf("Minimise: %w", err) + } + if fe < fr { + copy(worst, expanded) + values[n] = fe + } else { + copy(worst, reflected) + values[n] = fr + } + case fr < values[n-1]: + copy(worst, reflected) + values[n] = fr + default: + // Reflected worse than the second worst: contract. + for j := range n { + contracted[j] = centre[j] + 0.5*(worst[j]-centre[j]) + } + fc, err := eval(contracted) + if err != nil { + return nil, 0, base.Errf("Minimise: %w", err) + } + if fc < values[n] { + copy(worst, contracted) + values[n] = fc + break + } + // Shrink everything towards the best vertex. + if err := shrinkSimplex(simplex, values, eval); err != nil { + return nil, 0, base.Errf("Minimise: %w", err) + } + } + } + // Falling out of the loop means the budget ran out, not that a + // minimum was found: the simplex still moves, and reporting its + // best vertex as the answer is the silent wrongness the level-set + // and stall exits refuse. The final move happened after the top-of- + // loop test, so the convergence pair is re-checked once first. + orderSimplex(simplex, values) + if values[n]-values[0] <= opts.Tolerance*math.Max(1, math.Abs(values[0])) && + simplexDiameter(simplex) <= opts.Tolerance*math.Max(1, maxAbs(simplex[0])) { + converged = true + } + if !converged && !opts.AllowBudgetExit { + return nil, 0, base.Errf("Minimise: the iteration budget of %d ran out without the simplex converging", opts.MaxIterations) + } + // values[0] is f(simplex[0]) by construction: orderSimplex keeps the + // pairs together and every update writes both, so the answer is the + // held value, not a fresh evaluation of the best vertex. + fv := values[0] + out, err := core.FromFloats(simplex[0], n) + if err != nil { + return nil, 0, err + } + return out, fv, nil +} + +// orderSimplex sorts the vertices together with their values by +// ascending value. +func orderSimplex(simplex [][]float64, values []float64) { + for i := 1; i < len(values); i++ { + for j := i; j > 0 && values[j] < values[j-1]; j-- { + simplex[j], simplex[j-1] = simplex[j-1], simplex[j] + values[j], values[j-1] = values[j-1], values[j] + } + } +} + +// simplexDiameter returns the largest L∞ distance from the simplex's +// best vertex to another one: the spatial extent of the simplex, which +// the convergence test pairs with the value spread so that a set of +// vertices lying on one level set of the objective is never read as a +// converged answer. +func simplexDiameter(simplex [][]float64) float64 { + n := len(simplex[0]) + diameter := 0.0 + for i := 1; i < len(simplex); i++ { + for j := range n { + if d := math.Abs(simplex[i][j] - simplex[0][j]); d > diameter { + diameter = d + } + } + } + return diameter +} + +// shrinkSimplex halves every vertex's distance to the best one and +// re-evaluates the moved vertices. It is both the classic Nelder-Mead +// shrink, taken when a contraction failed, and the remedy for a +// level-set stall, where the vertices agree on the value without +// spanning a small neighbourhood. Each call halves the diameter, so a +// stalled simplex collapses onto its best point in a bounded number of +// rounds. +func shrinkSimplex(simplex [][]float64, values []float64, eval func([]float64) (float64, error)) error { + n := len(simplex) - 1 + for i := 1; i <= n; i++ { + for j := range n { + simplex[i][j] = simplex[0][j] + 0.5*(simplex[i][j]-simplex[0][j]) + } + fv, err := eval(simplex[i]) + if err != nil { + return err + } + values[i] = fv + } + return nil +} + +// signOf returns ±1 for the sign of x (0 counts as +). +func signOf(x float64) float64 { + if x < 0 { + return -1 + } + return 1 +} + +// cloneDense copies an array's elements into a plain float64 slice. +func cloneDense(a *core.Array) []float64 { + vals := make([]float64, a.Len()) + for i := range vals { + vals[i] = a.FloatAt(i) + } + return vals +} + +// requireReal refuses a complex array a caller supplied as a +// constraint matrix or as a callback payload. The complex payload has +// no real part to read, so the library's FloatAt dereferences a nil +// payload and panics; every optim entry point answers complex input +// with an error instead. name identifies the entry point, what the +// payload, so the message reads like the dtype refusals the family +// already raises ("complex starting points are not supported"). +func requireReal(name, what string, a *core.Array) error { + if a.Dtype() != core.Complex { + return nil + } + return base.Errf("%s: complex %s are not supported", name, what) +} diff --git a/optim/optimise_test.go b/optim/optimise_test.go new file mode 100644 index 0000000..9c9f6da --- /dev/null +++ b/optim/optimise_test.go @@ -0,0 +1,175 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestFindRoot checks Brent's method on classic brackets, including +// the polynomial Whittaker and Robinson used to sell the method. +func TestFindRoot(t *testing.T) { + // x³ − 2x − 5 has its root near 2.0945514815423265 in [2, 3]. + root, err := FindRoot(func(x float64) float64 { return x*x*x - 2*x - 5 }, 2, 3, 0) + if err != nil { + t.Fatalf("FindRoot: %v", err) + } + want := 2.0945514815423265 + if math.Abs(root-want) > 1e-12 { + t.Fatalf("root = %.16g, want %.16g", root, want) + } + // sin(x) on [3, 4] brackets π. + pi, err := FindRoot(math.Sin, 3, 4, 0) + if err != nil { + t.Fatalf("FindRoot sin: %v", err) + } + if math.Abs(pi-math.Pi) > 1e-12 { + t.Fatalf("sin root = %.16g, want π", pi) + } + // A root exactly at an endpoint is returned. + zero, err := FindRoot(func(x float64) float64 { return x - 1 }, 1, 4, 0) + if err != nil { + t.Fatalf("FindRoot endpoint: %v", err) + } + if zero != 1 { + t.Fatalf("endpoint root = %v, want 1", zero) + } + // A bracket that does not change sign is refused. + if _, err := FindRoot(func(x float64) float64 { return x*x - 1 }, -2, 2, 0); err == nil { + t.Fatal("expected an error for a bracket without a sign change") + } + // NaN at the bracket is refused. + if _, err := FindRoot(func(x float64) float64 { return math.NaN() }, -1, 1, 0); err == nil { + t.Fatal("expected an error for a NaN bracket") + } +} + +// TestFindRootNewton checks the Newton iteration on the square root +// of two and pins the two failure contracts. +func TestFindRootNewton(t *testing.T) { + root, err := FindRootNewton(func(x float64) float64 { return x*x - 2 }, + func(x float64) float64 { return 2 * x }, 1, 1e-14, 50) + if err != nil { + t.Fatalf("FindRootNewton: %v", err) + } + if math.Abs(root-math.Sqrt2) > 1e-12 { + t.Fatalf("root = %.16g, want √2", root) + } + // A vanishing derivative is refused rather than dividing by zero. + if _, err := FindRootNewton(func(x float64) float64 { return x*x - 2 }, + func(x float64) float64 { return 0 }, 1, 0, 0); err == nil { + t.Fatal("expected an error for a vanishing derivative") + } + // An exhausted budget is an error. + if _, err := FindRootNewton(func(x float64) float64 { return x*x*x - 2 }, + func(x float64) float64 { return 3 * x * x }, 1000, 1e-14, 2); err == nil { + t.Fatal("expected an error for an exhausted iteration budget") + } +} + +// TestMinimise checks the simplex method on a convex sphere and on +// the Rosenbrock valley, the standard trap for naive descent. +func TestMinimise(t *testing.T) { + // Sphere: the minimum is the origin. + sphere := func(p *core.Array) (float64, error) { + sum := 0.0 + for i := range p.Len() { + sum += p.FloatAt(i) * p.FloatAt(i) + } + return sum, nil + } + start, err := core.FromFloats([]float64{2, -3, 1.5}, 3) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + point, value, err := Minimise(sphere, start, MinimiseOptions{}) + if err != nil { + t.Fatalf("Minimise sphere: %v", err) + } + // The default stopping rule bounds the function spread at 1e-10, + // which on a quadratic puts the point at roughly 1e-5 per + // coordinate; the contract is the spread, not exact zero. + if value > 1e-8 { + t.Fatalf("sphere minimum value = %v, want <= 1e-8", value) + } + for i := range point.Len() { + if math.Abs(point.FloatAt(i)) > 1e-3 { + t.Fatalf("sphere minimiser[%d] = %v, want ≈ 0", i, point.FloatAt(i)) + } + } + // A tighter tolerance buys a tighter minimum on the same objective. + point, value, err = Minimise(sphere, start, MinimiseOptions{Tolerance: 1e-22}) + if err != nil { + t.Fatalf("Minimise sphere tight: %v", err) + } + if value > 1e-20 { + t.Fatalf("tight sphere minimum value = %v, want <= 1e-20", value) + } + + // Rosenbrock in two dimensions: minimum at (1, 1). + rosenbrock := func(p *core.Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + } + rs, _ := core.FromFloats([]float64{-1.2, 1}, 2) + point, value, err = Minimise(rosenbrock, rs, MinimiseOptions{}) + if err != nil { + t.Fatalf("Minimise Rosenbrock: %v", err) + } + if math.Abs(point.FloatAt(0)-1) > 1e-4 || math.Abs(point.FloatAt(1)-1) > 1e-4 { + t.Fatalf("Rosenbrock minimiser = (%.8g, %.8g), want (1, 1)", + point.FloatAt(0), point.FloatAt(1)) + } + if value > 1e-8 { + t.Fatalf("Rosenbrock minimum value = %v, want ≈ 0", value) + } + + // An objective that errors must surface the error. + failing := func(p *core.Array) (float64, error) { + if p.FloatAt(0) > 0.5 { + return 0, base.Errf("objective exploded") + } + return p.FloatAt(0) * p.FloatAt(0), nil + } + fs, _ := core.FromFloats([]float64{0}, 1) + if _, _, err := Minimise(failing, fs, MinimiseOptions{}); err == nil { + t.Fatal("expected the objective's error to propagate") + } + + // An empty starting point is refused. + empty, _ := core.FromFloats(nil, 0) + if _, _, err := Minimise(sphere, empty, MinimiseOptions{}); err == nil { + t.Fatal("expected an error for an empty starting point") + } + if _, _, err := Minimise(sphere, mustComplexPoint(t), MinimiseOptions{}); err == nil { + t.Fatal("expected an error for a complex starting point") + } +} + +func mustComplexPoint(t *testing.T) *core.Array { + t.Helper() + a, err := core.FromComplexes([]complex128{1}, 1) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + return a +} + +// TestFindRootNewtonLargeScale pins the relative step tolerance: for a +// root at 1e8 the floating-point step granularity (eps·|x| ≈ 2e-8) +// exceeds an absolute 1e-12 tolerance, which used to report +// non-convergence from an accurate answer. +func TestFindRootNewtonLargeScale(t *testing.T) { + root, err := FindRootNewton(func(x float64) float64 { return x*x - 1e16 }, + func(x float64) float64 { return 2 * x }, 3e8, 1e-12, 100) + if err != nil { + t.Fatalf("FindRootNewton: %v", err) + } + if math.Abs(root-1e8) > 1e-4 { + t.Fatalf("root = %.12g, want 1e8", root) + } +} diff --git a/optim/qp.go b/optim/qp.go new file mode 100644 index 0000000..fabc7d9 --- /dev/null +++ b/optim/qp.go @@ -0,0 +1,447 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Convex quadratic programming by the primal active-set method on the +// house two-sided rows: +// +// min ½xᵀHx + c·x subject to l ≤ A·x ≤ u. +// +// The method walks the active sets: each iterate solves the +// equality-constrained subproblem that holds the working rows exactly +// at their walls, moves along the solution until a new row blocks, and +// adds the blocker to the working set. When the step vanishes, the +// working rows' multipliers are inspected: a row whose multiplier +// turned negative was pushed into the set by an earlier wall and is +// released. The subproblem is the KKT system +// +// [H Ãᵀ] [p] [−g] +// [à 0] [ν] = [ 0], +// +// with g = Hx + c the gradient at the iterate and à the working rows +// each carrying the sign of the wall it holds, written as the +// constraint's own gradient: a row held at its upper wall enters as +// +a, one at its lower wall as −a (the constraint reads l − a·x ≤ 0 +// there), and an equality as +a. The ν that come out are then the +// multipliers themselves: non-negative on the inequality rows, signed +// on the equalities. The system is factored by the same dense LU the +// revised simplex in simplex.go uses: one factorisation per iterate. +// +// Termination is the classic result for the method (Nocedal and +// Wright, Numerical Optimization, chapter 16.5): with H positive +// definite each subproblem has a unique solution, the objective +// strictly decreases on every non-zero step, the objective is the same +// on the zero steps that release a row, and the working set never +// repeats with a lower objective, so the loop reaches the unique KKT +// point in finitely many iterations on a nondegenerate problem. +// Degenerate ties (several rows blocking at one step, several rows +// tied at the most negative multiplier) are broken by the lowest row +// index, in the spirit of Bland's rule; a degenerate configuration +// that still cannot progress is stopped by the iteration budget as an +// error, never returned as a solution. Without positive definiteness +// none of that holds, so the Hessian is factorised at entry and a +// matrix that refuses the Cholesky factorisation is refused. + +// QPOptions tunes MinimiseQP. MaxIterations ≤ 0 means 1000 active-set +// rounds, Tolerance ≤ 0 means 1e-10. The tolerance is the threshold on +// the KKT step (below it the iterate is stationary on its working set), +// on the multiplier that releases a row, and on the symmetry of H; it +// is absolute, like the other tolerances in the package. +type QPOptions struct { + MaxIterations int + Tolerance float64 +} + +// MinimiseQP returns the point, the value ½xᵀHx + c·x and the row +// multipliers of the minimum of the strictly convex quadratic over the +// two-sided rows. The multipliers hold one entry per row of cons: an +// equality row carries its signed multiplier, an inequality row the +// non-negative multiplier of its active wall and zero when the row is +// slack, so complementary slackness reads directly off the slice. +// +// H must be symmetric and positive definite: anything else is refused +// at entry with the failed pivot named, because the termination +// guarantee and the uniqueness of the answer both rest on it. The +// starting point x0 may be nil, in which case a feasible point is +// found by running MinimiseLinearRows on the same rows with a zero +// cost, whose phase-1 refusal is the honest answer for an infeasible +// set; a supplied x0 must be feasible within 1e-9 and is refused with +// the worst violation otherwise. A nil or empty constraint set is the +// unconstrained quadratic and solves in one Newton step. +func MinimiseQP(h, c *core.Array, cons LinearConstraints, x0 *core.Array, opts QPOptions) (*core.Array, float64, []float64, error) { + const name = "MinimiseQP" + tol := opts.Tolerance + if tol <= 0 { + tol = 1e-10 + } + maxIter := opts.MaxIterations + if maxIter <= 0 { + maxIter = 1000 + } + if c.NDim() != 1 || c.Len() == 0 { + return nil, 0, nil, base.Errf("%s: c must be a non-empty rank-1 cost vector", name) + } + if c.Dtype() == core.Complex { + return nil, 0, nil, base.Errf("%s: complex costs are not supported", name) + } + n := c.Len() + qc := make([]float64, n) + for j := range n { + v := c.FloatAt(j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, 0, nil, base.Errf("%s: the cost carries a non-finite entry at %d", name, j+1) + } + qc[j] = v + } + if h == nil { + return nil, 0, nil, base.Errf("%s: the Hessian matrix is nil", name) + } + if err := requireReal(name, "Hessian matrices", h); err != nil { + return nil, 0, nil, err + } + if h.NDim() != 2 || h.Shape()[0] != n || h.Shape()[1] != n { + return nil, 0, nil, base.Errf("%s: the Hessian is %s, want %d×%d", name, base.ShapeText(h.Shape()), n, n) + } + hm := make([]float64, n*n) + for i := range n { + for j := range n { + v := h.FloatAt(i*n + j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, 0, nil, base.Errf("%s: the Hessian carries a non-finite entry at (%d, %d)", name, i+1, j+1) + } + hm[i*n+j] = v + } + } + for i := range n { + for j := i + 1; j < n; j++ { + if d := math.Abs(hm[i*n+j] - hm[j*n+i]); d > tol*math.Max(1, math.Abs(hm[i*n+j])) { + return nil, 0, nil, base.Errf("%s: the Hessian is not symmetric at (%d, %d): %g against %g", + name, i+1, j+1, hm[i*n+j], hm[j*n+i]) + } + } + } + if err := requirePositiveDefinite(hm, n, name); err != nil { + return nil, 0, nil, err + } + rows, lower, upper, equal, err := qpRows(cons, n, name) + if err != nil { + return nil, 0, nil, err + } + r := len(rows) + + // The start: a feasible point from the linear wrapper when none is + // supplied, or the caller's own point checked against the rows. + x := make([]float64, n) + if x0 == nil { + if r > 0 { + // The zero cost makes the wrapper return any feasible + // point; its phase-1 refusal is the honest answer for an + // infeasible set. + feas, _, err := MinimiseLinearRows(core.New(core.Float, n), cons, LinearProgramOptions{Tolerance: math.Max(tol, 1e-9)}) + if err != nil { + return nil, 0, nil, base.Errf("%s: %w", name, err) + } + for i := range n { + x[i] = feas.FloatAt(i) + } + } + } else { + if x0.Dtype() == core.Complex { + return nil, 0, nil, base.Errf("%s: complex starting points are not supported", name) + } + if x0.Len() != n { + return nil, 0, nil, base.Errf("%s: the starting point holds %d elements for %d variables", name, x0.Len(), n) + } + for i := range n { + x[i] = x0.FloatAt(i) + } + if worst, row := worstViolation(rows, lower, upper, x); worst > 1e-9 { + return nil, 0, nil, base.Errf("%s: the starting point violates row %d by %g", name, row+1, worst) + } + } + + // The working set: every equality row starts pinned to its wall; + // the inequality rows join as the steps meet them. wallSign is the + // KKT column sign: +1 on an upper wall or an equality, -1 on a + // lower wall. + active := make([]int, 0, r) + wallSign := make([]float64, 0, r) + inW := make([]bool, r) + for i := range r { + if equal[i] { + active = append(active, i) + wallSign = append(wallSign, 1) + inW[i] = true + } + } + + g := make([]float64, n) + kkt := make([]float64, (n+r)*(n+r)) + rhs := make([]float64, n+r) + sol := make([]float64, n+r) + // The KKT system is refactored once per iterate, so one factor + // serves the whole run: the factorisation workspace is rewritten + // from kkt every time. + var fac lu + + for range maxIter { + // The gradient at the iterate and the KKT system for the step + // that stays stationary on the working set. + matVec(hm, x, n, n, g) + for j := range n { + g[j] += qc[j] + } + w := len(active) + k := n + w + clear(rhs[:k]) + // Only the working rows' own block survives from the previous + // iteration: the Hessian block and the two constraint blocks are + // written whole just below, the w×w zero block of the KKT system + // is not written at all. + for i := n; i < k; i++ { + for j := n; j < k; j++ { + kkt[i*k+j] = 0 + } + } + for i := range n { + for j := range n { + kkt[i*k+j] = hm[i*n+j] + } + rhs[i] = -g[i] + } + for t := range w { + row := rows[active[t]] + for i := range n { + v := wallSign[t] * row[i] + kkt[i*k+n+t] = v + kkt[(n+t)*k+i] = v + } + } + if err := fac.factor(kkt[:k*k], k); err != nil { + return nil, 0, nil, base.Errf("%s: the working set has lost rank: %w", name, err) + } + fac.solve(rhs, sol) + pNorm := 0.0 + for j := range n { + pNorm = math.Max(pNorm, math.Abs(sol[j])) + } + if pNorm <= tol*math.Max(1, maxAbs(x)) { + // Stationary on the working set: an inequality row whose + // multiplier turned negative belongs off the set. The most + // negative multiplier leaves, ties to the lowest row + // index; with none left the KKT point is reached. + tie, worst := -1, 0.0 + for t := range w { + if equal[active[t]] { + continue + } + mu := sol[n+t] + if mu >= -tol { + continue + } + switch { + case tie == -1 || mu < worst-1e-12: + worst, tie = mu, t + case mu <= worst+1e-12 && active[t] < active[tie]: + tie = t + } + } + if tie == -1 { + var multipliers []float64 + if r > 0 { + multipliers = make([]float64, r) + for t := range w { + multipliers[active[t]] = sol[n+t] + } + } + if worst, row := worstViolation(rows, lower, upper, x); worst > 1e-9 { + return nil, 0, nil, base.Errf("%s: the solution violates row %d by %g", name, row+1, worst) + } + value := 0.5 * dot(x, matVecR(hm, x, n)) + for j := range n { + value += qc[j] * x[j] + } + out, fv := packResult(x, value) + return out, fv, multipliers, nil + } + inW[active[tie]] = false + active = append(active[:tie], active[tie+1:]...) + wallSign = append(wallSign[:tie], wallSign[tie+1:]...) + continue + } + // The largest step before an off-set row blocks: a row can only + // be crossed towards a finite wall, and the blocker hit first, + // ties to the lowest row index, joins the working set there. + alpha := 1.0 + blocker, blockSign := -1, 0.0 + for i := range r { + if inW[i] || equal[i] { + continue + } + d := dot(rows[i], sol[:n]) + var ai float64 + ai = math.Inf(1) + switch { + case d > 0 && upper[i] < math.Inf(1): + ai = (upper[i] - dot(rows[i], x)) / d + case d < 0 && lower[i] > math.Inf(-1): + ai = (lower[i] - dot(rows[i], x)) / d + } + if ai < 0 { + ai = 0 + } + if ai < 1 { + if ai < alpha-1e-12 { + alpha, blocker, blockSign = ai, i, wallKKTSign(d) + } else if ai <= alpha+1e-12 && (blocker == -1 || i < blocker) { + if ai < alpha { + alpha = ai + } + blocker, blockSign = i, wallKKTSign(d) + } + } + } + for j := range n { + x[j] += alpha * sol[j] + } + if blocker >= 0 && alpha < 1 { + active = append(active, blocker) + wallSign = append(wallSign, blockSign) + inW[blocker] = true + } + } + return nil, 0, nil, base.Errf("%s: the active-set budget of %d rounds ran out without reaching the KKT point", name, maxIter) +} + +// qpRows extracts the two-sided rows into plain slices and validates +// them with the same gates MinimiseConstrained applies to its matrix. +// A nil or empty set is no rows at all. +func qpRows(cons LinearConstraints, n int, name string) (rows [][]float64, lower, upper []float64, equal []bool, err error) { + if cons.A == nil { + return nil, nil, nil, nil, nil + } + if err := requireReal(name, "constraint matrices", cons.A); err != nil { + return nil, nil, nil, nil, err + } + if cons.A.NDim() != 2 || cons.A.Shape()[1] != n { + return nil, nil, nil, nil, base.Errf("%s: the constraint matrix is %s, want r×%d", name, base.ShapeText(cons.A.Shape()), n) + } + r := cons.A.Shape()[0] + if r == 0 { + return nil, nil, nil, nil, nil + } + if len(cons.Lower) != r || len(cons.Upper) != r { + return nil, nil, nil, nil, base.Errf("%s: the bounds hold %d and %d entries for %d rows", + name, len(cons.Lower), len(cons.Upper), r) + } + rows = make([][]float64, r) + lower = make([]float64, r) + upper = make([]float64, r) + equal = make([]bool, r) + // One backing block for every row: the rows are read-only after + // this build, so one allocation carries the whole block. + back := make([]float64, r*n) + for i := range r { + row := back[i*n : (i+1)*n] + for j := range n { + v := cons.A.FloatAt(i*n + j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, nil, nil, nil, base.Errf("%s: row %d carries a non-finite coefficient", name, i+1) + } + row[j] = v + } + lo, up := cons.Lower[i], cons.Upper[i] + if math.IsNaN(lo) || math.IsNaN(up) || lo > up { + return nil, nil, nil, nil, base.Errf("%s: row %d has bounds [%g, %g]", name, i+1, lo, up) + } + rows[i], lower[i], upper[i], equal[i] = row, lo, up, lo == up + } + return rows, lower, upper, equal, nil +} + +// worstViolation measures the rows at x and returns the largest +// violation with the row that carries it. +func worstViolation(rows [][]float64, lower, upper []float64, x []float64) (float64, int) { + worst, row := 0.0, 0 + for i := range rows { + ax := dot(rows[i], x) + v := math.Max(ax-upper[i], lower[i]-ax) + if v > worst { + worst, row = v, i + } + } + return worst, row +} + +// requirePositiveDefinite runs a Cholesky factorisation as the test: +// a pivot at or below the relative floor means the matrix is singular +// or indefinite, and both refuse the problem. +func requirePositiveDefinite(hm []float64, n int, name string) error { + scale := 1.0 + for i := range n { + scale = math.Max(scale, math.Abs(hm[i*n+i])) + } + work := make([]float64, n*n) + copy(work, hm) + for k := range n { + p := work[k*n+k] + for j := range k { + p -= work[k*n+j] * work[k*n+j] + } + if p <= 1e-13*scale { + return base.Errf("%s: the Hessian is not positive definite (pivot %g at %d)", name, p, k+1) + } + p = math.Sqrt(p) + work[k*n+k] = p + for i := k + 1; i < n; i++ { + s := work[i*n+k] + for j := range k { + s -= work[i*n+j] * work[k*n+j] + } + work[i*n+k] = s / p + } + } + return nil +} + +// matVec writes A·x into dst for a row-major m×n A. +func matVec(a []float64, x []float64, m, n int, dst []float64) { + for i := range m { + dst[i] = dot(a[i*n:i*n+n], x) + } +} + +// matVecR returns A·x as a fresh slice, the value form of matVec. +func matVecR(a []float64, x []float64, n int) []float64 { + dst := make([]float64, n) + matVec(a, x, n, n, dst) + return dst +} + +// dot is the plain inner product of two equal-length slices. +func dot(a, b []float64) float64 { + s := 0.0 + for i := range a { + s += a[i] * b[i] + } + return s +} + +// wallKKTSign maps a step's slope along a row to the KKT column sign +// of the wall it is heading for: a positive slope meets the upper wall +// (the constraint reads a·x − u ≤ 0 there), a negative one the lower +// wall (the constraint reads l − a·x ≤ 0). +func wallKKTSign(d float64) float64 { + if d > 0 { + return 1 + } + return -1 +} diff --git a/optim/qp_test.go b/optim/qp_test.go new file mode 100644 index 0000000..0c597e3 --- /dev/null +++ b/optim/qp_test.go @@ -0,0 +1,381 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// qpBowl builds the canonical data of ½xᵀHx + c·x for the squared +// distance to centre: H = 2I and c = −2·centre, the shape most of the +// hand-solved pins below use. +func qpBowl(t *testing.T, cx, cy float64) (*core.Array, *core.Array) { + t.Helper() + h := mustFloats(t, []float64{2, 0, 0, 2}, 2, 2) + c := mustFloats(t, []float64{-2 * cx, -2 * cy}) + return h, c +} + +// TestMinimiseQPActiveUpperWall pins the analytic case: the bowl's +// unconstrained minimum (2, 2) lies beyond x + y ≤ 2, the constrained +// optimum is the wall point (1, 1) with value −6 in the canonical +// form, and the row's multiplier is the hand-solved 2, positive as the +// KKT conditions demand for an active upper wall. +func TestMinimiseQPActiveUpperWall(t *testing.T) { + h, c := qpBowl(t, 2, 2) + cons := LinearConstraints{A: mustFloats(t, []float64{1, 1}, 1, 2), + Lower: []float64{math.Inf(-1)}, Upper: []float64{2}} + for _, x0 := range []*core.Array{nil, mustFloats(t, []float64{0, 0})} { + x, value, multipliers, err := MinimiseQP(h, c, cons, x0, QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-8 || math.Abs(x.FloatAt(1)-1) > 1e-8 { + t.Fatalf("point = (%.10g, %.10g), want (1, 1)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(value+6) > 1e-8 { + t.Fatalf("value = %.12g, want −6", value) + } + if len(multipliers) != 1 || math.Abs(multipliers[0]-2) > 1e-7 { + t.Fatalf("multipliers = %v, want [2]", multipliers) + } + } +} + +// TestMinimiseQPActiveLowerWall pins the lower-wall case: the bowl +// around (−3, −3) with y ≥ 0 bottoms out at (−3, 0) and the row's +// multiplier is the hand-solved 6. +func TestMinimiseQPActiveLowerWall(t *testing.T) { + h, c := qpBowl(t, -3, -3) + cons := LinearConstraints{A: mustFloats(t, []float64{0, 1}, 1, 2), + Lower: []float64{0}, Upper: []float64{math.Inf(1)}} + x, value, multipliers, err := MinimiseQP(h, c, cons, mustFloats(t, []float64{-3, 2}), QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + if math.Abs(x.FloatAt(0)+3) > 1e-8 || math.Abs(x.FloatAt(1)) > 1e-8 { + t.Fatalf("point = (%.10g, %.10g), want (−3, 0)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(value+9) > 1e-8 { + t.Fatalf("value = %.12g, want −9", value) + } + if len(multipliers) != 1 || math.Abs(multipliers[0]-6) > 1e-7 { + t.Fatalf("multipliers = %v, want [6]", multipliers) + } +} + +// TestMinimiseQPEqualityAndSlack pins complementary slackness on a +// problem with an equality row and a slack inequality: min x² + y² on +// x + y = 2 sits at (1, 1) with the signed equality multiplier −2, and +// the inactive wall y ≤ 3 must report exactly zero. +func TestMinimiseQPEqualityAndSlack(t *testing.T) { + h := mustFloats(t, []float64{2, 0, 0, 2}, 2, 2) + c := mustFloats(t, []float64{0, 0}) + cons := LinearConstraints{ + A: mustFloats(t, []float64{1, 1, 0, 1}, 2, 2), + Lower: []float64{2, math.Inf(-1)}, + Upper: []float64{2, 3}, + } + x, value, multipliers, err := MinimiseQP(h, c, cons, nil, QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-8 || math.Abs(x.FloatAt(1)-1) > 1e-8 { + t.Fatalf("point = (%.10g, %.10g), want (1, 1)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(value-2) > 1e-8 { + t.Fatalf("value = %.12g, want 2", value) + } + if len(multipliers) != 2 { + t.Fatalf("multipliers = %v, want one entry per row", multipliers) + } + if math.Abs(multipliers[0]+2) > 1e-7 { + t.Fatalf("equality multiplier = %.12g, want −2", multipliers[0]) + } + if multipliers[1] != 0 { + t.Fatalf("slack inequality multiplier = %.12g, want 0", multipliers[1]) + } +} + +// TestMinimiseQPRotated pins a genuinely coupled Hessian: min ½xᵀHx + +// c·x with H = [[2, 1], [1, 2]] and c = (−2, −2) on x − y ≥ ½. The +// hand-solved KKT point is (11/12, 5/12) with multiplier ¼: the +// stationarity Hx + c = ν(1, −1) and the row give three linear +// equations with exactly that solution. +func TestMinimiseQPRotated(t *testing.T) { + h := mustFloats(t, []float64{2, 1, 1, 2}, 2, 2) + c := mustFloats(t, []float64{-2, -2}) + cons := LinearConstraints{A: mustFloats(t, []float64{1, -1}, 1, 2), + Lower: []float64{0.5}, Upper: []float64{math.Inf(1)}} + x, _, multipliers, err := MinimiseQP(h, c, cons, mustFloats(t, []float64{0.5, 0}), QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + if math.Abs(x.FloatAt(0)-11.0/12.0) > 1e-8 || math.Abs(x.FloatAt(1)-5.0/12.0) > 1e-8 { + t.Fatalf("point = (%.10g, %.10g), want (11/12, 5/12)", x.FloatAt(0), x.FloatAt(1)) + } + if len(multipliers) != 1 || math.Abs(multipliers[0]-0.25) > 1e-7 { + t.Fatalf("multipliers = %v, want [0.25]", multipliers) + } +} + +// TestMinimiseQPUnconstrained pins the nil-or-empty constraint set: +// one Newton step to −H⁻¹c with no multipliers. +func TestMinimiseQPUnconstrained(t *testing.T) { + h := mustFloats(t, []float64{2, 0, 0, 4}, 2, 2) + c := mustFloats(t, []float64{-2, -8}) + x, value, multipliers, err := MinimiseQP(h, c, LinearConstraints{}, nil, QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-8 || math.Abs(x.FloatAt(1)-2) > 1e-8 { + t.Fatalf("point = (%.10g, %.10g), want (1, 2)", x.FloatAt(0), x.FloatAt(1)) + } + // ½(2·1 + 4·4) + (−2 −16) = 9 − 18 = −9. + if math.Abs(value+9) > 1e-8 { + t.Fatalf("value = %.12g, want −9", value) + } + if multipliers != nil { + t.Fatalf("multipliers = %v, want nil", multipliers) + } + // The unconstrained quadratic in three variables walks the + // Cholesky test past the first pivot: H = diag(2, 4, 6) and + // c = (−2, −8, −18) give the analytic minimum (1, 2, 3) with + // value −36. + h3 := mustFloats(t, []float64{2, 0, 0, 0, 4, 0, 0, 0, 6}, 3, 3) + c3 := mustFloats(t, []float64{-2, -8, -18}) + x3, v3, mult3, err := MinimiseQP(h3, c3, LinearConstraints{}, nil, QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + for i, w := range []float64{1, 2, 3} { + if math.Abs(x3.FloatAt(i)-w) > 1e-8 { + t.Fatalf("x3[%d] = %.10g, want %g", i, x3.FloatAt(i), w) + } + } + if math.Abs(v3+36) > 1e-8 { + t.Fatalf("value = %.10g, want −36", v3) + } + if mult3 != nil { + t.Fatalf("multipliers = %v, want nil", mult3) + } + // A constraint matrix with zero rows is no constraints: valid. + cons0 := LinearConstraints{A: core.New(core.Float, 0, 2)} + _, _, _, err = MinimiseQP(h, c, cons0, nil, QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP with zero rows: %v", err) + } +} + +// TestMinimiseQPCornerRows pins a two-row corner: min ½x² + 2y² − +// 3x − 3y on x + 2y ≤ 2, y ≥ 0.5. The unconstrained minimum (3, ¾) +// violates the first row, the A-face minimiser (1.75, 0.125) violates +// the second, so the answer is the corner (1, 0.5) with the +// hand-solved multipliers 2 and 3, both positive as an active corner +// requires. +func TestMinimiseQPCornerRows(t *testing.T) { + h := mustFloats(t, []float64{1, 0, 0, 4}, 2, 2) + c := mustFloats(t, []float64{-3, -3}) + cons := LinearConstraints{ + A: mustFloats(t, []float64{1, 2, 0, 1}, 2, 2), + Lower: []float64{math.Inf(-1), 0.5}, + Upper: []float64{2, math.Inf(1)}, + } + x, value, multipliers, err := MinimiseQP(h, c, cons, mustFloats(t, []float64{0, 1}), QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-7 || math.Abs(x.FloatAt(1)-0.5) > 1e-7 { + t.Fatalf("point = (%.10g, %.10g), want (1, 0.5)", x.FloatAt(0), x.FloatAt(1)) + } + canonical := 0.5*(x.FloatAt(0)*x.FloatAt(0)+4*x.FloatAt(1)*x.FloatAt(1)) - 3*x.FloatAt(0) - 3*x.FloatAt(1) + if math.Abs(value-canonical) > 1e-12 { + t.Fatalf("value %.12g disagrees with the canonical form %.12g", value, canonical) + } + if len(multipliers) != 2 { + t.Fatalf("multipliers = %v, want one entry per row", multipliers) + } + if math.Abs(multipliers[0]-2) > 1e-7 || math.Abs(multipliers[1]-3) > 1e-7 { + t.Fatalf("multipliers = (%.12g, %.12g), want (2, 3)", multipliers[0], multipliers[1]) + } +} + +// TestMinimiseQPReleaseRow drives a working-set release end to end: +// min ½(x² + 10y²) − 4x − 40y on x ≤ 1, x + y ≤ 2 from the origin. The +// Newton pull (4, 4) blocks x ≤ 1 first, the face step meets x + y = 2 +// at a zero-length step, and at that corner the first row's multiplier +// is negative, so it is released and the KKT point lands on the second +// row alone: x = (−16/11, 38/11) with multiplier 60/11 there and zero +// on the released row, all three figures hand-solved. +func TestMinimiseQPReleaseRow(t *testing.T) { + h := mustFloats(t, []float64{1, 0, 0, 10}, 2, 2) + c := mustFloats(t, []float64{-4, -40}) + cons := LinearConstraints{ + A: mustFloats(t, []float64{1, 0, 1, 1}, 2, 2), + Lower: []float64{math.Inf(-1), math.Inf(-1)}, + Upper: []float64{1, 2}, + } + x, value, multipliers, err := MinimiseQP(h, c, cons, mustFloats(t, []float64{0, 0}), QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } + if math.Abs(x.FloatAt(0)+16.0/11.0) > 1e-7 || math.Abs(x.FloatAt(1)-38.0/11.0) > 1e-7 { + t.Fatalf("point = (%.10g, %.10g), want (−16/11, 38/11)", x.FloatAt(0), x.FloatAt(1)) + } + canonical := 0.5*(x.FloatAt(0)*x.FloatAt(0)+10*x.FloatAt(1)*x.FloatAt(1)) - 4*x.FloatAt(0) - 40*x.FloatAt(1) + if math.Abs(value-canonical) > 1e-9 { + t.Fatalf("value %.12g disagrees with the canonical form %.12g", value, canonical) + } + if len(multipliers) != 2 { + t.Fatalf("multipliers = %v, want one entry per row", multipliers) + } + if multipliers[0] != 0 { + t.Fatalf("released row's multiplier = %.12g, want 0", multipliers[0]) + } + if math.Abs(multipliers[1]-60.0/11.0) > 1e-6 { + t.Fatalf("active row's multiplier = %.12g, want 60/11", multipliers[1]) + } +} + +// TestMinimiseQPRefusals checks the entry gates: indefinite or +// singular Hessians, asymmetry, shape mismatches, infeasible rows and +// an infeasible starting point. +func TestMinimiseQPRefusals(t *testing.T) { + goodH := mustFloats(t, []float64{2, 0, 0, 2}, 2, 2) + goodC := mustFloats(t, []float64{-2, -2}) + cases := []struct { + name string + run func() error + }{ + {"nil Hessian", func() error { + _, _, _, err := MinimiseQP(nil, goodC, LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"indefinite Hessian", func() error { + h := mustFloats(t, []float64{2, 0, 0, -2}, 2, 2) + _, _, _, err := MinimiseQP(h, goodC, LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"singular Hessian", func() error { + h := mustFloats(t, []float64{2, 0, 0, 0}, 2, 2) + _, _, _, err := MinimiseQP(h, goodC, LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"asymmetric Hessian", func() error { + h := mustFloats(t, []float64{2, 1, 0, 2}, 2, 2) + _, _, _, err := MinimiseQP(h, goodC, LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"Hessian shape", func() error { + h := mustFloats(t, []float64{2, 0, 0, 2, 0, 0, 0, 2, 0}, 3, 3) + _, _, _, err := MinimiseQP(h, goodC, LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"empty cost", func() error { + _, _, _, err := MinimiseQP(goodH, core.New(core.Float, 0), LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"complex cost", func() error { + _, _, _, err := MinimiseQP(goodH, mustComplexPoint(t), LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"NaN cost", func() error { + _, _, _, err := MinimiseQP(goodH, mustFloats(t, []float64{math.NaN(), -2}), LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"matrix shape", func() error { + cons := LinearConstraints{A: mustFloats(t, []float64{1, 1}), Lower: []float64{0}, Upper: []float64{1}} + _, _, _, err := MinimiseQP(goodH, goodC, cons, nil, QPOptions{}) + return err + }}, + {"crossed bounds", func() error { + cons := LinearConstraints{A: mustFloats(t, []float64{1, 1}, 1, 2), Lower: []float64{3}, Upper: []float64{1}} + _, _, _, err := MinimiseQP(goodH, goodC, cons, nil, QPOptions{}) + return err + }}, + {"NaN coefficient", func() error { + cons := LinearConstraints{A: mustFloats(t, []float64{math.NaN(), 1}, 1, 2), Lower: []float64{0}, Upper: []float64{1}} + _, _, _, err := MinimiseQP(goodH, goodC, cons, nil, QPOptions{}) + return err + }}, + {"complex Hessian", func() error { + ch, _ := core.FromComplexes([]complex128{2, 0, 0, 2}, 2, 2) + _, _, _, err := MinimiseQP(ch, goodC, LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"NaN Hessian entry", func() error { + h := mustFloats(t, []float64{2, 0, 0, math.NaN()}, 2, 2) + _, _, _, err := MinimiseQP(h, goodC, LinearConstraints{}, nil, QPOptions{}) + return err + }}, + {"duplicate equality rows", func() error { + // Two identical equalities pin the same wall twice: the + // KKT system loses rank and the refusal is honest. + cons := LinearConstraints{A: mustFloats(t, []float64{1, 1, 1, 1}, 2, 2), + Lower: []float64{2, 2}, Upper: []float64{2, 2}} + _, _, _, err := MinimiseQP(goodH, goodC, cons, mustFloats(t, []float64{1, 1}), QPOptions{}) + return err + }}, + {"short bounds with rows", func() error { + cons := LinearConstraints{A: mustFloats(t, []float64{1, 1, 0, 1}, 2, 2), + Lower: []float64{0}, Upper: []float64{1}} + _, _, _, err := MinimiseQP(goodH, goodC, cons, nil, QPOptions{}) + return err + }}, + {"complex constraint matrix", func() error { + ca, _ := core.FromComplexes([]complex128{1, 1, 1, 1}, 2, 2) + _, _, _, err := MinimiseQP(goodH, goodC, LinearConstraints{A: ca, Lower: []float64{0, 0}, Upper: []float64{1, 1}}, nil, QPOptions{}) + return err + }}, + {"infeasible rows", func() error { + cons := LinearConstraints{ + A: mustFloats(t, []float64{1, 0, 1, 0}, 2, 2), + Lower: []float64{1, math.Inf(-1)}, + Upper: []float64{math.Inf(1), 0}, + } + _, _, _, err := MinimiseQP(goodH, goodC, cons, nil, QPOptions{}) + return err + }}, + {"infeasible start", func() error { + cons := LinearConstraints{A: mustFloats(t, []float64{1, 1}, 1, 2), + Lower: []float64{math.Inf(-1)}, Upper: []float64{1}} + _, _, _, err := MinimiseQP(goodH, goodC, cons, mustFloats(t, []float64{1, 1}), QPOptions{}) + return err + }}, + {"complex start", func() error { + _, _, _, err := MinimiseQP(goodH, goodC, LinearConstraints{}, mustComplexPoint(t), QPOptions{}) + return err + }}, + {"start length", func() error { + _, _, _, err := MinimiseQP(goodH, goodC, LinearConstraints{}, mustFloats(t, []float64{1}), QPOptions{}) + return err + }}, + } + for _, c := range cases { + if err := c.run(); err == nil { + t.Fatalf("%s accepted", c.name) + } + } +} + +// TestMinimiseQPBudget pins the honest refusal when the active-set +// budget is too small for the two rounds the wall case needs, and the +// infeasibility message the linear phase propagates. +func TestMinimiseQPBudget(t *testing.T) { + h, c := qpBowl(t, 2, 2) + cons := LinearConstraints{A: mustFloats(t, []float64{1, 1}, 1, 2), + Lower: []float64{math.Inf(-1)}, Upper: []float64{2}} + _, _, _, err := MinimiseQP(h, c, cons, mustFloats(t, []float64{0, 0}), QPOptions{MaxIterations: 1}) + if err == nil || !strings.Contains(err.Error(), "budget") { + t.Fatalf("error = %v, want the budget refusal", err) + } + _, _, _, err = MinimiseQP(h, c, cons, mustFloats(t, []float64{0, 0}), QPOptions{}) + if err != nil { + t.Fatalf("MinimiseQP: %v", err) + } +} diff --git a/optim/rootsystem.go b/optim/rootsystem.go new file mode 100644 index 0000000..7b805ed --- /dev/null +++ b/optim/rootsystem.go @@ -0,0 +1,437 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "fmt" + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" + "sourcedock.dev/petrbalvin/tensor/linalg" +) + +// Root finding for systems of nonlinear equations r(x) = 0 with as +// many equations as unknowns. FindRoot and FindRootNewton cover the +// one-dimensional case; implicit solvers, equilibrium conditions and +// closure relations live in many dimensions, where a residual vector +// replaces the single equation. +// +// The method is the damped Newton iteration: each step solves the +// linearised system J·δ = −r with a central-difference Jacobian and +// the library's LU solver, then backtracks along δ until the residual +// norm actually falls, which keeps the iteration from leaping out of +// the basin of attraction when the linear model overshoots. +// +// RootSystemOptions.UseBroyden switches how the Jacobian is +// maintained: it is built by central differences once at the start +// and carried between steps by the rank-one Broyden update on +// its inverse, with a numerical rebuild whenever the update degrades +// (broydenMaintain documents the triggers). The default false keeps +// the per-step Jacobian and the iteration exactly as described here. + +// RootSystemOptions tunes FindRootSystem. Tolerance ≤ 0 means 1e-10 +// (an infinity-norm threshold on both the residual and the scaled +// step), MaxIterations ≤ 0 means 100. +type RootSystemOptions struct { + Tolerance float64 + MaxIterations int + // UseBroyden maintains the Jacobian across steps instead of + // rebuilding it every round: one central-difference Jacobian at + // the start, then the rank-one Broyden update on its inverse + // after every accepted step, with a numerical rebuild whenever the + // update degrades (the triggers are documented on + // broydenMaintain). The default false leaves the per-step + // Jacobian and the iteration's results exactly as they are + // without the option. + UseBroyden bool + // ParallelJacobian lets the central-difference Jacobian sweep its + // columns on several goroutines, the build at the start and every + // Broyden restart included. Setting it is the caller's consent + // that the residual callback may run concurrently from more than + // one goroutine: the default false keeps every evaluation on the + // caller's goroutine, and the run is bit for bit the same either + // way, because the columns are independent and each one is + // differenced by the same stencil. + ParallelJacobian bool +} + +// FindRootSystem solves r(x) = 0 for a vector function r of an n-vector, +// returning the solution and the residual infinity norm at it. The +// function must return a vector of the same length as x0. A singular +// Jacobian, a step that collapses before the tolerance is met or an +// exhausted iteration budget are errors, never silent answers. +// +// The iteration is a local solver: it follows the residual norm +// downhill, so a start whose basin contains no root, or one sitting +// in a parasitic minimum of the residual norm, reports the unfinished +// residual instead of pretending to converge. When the Newton system +// is singular the step falls back to the steepest-descent direction +// of the residual norm, which lets the iteration slide off the +// singular locus rather than dying there. +func FindRootSystem(f func(x *core.Array) (*core.Array, error), x0 *core.Array, + opts RootSystemOptions) (*core.Array, float64, error) { + if opts.Tolerance <= 0 { + opts.Tolerance = 1e-10 + } + if opts.MaxIterations <= 0 { + opts.MaxIterations = 100 + } + if x0.Dtype() == core.Complex { + return nil, 0, base.Errf("FindRootSystem: complex unknowns are not supported") + } + n := x0.Len() + if n == 0 { + return nil, 0, base.Errf("FindRootSystem: the initial guess must not be empty") + } + // eval reads the residual at x. The finiteness gate is strict for + // the states the iteration adopts (the start point and accepted + // iterates): a non-finite residual there is invisible to the + // norm's strict comparisons and would read as a converged root, + // the same refusal LevenbergMarquardt makes. Backtracking trials + // take the lenient variant: a trial that stepped into overflow is + // a rejected candidate to halve past, not a dead run, exactly the + // recovery the damping exists for. + eval := func(x []float64, strict bool, dst []float64) ([]float64, error) { + // The callback receives a copy, not an aliasing view of the + // reused iterate or stencil buffers: arrays are contractually + // immutable, and a callback that retains its argument must not + // observe the later writes. Minimise, L-BFGS and Levenberg- + // Marquardt copy for the same reason. + xa, aerr := core.FromFloats(x, len(x)) + if aerr != nil { + return nil, aerr + } + out, err := f(xa) + if err != nil { + return nil, err + } + if out.NDim() != 1 || out.Len() != n { + return nil, base.Errf("residual has shape %s, want a vector of length %d", + base.ShapeText(out.Shape()), n) + } + if err := requireReal("FindRootSystem", "residuals", out); err != nil { + return nil, err + } + // dst carries a buffer the Jacobian's stencil reuses across + // columns; the states the iteration keeps come back freshly + // allocated. Every entry of the buffer is written before it is + // read. + res := dst + if cap(res) < n { + res = make([]float64, n) + } + r := res[:n] + for i := range n { + r[i] = out.FloatAt(i) + if strict && (math.IsNaN(r[i]) || math.IsInf(r[i], 0)) { + // Bare, as Minimise's eval explains: every call site + // wraps once with the entry-point name, so the surfaced + // message carries exactly one prefix. + return nil, fmt.Errorf("the residual returned the non-finite value %g at coordinate %d", r[i], i) + } + } + return r, nil + } + + // cloneDense promotes through FloatAt, so Int and Float32 starting + // vectors behave exactly like Float64 ones (a RawFloats fast path + // would hand the iteration a nil slice for those dtypes). + x := cloneDense(x0) + r, err := eval(x, true, nil) + if err != nil { + return nil, 0, base.Errf("FindRootSystem: %w", err) + } + res := normInfOfStep(r) + jac := make([][]float64, n) + for i := range n { + jac[i] = make([]float64, n) + } + // The Newton solve runs on a working copy of the Jacobian, refilled + // from jac every round: base.Factor consumes its argument in place + // before it can fail, and the fallback below needs the pristine + // central-difference matrix. + jacWork := make([][]float64, n) + for i := range n { + jacWork[i] = make([]float64, n) + } + // Buffers reused across iterations: the two perturbed stencils, the + // backtracking candidate, the Newton column and the step. Each is + // fully rewritten before it is read, and only transiently wrapped + // views of them reach the callback. The two residual vectors of the + // Jacobian's stencil are carried the same way: a column is + // differenced from them and they are dead once it is written, so + // one pair per round replaces one pair per column. + xp := make([]float64, n) + xm := make([]float64, n) + candidate := make([]float64, n) + col := make([]float64, n) + colWrap := [][]float64{nil} + step := make([]float64, n) + var resPlus, resMinus []float64 + // Broyden state, allocated only under UseBroyden: the maintained + // inverse of the Jacobian and the flag saying it is fit to use, + // false until the first Jacobian is inverted and again after any + // restart, which sends the next round down the same + // build-and-invert path as the first. + useBroyden := opts.UseBroyden + var invJac [][]float64 + var sVec, yVec []float64 + haveInv := false + stalled := 0 + if useBroyden { + invJac = make([][]float64, n) + for i := range n { + invJac[i] = make([]float64, n) + } + sVec = make([]float64, n) + yVec = make([]float64, n) + } + for range opts.MaxIterations { + if res <= opts.Tolerance { + return linalg.ArrayFromFloatsSafe(x, len(x)), res, nil + } + var derr error + if useBroyden && haveInv { + // The maintained inverse turns the linearised solve into a + // matrix-vector product, δ = −H·r: no Jacobian build and + // no factorisation this round. + broydenStep(invJac, r, step) + } else { + // Central-difference Jacobian, one column per unknown. The + // stencil carries the offset on one unknown at a time, + // restored as soon as the column is differenced, so the + // whole point is copied once per round rather than once per + // column. column walks one unknown's stencil and writes + // that column of jac, and nothing else, so the bits it + // produces do not depend on which driver walks the columns. + column := func(j int, sp, sm, rp, rm []float64) ([]float64, []float64, error) { + eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(x[j])) + sp[j] += eps + sm[j] -= eps + rp, e1 := eval(sp, true, rp) + rm, e2 := eval(sm, true, rm) + sp[j], sm[j] = x[j], x[j] + if e1 != nil || e2 != nil { + return rp, rm, firstError(e1, e2) + } + for i := range n { + jac[i][j] = (rp[i] - rm[i]) / (2 * eps) + } + return rp, rm, nil + } + if opts.ParallelJacobian { + // The consent the option records lets the columns go to + // the engine's workers: each goroutine owns a disjoint + // run of columns, writes only into those columns of jac + // and reads only x, so no two workers write the same + // address and the sweep needs no locks. The build runs + // once at the start and again at every Broyden restart, + // both through this sweep. A failing column records its + // error and abandons the rest of its run; the reported + // one is the lowest failing column, the one the serial + // walk would hit first. + colErrs := make([]error, n) + engine.ParallelMin(n, 1, func(start, end int) { + sp, sm := make([]float64, n), make([]float64, n) + copy(sp, x) + copy(sm, x) + var rp, rm []float64 + for j := start; j < end; j++ { + var err error + rp, rm, err = column(j, sp, sm, rp, rm) + if err != nil { + colErrs[j] = err + return + } + } + }) + for _, err := range colErrs { + if err != nil { + return nil, 0, base.Errf("FindRootSystem: %w", err) + } + } + } else { + copy(xp, x) + copy(xm, x) + for j := range n { + var err error + resPlus, resMinus, err = column(j, xp, xm, resPlus, resMinus) + if err != nil { + return nil, 0, base.Errf("FindRootSystem: %w", err) + } + } + } + for i := range n { + col[i] = -r[i] + } + // Factor consumes the matrix in place: the pivot search swaps + // the rows and the elimination overwrites the subdiagonal with + // the multipliers, and the mutation lands in base.go before + // CheckSingular can reject the matrix. The solve therefore runs + // on the copy, so the steepest-descent fallback reads the + // pristine Jacobian: a direction formed from the factored + // matrix applies the permuted LU factors to r, which is not + // −Jᵀr and does not descend. + for i := range n { + copy(jacWork[i], jac[i]) + } + colWrap[0] = col + if useBroyden { + // First round or restart after a degraded update: + // rebuild the inverse and step with it at once, so a + // fresh Jacobian is never paid for without moving. A + // singular Jacobian takes the same steepest-descent + // escape the Newton path uses, with the inverse left + // unfit so the next round rebuilds. + if ierr := broydenInvert(jac, jacWork, invJac); ierr != nil { + derr = ierr + steepestDescentStep(jac, r, step) + } else { + haveInv = true + broydenStep(invJac, r, step) + } + } else { + sol, serr := base.SolveSystem("FindRootSystem", jacWork, colWrap) + derr = serr + if serr != nil { + // A singular Jacobian has no Newton direction, but the + // steepest-descent direction of the merit function ‖r‖² + // always exists when r ≠ 0 and always decreases it, so + // the iteration escapes the singular locus instead of + // dying there. + steepestDescentStep(jac, r, step) + } else { + for i := range n { + step[i] = sol[0][i] + } + } + } + } + // Backtrack along the step until the residual norm falls, at + // the Armijo fraction of the predicted linear decrease. + resSq := 0.0 + for i := range n { + resSq += r[i] * r[i] + } + // Scale the descent fallback down to Newton's magnitude so the + // first trial is comparable. + if derr != nil { + stepNorm := normInfOfStep(step) + if stepNorm > 0 { + scale := normInfOfStep(col) / stepNorm + if scale > 0 && !math.IsInf(scale, 0) { + for i := range step { + step[i] *= scale + } + } + } + } + alpha := 1.0 + improved := false + for range 60 { + for j := range n { + candidate[j] = x[j] + alpha*step[j] + } + rc, cerr := eval(candidate, false, nil) + if cerr != nil { + return nil, 0, base.Errf("FindRootSystem: %w", cerr) + } + candSq := 0.0 + for i := range n { + candSq += rc[i] * rc[i] + } + // A trial whose residual overflowed is a rejected candidate: + // its norm is Inf or NaN, compares false below, and the + // halving continues, which is what the damping exists for. + // A finite sum of squares admits only finite components, so + // whatever passes both comparisons is adoptable. + if candSq <= (1-1e-4*alpha)*resSq || candSq < opts.Tolerance*opts.Tolerance { + // The candidate buffer is reused by the next round's + // backtracking, so the accepted point is copied into + // the iteration's own state. + if useBroyden { + // The update's step and residual differences come + // from the state the accepted step leaves, so they + // are read before the copy overwrites it. + for j := range n { + sVec[j] = candidate[j] - x[j] + yVec[j] = rc[j] - r[j] + } + } + copy(x, candidate) + r = rc + resNew := normInfOfStep(r) + // The update corrects a maintained inverse, so it only + // runs when one is fit to be corrected: after a + // singular round the rebuild next loop takes over. + if useBroyden && haveInv { + // A false report leaves the inverse untouched and + // sends the next round back to a fresh Jacobian. + haveInv, stalled = broydenMaintain(invJac, sVec, yVec, r, resNew, res, stalled) + } + res = resNew + improved = true + break + } + alpha /= 2 + } + if !improved { + return nil, 0, base.Errf("FindRootSystem: the residual cannot be reduced below %g by damping", res) + } + // The documented success criterion is both a small step and a + // residual under the tolerance: a vanished step alone can come + // from a singular Jacobian whose steepest-descent fallback is + // numerically zero, and that point is a stall to keep working + // on, never a root to publish. + if res <= opts.Tolerance && normInfOfStep(step) <= opts.Tolerance*(1+normInfOfStep(x)) { + return linalg.ArrayFromFloatsSafe(x, len(x)), res, nil + } + } + // The last accepted step updated r and res after the loop-top test, + // so a run that met the tolerance exactly on the final iteration + // must re-test before the budget refusal reports it. + if res <= opts.Tolerance { + return linalg.ArrayFromFloatsSafe(x, len(x)), res, nil + } + return nil, 0, base.Errf("FindRootSystem: reached MaxIterations=%d with residual %g", + opts.MaxIterations, res) +} + +// 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 +} + +// steepestDescentStep writes the negative gradient of the merit +// function ‖r‖², −Jᵀr, into step: entry i sums column i of the +// Jacobian against the residual, so entry j of the Jacobian's row k +// carries ∂r_k/∂x_j. jac must be the pristine central-difference +// matrix, never a factored one: Factor overwrites the subdiagonal +// with the multipliers and swaps the rows in place, so a factored +// matrix yields a direction that is no descent direction at all. +func steepestDescentStep(jac [][]float64, r, step []float64) { + for i := range step { + s := 0.0 + for k := range r { + s += jac[k][i] * r[k] + } + step[i] = -s + } +} + +// firstError returns the first non-nil of two errors. +func firstError(e1, e2 error) error { + if e1 != nil { + return e1 + } + return e2 +} diff --git a/optim/rootsystem_test.go b/optim/rootsystem_test.go new file mode 100644 index 0000000..6732263 --- /dev/null +++ b/optim/rootsystem_test.go @@ -0,0 +1,188 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/linalg" + "testing" +) + +// TestFindRootSystemLinear checks the Newton step on a linear system: +// one iteration must land on the dense solve's answer. +func TestFindRootSystemLinear(t *testing.T) { + vals := []float64{ + 3, 1, 1, + 1, 4, 2, + 0, 1, 5, + } + b := []float64{1, 2, 3} + residual := func(x *core.Array) (*core.Array, error) { + out := core.New(core.Float, 3) + for i := range 3 { + s := 0.0 + for j := range 3 { + s += vals[i*3+j] * x.FloatAt(j) + } + out.RawFloats()[i] = s - b[i] + } + return out, nil + } + x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0, 0, 0}), RootSystemOptions{}) + if err != nil { + t.Fatalf("FindRootSystem: %v", err) + } + if res > 1e-12 { + t.Fatalf("residual %g, want ≤ 1e-12", res) + } + ref, err := linalg.Solve(mustFloats(t, vals, 3, 3), mustFloats(t, b)) + if err != nil { + t.Fatalf("Solve: %v", err) + } + for i := range 3 { + if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-12 { + t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i)) + } + } +} + +// TestFindRootSystemCircleLine solves x² + y² = 4 with x = y from the +// interior: the root is x = y = √2. +func TestFindRootSystemCircleLine(t *testing.T) { + residual := func(x *core.Array) (*core.Array, error) { + cx, cy := x.FloatAt(0), x.FloatAt(1) + return mustFloats(t, []float64{cx*cx + cy*cy - 4, cx - cy}), nil + } + x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0.5, 0.5}), RootSystemOptions{}) + if err != nil { + t.Fatalf("FindRootSystem: %v", err) + } + if res > 1e-10 { + t.Fatalf("residual %g, want ≤ 1e-10", res) + } + root := math.Sqrt2 + if math.Abs(x.FloatAt(0)-root) > 1e-10 || math.Abs(x.FloatAt(1)-root) > 1e-10 { + t.Fatalf("solution = (%.12g, %.12g), want (√2, √2)", x.FloatAt(0), x.FloatAt(1)) + } +} + +// TestFindRootSystemCircleHyperbola solves x² + y² = 5 with xy = 2, +// whose two roots are (1, 2) and (2, 1); from an interior start the +// iteration must land on one of them with a vanishing residual. The +// damped step is exercised by the curved crossing of the two curves. +func TestFindRootSystemCircleHyperbola(t *testing.T) { + residual := func(x *core.Array) (*core.Array, error) { + cx, cy := x.FloatAt(0), x.FloatAt(1) + return mustFloats(t, []float64{cx*cx + cy*cy - 5, cx*cy - 2}), nil + } + x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0.4, 2.2}), RootSystemOptions{}) + if err != nil { + t.Fatalf("FindRootSystem: %v", err) + } + if res > 1e-10 { + t.Fatalf("residual %g, want ≤ 1e-10", res) + } + root1 := []float64{1, 2} + root2 := []float64{2, 1} + hit := func(root []float64) bool { + return math.Abs(x.FloatAt(0)-root[0]) < 1e-9 && math.Abs(x.FloatAt(1)-root[1]) < 1e-9 + } + if !hit(root1) && !hit(root2) { + t.Fatalf("solution = (%.12g, %.12g), want (1, 2) or (2, 1)", x.FloatAt(0), x.FloatAt(1)) + } +} + +// TestFindRootSystemTrig mixes transcendental equations: cos x = y and +// sin x = y meet where x = π/4. +func TestFindRootSystemTrig(t *testing.T) { + residual := func(x *core.Array) (*core.Array, error) { + cx, cy := x.FloatAt(0), x.FloatAt(1) + return mustFloats(t, []float64{math.Cos(cx) - cy, math.Sin(cx) - cy}), nil + } + x, res, err := FindRootSystem(residual, mustFloats(t, []float64{0.3, 0.1}), RootSystemOptions{}) + if err != nil { + t.Fatalf("FindRootSystem: %v", err) + } + if res > 1e-10 { + t.Fatalf("residual %g, want ≤ 1e-10", res) + } + if math.Abs(x.FloatAt(0)-math.Pi/4) > 1e-9 || math.Abs(x.FloatAt(1)-math.Sqrt2/2) > 1e-9 { + t.Fatalf("solution = (%.12g, %.12g), want (π/4, 1/√2)", x.FloatAt(0), x.FloatAt(1)) + } +} + +func TestFindRootSystemErrors(t *testing.T) { + // The origin is a genuine root of (x², y²) and must be answered. + x, res, err := FindRootSystem(func(x *core.Array) (*core.Array, error) { + return mustFloats(t, []float64{x.FloatAt(0) * x.FloatAt(0), x.FloatAt(1) * x.FloatAt(1)}), nil + }, mustFloats(t, []float64{0, 0}), RootSystemOptions{}) + if err != nil || res > 0 || x.FloatAt(0) != 0 { + t.Fatalf("root at the origin: x = %v, res = %v, %v", x.RawFloats(), res, err) + } + // Duplicated equations have a rank-one Jacobian everywhere: the + // Newton system is singular while the residual is nonzero. + rankOne := func(x *core.Array) (*core.Array, error) { + cx := x.FloatAt(0) + return mustFloats(t, []float64{cx*cx - 1, cx*cx - 1}), nil + } + if _, _, err := FindRootSystem(rankOne, mustFloats(t, []float64{2}), RootSystemOptions{}); err == nil { + t.Fatal("hopeless rank-one system: want an error") + } + wrongShape := func(x *core.Array) (*core.Array, error) { + return mustFloats(t, []float64{1, 1, 1}), nil + } + if _, _, err := FindRootSystem(wrongShape, mustFloats(t, []float64{1, 1}), RootSystemOptions{}); err == nil { + t.Fatal("wrong residual shape: want an error") + } + fails := func(*core.Array) (*core.Array, error) { return nil, base.Errf("residual failed") } + if _, _, err := FindRootSystem(fails, mustFloats(t, []float64{1, 1}), RootSystemOptions{}); err == nil { + t.Fatal("residual error: want an error") + } + // An exhausted budget must be reported, not approximated away. + circleResidual := func(x *core.Array) (*core.Array, error) { + cx, cy := x.FloatAt(0), x.FloatAt(1) + return mustFloats(t, []float64{cx*cx + cy*cy - 4, cx - cy}), nil + } + if _, _, err := FindRootSystem(circleResidual, mustFloats(t, []float64{1, 1}), + RootSystemOptions{MaxIterations: 1, Tolerance: 1e-20}); err == nil { + t.Fatal("exhausted budget: want an error") + } + if _, _, err := FindRootSystem(circleResidual, mustFloats(t, []float64{}), RootSystemOptions{}); err == nil { + t.Fatal("empty guess: want an error") + } +} + +// TestFindRootSystemDtypes pins the starting-vector promotion: Int and +// Float32 starts must behave exactly like Float64 ones. The old +// RawFloats() fast path handed the iteration a nil or zero slice for +// those dtypes. +func TestFindRootSystemDtypes(t *testing.T) { + residual := func(x *core.Array) (*core.Array, error) { + cx, cy := x.FloatAt(0), x.FloatAt(1) + return mustFloats(t, []float64{cx*cx + cy*cy - 4, cx - cy}), nil + } + ints, err := core.FromInts([]int64{1, 1}, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + thirtyTwo, err := core.FromFloat32s([]float32{0.5, 0.5}, 2) + if err != nil { + t.Fatalf("FromFloat32s: %v", err) + } + for name, x0 := range map[string]*core.Array{"int": ints, "float32": thirtyTwo} { + x, res, err := FindRootSystem(residual, x0, RootSystemOptions{}) + if err != nil { + t.Fatalf("FindRootSystem(%s start): %v", name, err) + } + if res > 1e-10 { + t.Fatalf("FindRootSystem(%s start): residual %g, want ≤ 1e-10", name, res) + } + if math.Abs(x.FloatAt(0)-math.Sqrt2) > 1e-8 || math.Abs(x.FloatAt(1)-math.Sqrt2) > 1e-8 { + t.Fatalf("FindRootSystem(%s start): solution = (%.12g, %.12g), want (√2, √2)", + name, x.FloatAt(0), x.FloatAt(1)) + } + } +} diff --git a/optim/rowblocks_pin_test.go b/optim/rowblocks_pin_test.go new file mode 100644 index 0000000..cad137f --- /dev/null +++ b/optim/rowblocks_pin_test.go @@ -0,0 +1,48 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// MinimiseLinearRows hands both standard rows of a two-sided constraint +// slots in the one backing block, taken in build order. A fixture with +// a two-sided row and a one-sided row pins the hand-computed optimum, +// which a row that aliased its neighbour could not reach. + +func TestMinimiseLinearRowsTwoSidedConstraint(t *testing.T) { + // minimise x1 - 2*x2 over 2 ≤ x1 + x2 ≤ 3 and 0 ≤ x1 ≤ 1: the + // objective pushes x2 up and x1 down, so x1 = 0, x2 = 3 and the + // optimum is -6. + cost, err := core.FromFloats([]float64{1, -2}, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + a, err := core.FromFloats([]float64{1, 1, 1, 0}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + cons := LinearConstraints{ + A: a, + Lower: []float64{2, 0}, + Upper: []float64{3, 1}, + } + x, value, merr := MinimiseLinearRows(cost, cons, LinearProgramOptions{}) + if merr != nil { + t.Fatalf("MinimiseLinearRows: %v", merr) + } + if math.Abs(value+6) > 1e-9 { + t.Fatalf("optimum = %v, want -6", value) + } + if v := x.FloatAt(0); math.Abs(v) > 1e-9 { + t.Fatalf("x1 = %v, want 0", v) + } + if v := x.FloatAt(1); math.Abs(v-3) > 1e-9 { + t.Fatalf("x2 = %v, want 3", v) + } +} diff --git a/optim/simplex.go b/optim/simplex.go new file mode 100644 index 0000000..fac0449 --- /dev/null +++ b/optim/simplex.go @@ -0,0 +1,814 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Linear programming by the revised simplex method on the standard +// form +// +// min c·x subject to A·x = b, x ≥ 0. +// +// Free or two-sided quantities belong to the caller's own conversion: +// the wrapper MinimiseLinearRows turns the house rows l ≤ A·x ≤ u into +// this form mechanically (a free variable splits into the difference +// of two non-negative ones, each finite row side gains a slack), so a +// caller with ordinary bounds never touches the standard form at all. +// +// The method is the two-phase revised simplex. Phase 1 minimises the +// sum of the artificial variables that carry the starting basis, so +// its optimum is either zero, which leaves a feasible basis in hand, +// or the total infeasibility of the rows, which refuses the problem +// with that figure as the evidence. Phase 2 prices the real columns +// from the feasible basis and walks along vertices to the optimum. +// +// Both phases pick the entering column by Bland's rule: the +// lowest-indexed column whose reduced cost is negative, and, among the +// rows tied at the minimum ratio, the lowest-indexed basic variable to +// leave. The rule is slower than Dantzig's most-negative pricing but +// it cannot cycle: on a degenerate problem, where several bases carry +// the same vertex and the classic rule can pivot forever, Bland's rule +// is guaranteed to terminate (Bland, 1977). Redundant rows surface in +// phase 1 as artificial columns that will not leave: a row no real +// column can pivot out is a linear combination of the others, so the +// row and its artificial leave the problem together and the reduced +// basis stays valid. +// +// The basis is refactorised by a dense LU with partial pivoting at +// every pivot. The solver targets the small dense problems a library +// of this shape meets, where O(m³) per pivot is cheap and a fresh +// factorisation keeps the iteration honest where an updated inverse +// would drift. The same factorisation machinery carries the +// active-set solver in qp.go. + +// LinearProgramOptions tunes MinimiseLinear and MinimiseLinearRows. +// MaxIterations ≤ 0 means 10000 pivots, Tolerance ≤ 0 means 1e-9. The +// tolerance prices reduced costs and separates ratio-test ties, and it +// is absolute in the scale the caller's costs and rows carry, so a +// badly scaled problem should be rescaled to O(1) first, as with the +// other tolerances in the package. +type LinearProgramOptions struct { + MaxIterations int + Tolerance float64 +} + +// MinimiseLinear returns the point and value of the minimum of c·x +// over the standard-form polytope A·x = b with x ≥ 0. The contract is +// the standard form exactly: every variable is non-negative, every row +// is an equality, and a caller holding inequalities, free variables or +// bounds converts them first (MinimiseLinearRows does that conversion +// for the house two-sided rows). The returned point has one entry per +// column of A, slack columns included when the caller built them into +// the standard form. +// +// An infeasible problem is refused with the phase-1 evidence: the +// total infeasibility the artificial phase ended with and the row that +// carries the worst of it. An unbounded objective is refused with the +// column that prices out as a profitable ray no row limits. A run that +// spends the pivot budget without pricing out is an error, never a +// silent answer: under Bland's rule an exhausted budget on a +// well-scaled problem is the signature of a tolerance the data does +// not support. A problem with no rows is the simplex over x ≥ 0: it +// returns the origin when every cost is non-negative and refuses as +// unbounded when one is not. +func MinimiseLinear(c, a, b *core.Array, opts LinearProgramOptions) (*core.Array, float64, error) { + const name = "MinimiseLinear" + if c.NDim() != 1 || c.Len() == 0 { + return nil, 0, base.Errf("%s: c must be a non-empty rank-1 cost vector", name) + } + if c.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex costs are not supported", name) + } + n := c.Len() + if a == nil { + return nil, 0, base.Errf("%s: the constraint matrix is nil", name) + } + if err := requireReal(name, "constraint matrices", a); err != nil { + return nil, 0, err + } + if a.NDim() != 2 || a.Shape()[1] != n { + return nil, 0, base.Errf("%s: the constraint matrix is %s, want m×%d", name, base.ShapeText(a.Shape()), n) + } + m := a.Shape()[0] + if b.NDim() != 1 || b.Len() != m { + return nil, 0, base.Errf("%s: b must be a rank-1 vector with one entry per row (%d)", name, m) + } + if b.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex right-hand sides are not supported", name) + } + cost := make([]float64, n) + for j := range n { + v := c.FloatAt(j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, 0, base.Errf("%s: the cost carries a non-finite entry at %d", name, j+1) + } + cost[j] = v + } + // One backing block for every standard row: a row is built once + // here, extended in place by graftArtificials and never outgrows + // its slot, so one allocation carries the whole block. + rows := make([][]float64, m) + back := make([]float64, m*(n+m)) + rhs := make([]float64, m) + for i := range m { + // Rows carry their artificial column from the start: the tail + // stays zero until graftArtificials writes the unit entry, so + // the phases read the same values a freshly extended row held. + row := back[i*(n+m) : (i+1)*(n+m)] + for j := range n { + v := a.FloatAt(i*n + j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, 0, base.Errf("%s: row %d carries a non-finite coefficient", name, i+1) + } + row[j] = v + } + v := b.FloatAt(i) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, 0, base.Errf("%s: the right-hand side carries a non-finite entry at %d", name, i+1) + } + // The artificial basis needs b ≥ 0, so a negative row is + // negated whole: the feasible set is unchanged. + if v < 0 { + for j := range n { + row[j] = -row[j] + } + v = -v + } + rows[i], rhs[i] = row, v + } + prob := &standardForm{rows: rows, b: rhs, nreal: n} + x, value, err := solveTwoPhase(prob, cost, opts, name) + if err != nil { + return nil, 0, err + } + out, fv := packResult(x, value) + return out, fv, nil +} + +// MinimiseLinearRows returns the point and value of the minimum of c·x +// subject to the two-sided rows l ≤ A·x ≤ u carried by cons, the same +// rows LinearConstraints holds for MinimiseConstrained. The variables +// are free: a bound on a variable is just a row with a unit +// coefficient, as the linear-constraint tests build them. A row with +// Lower = Upper is an equality; an infinite bound opens that side; a +// row open at both ends constrains nothing and is dropped from the +// standard form. +// +// The conversion is mechanical and exact: each variable x splits into +// the difference of two non-negative columns, each finite upper side +// gains a slack column added to the row, each finite lower side a +// slack subtracted, and an equality row passes through bare. The two +// entries a variable splits into cancel in the objective, so the +// standard-form optimum back-substitutes to the original variables and +// the reported value is c·x computed on them. +// +// Infeasibility, unboundedness and budget exhaustion are refused +// exactly as MinimiseLinear refuses them. +func MinimiseLinearRows(c *core.Array, cons LinearConstraints, opts LinearProgramOptions) (*core.Array, float64, error) { + const name = "MinimiseLinearRows" + if c.NDim() != 1 || c.Len() == 0 { + return nil, 0, base.Errf("%s: c must be a non-empty rank-1 cost vector", name) + } + if c.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex costs are not supported", name) + } + n := c.Len() + if cons.A == nil { + return nil, 0, base.Errf("%s: the constraint matrix is nil", name) + } + if err := requireReal(name, "constraint matrices", cons.A); err != nil { + return nil, 0, err + } + if cons.A.NDim() != 2 || cons.A.Shape()[1] != n { + return nil, 0, base.Errf("%s: the constraint matrix is %s, want r×%d", name, base.ShapeText(cons.A.Shape()), n) + } + r := cons.A.Shape()[0] + if r == 0 { + return nil, 0, base.Errf("%s: the constraint matrix has no rows", name) + } + if len(cons.Lower) != r || len(cons.Upper) != r { + return nil, 0, base.Errf("%s: the bounds hold %d and %d entries for %d rows", + name, len(cons.Lower), len(cons.Upper), r) + } + cost := make([]float64, n) + for j := range n { + v := c.FloatAt(j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, 0, base.Errf("%s: the cost carries a non-finite entry at %d", name, j+1) + } + cost[j] = v + } + // The standard form: n split pairs, then one slack per finite + // non-equality side. Count the slacks and the materialised rows + // first so every row slice is allocated once, wide enough for its + // artificial column. + slacks := 0 + built := 0 + for i := range r { + lo, up := cons.Lower[i], cons.Upper[i] + if math.IsNaN(lo) || math.IsNaN(up) || lo > up { + return nil, 0, base.Errf("%s: row %d has bounds [%g, %g]", name, i+1, lo, up) + } + if lo == up && math.IsInf(lo, 0) { + return nil, 0, base.Errf("%s: row %d is an equality at infinity", name, i+1) + } + if lo == up { + built++ + continue + } + if up < math.Inf(1) { + slacks++ + built++ + } + if lo > math.Inf(-1) { + slacks++ + built++ + } + } + for i := range r { + for j := range n { + if v := cons.A.FloatAt(i*n + j); math.IsNaN(v) || math.IsInf(v, 0) { + return nil, 0, base.Errf("%s: row %d carries a non-finite coefficient", name, i+1) + } + } + } + cols := 2*n + slacks + prob := &standardForm{nreal: cols} + rows := make([][]float64, 0, r) + // One backing block for every built row: a row is written once + // here, extended in place by graftArtificials and never outgrows + // its slot, so one allocation carries the whole block. + back := make([]float64, built*(cols+built)) + rhs := make([]float64, 0, r) + slackCol := 2 * n + for i := range r { + lo, up := cons.Lower[i], cons.Upper[i] + // build materialises one standard row for one finite side. The + // slack argument is +1 on an upper side, -1 on a lower one and + // 0 on a bare equality. A negative right-hand side is negated + // whole, coefficients, slack and all, because the artificial + // basis the two-phase start needs requires b >= 0 in every + // row; negating flips the slack's side but the sign convention + // of the bound row survives the flip. + build := func(slack float64, bound float64) { + // The row carries its artificial column from the start, the + // same in-place extension MinimiseLinear builds. + row := back[len(rows)*(cols+built) : (len(rows)+1)*(cols+built)] + for j := range n { + v := cons.A.FloatAt(i*n + j) + row[j], row[n+j] = v, -v + } + if slack != 0 { + row[slackCol] = slack + slackCol++ + } + if bound < 0 { + for j := range row { + row[j] = -row[j] + } + bound = -bound + } + rows = append(rows, row) + rhs = append(rhs, bound) + } + switch { + case lo == up: + build(0, up) + default: + if up < math.Inf(1) { + build(1, up) + } + if lo > math.Inf(-1) { + build(-1, lo) + } + } + } + prob.rows, prob.b = rows, rhs + stdCost := make([]float64, cols) + copy(stdCost, cost) + for j := range n { + stdCost[n+j] = -cost[j] + } + xStd, _, err := solveTwoPhase(prob, stdCost, opts, name) + if err != nil { + return nil, 0, err + } + // Back-substitute x = p − q and value the original cost on the + // original variables: the split's two halves cancel only in exact + // arithmetic, so the caller sees the recomputed figure. + x := make([]float64, n) + value := 0.0 + for j := range n { + x[j] = xStd[j] - xStd[n+j] + value += cost[j] * x[j] + } + out, fv := packResult(x, value) + return out, fv, nil +} + +// standardForm is the working copy the two-phase method runs on: the +// rows a·x = b with b ≥ 0 after negation, nreal real columns, and one +// artificial column per row appended behind them. Row drops during the +// phase transition shorten rows and b together with the basis. +// +// bm and fac are the reusable basis matrix and its factorisation: the +// basis is gathered afresh and refactorised at every pivot, which +// rewrites the whole m×m matrix, so one buffer per solve replaces one +// per pivot. Every entry of bm is written before it is read. The +// per-pivot vectors ride the same rule: the pricing, ratio and solution +// sweeps each overwrite the whole live prefix before reading it, so one +// set of buffers serves every pivot of one solve. +type standardForm struct { + rows [][]float64 + b []float64 + nreal int + bm []float64 + fac lu + xb []float64 + pi []float64 + cb []float64 + col []float64 + w []float64 + unit []float64 + y []float64 +} + +// growF returns buf at length n, allocating only when the current +// capacity falls short; every caller overwrites the whole prefix. +func growF(buf []float64, n int) []float64 { + if cap(buf) < n { + return make([]float64, n) + } + return buf[:n] +} + +// cols is the total column count: the real columns plus one artificial +// per row still carried. +func (s *standardForm) cols() int { return s.nreal + len(s.rows) } + +// graftArtificials extends every row with the artificial identity +// columns the artificial phase runs on: column nreal + r is the r-th +// unit vector. It runs once, before phase 1. A row the entry points +// built already wide enough for its artificial is extended in place: +// the tail slots hold zeros until the unit entry is written, so the +// values the phases read are the ones a freshly built row carried. +func (s *standardForm) graftArtificials() { + m := len(s.rows) + for i := range m { + if len(s.rows[i]) >= s.nreal+m { + s.rows[i] = s.rows[i][:s.nreal+m] + s.rows[i][s.nreal+i] = 1 + continue + } + row := make([]float64, s.nreal+m) + copy(row, s.rows[i]) + row[s.nreal+i] = 1 + s.rows[i] = row + } +} + +// solveTwoPhase runs the artificial phase, refuses an infeasible +// problem with its evidence, expels the surviving artificials, and +// runs the real phase. It returns the real part of the solution and +// the objective c·x valued on it. +func solveTwoPhase(s *standardForm, cost []float64, opts LinearProgramOptions, name string) ([]float64, float64, error) { + tol := opts.Tolerance + if tol <= 0 { + tol = 1e-9 + } + budget := opts.MaxIterations + if budget <= 0 { + budget = 10000 + } + m := len(s.rows) + s.graftArtificials() + basis := make([]int, m) + inBasic := make([]bool, s.cols()) + for i := range m { + basis[i] = s.nreal + i + inBasic[basis[i]] = true + } + if m > 0 { + // Phase 1: minimise the sum of the artificials. They start as + // the basis (the identity, with b ≥ 0), and once one leaves it + // never re-enters: canEnter admits the real columns only. + cost1 := make([]float64, s.cols()) + for j := s.nreal; j < s.cols(); j++ { + cost1[j] = 1 + } + enter1 := make([]bool, s.cols()) + for j := range s.nreal { + enter1[j] = true + } + if err := s.pivotLoop(basis, inBasic, cost1, enter1, tol, budget, name, "phase 1", true); err != nil { + return nil, 0, err + } + // The phase-1 optimum is the total infeasibility: anything + // above the tolerance is an infeasible problem, refused with + // the figure and the worst offending row as the evidence. + residual, worst, worstRow, aerr := s.artificialSum(basis) + if aerr != nil { + return nil, 0, base.Errf("%s: %w", name, aerr) + } + if residual > tol*math.Max(1, maxAbs(s.b)) { + return nil, 0, base.Errf("%s: the problem is infeasible: phase 1 ended with an infeasibility of %g (row %d still carries %g)", + name, residual, worstRow+1, worst) + } + var err error + if basis, err = s.expelArtificials(basis, inBasic, tol); err != nil { + return nil, 0, base.Errf("%s: %w", name, err) + } + } + // Phase 2: the real costs over a feasible basis. The artificials + // are gone from the basis and canEnter keeps them out of the + // pricing. + enter2 := make([]bool, s.cols()) + for j := range s.nreal { + enter2[j] = true + } + if err := s.pivotLoop(basis, inBasic, cost, enter2, tol, budget, name, "phase 2", false); err != nil { + return nil, 0, err + } + return s.solution(basis, cost) +} + +// basisMatrix gathers the basis columns into dst as a row-major m×m +// matrix for the factorisation. dst is grown to m² if it is too short +// and returned; every entry of the m×m block is written. +func (s *standardForm) basisMatrix(dst []float64, basis []int) []float64 { + m := len(s.rows) + if cap(dst) < m*m { + dst = make([]float64, m*m) + } + dst = dst[:m*m] + for r := range m { + row := s.rows[r] + for k, col := range basis { + dst[r*m+k] = row[col] + } + } + return dst +} + +// refactor gathers the basis columns and factors them into the form's +// own reusable factorisation, which is fully rewritten: the pivot loop, +// the artificial sum, the artificial expulsion and the final solution +// all read the basis this way. +func (s *standardForm) refactor(basis []int) (*lu, error) { + s.bm = s.basisMatrix(s.bm, basis) + if err := s.fac.factor(s.bm, len(s.rows)); err != nil { + return nil, err + } + return &s.fac, nil +} + +// pivotLoop is the revised simplex iteration: refactorise the basis, +// price the eligible non-basic columns, and pivot under Bland's rule +// until no eligible column prices out negatively. The phase1 flag +// shapes the diagnostics only: an unbounded ray is how phase 2 reports +// an unbounded objective and a contradiction in phase 1, whose +// objective is bounded below by zero. +func (s *standardForm) pivotLoop(basis []int, inBasic []bool, cost []float64, canEnter []bool, tol float64, budget int, + name, phase string, phase1 bool) error { + m := len(s.rows) + xb := growF(s.xb, m) + pi := growF(s.pi, m) + cb := growF(s.cb, m) + col := growF(s.col, m) + w := growF(s.w, m) + s.xb, s.pi, s.cb, s.col, s.w = xb, pi, cb, col, w + // A basic value that rounds a hair below zero after a solve is + // clamped; one that is genuinely negative means the basis lost its + // primal feasibility, which is a defect, not an answer. + floor := -1e-9 * math.Max(1, maxAbs(s.b)) + for piv := range budget { + f, err := s.refactor(basis) + if err != nil { + return base.Errf("%s: %s: %w after %d pivots", name, phase, err, piv) + } + f.solve(s.b, xb) + for i := range m { + if xb[i] < 0 { + if xb[i] < floor { + return base.Errf("%s: %s: the basis lost primal feasibility at row %d (%g) after %d pivots", + name, phase, i+1, xb[i], piv) + } + xb[i] = 0 + } + } + for i, c := range basis { + cb[i] = cost[c] + } + f.solveT(cb, pi) + // Bland's entering rule: the lowest-indexed eligible column + // whose reduced cost is negative. + enter := -1 + for j := range s.cols() { + if inBasic[j] || !canEnter[j] { + continue + } + d := cost[j] + for r := range m { + d -= pi[r] * s.rows[r][j] + } + if d < -tol { + enter = j + break + } + } + if enter == -1 { + return nil + } + for r := range m { + col[r] = s.rows[r][enter] + } + f.solve(col, w) + theta := math.Inf(1) + for i := range m { + if w[i] > tol { + theta = math.Min(theta, xb[i]/w[i]) + } + } + if math.IsInf(theta, 1) { + if phase1 { + return base.Errf("%s: %s: an unbounded ray contradicts the phase-1 objective, which is bounded below by zero", name, phase) + } + return base.Errf("%s: the objective is unbounded below: column %d prices out as a profitable ray no row limits", + name, enter+1) + } + // Bland's leaving rule: among the rows tied at the minimum + // ratio, the lowest-indexed basic variable leaves. The index, + // not the row position, is what the anti-cycling proof needs. + tie := 1e-9 * math.Max(1, math.Abs(theta)) + leave := -1 + for i := range m { + if w[i] > tol && xb[i]/w[i] <= theta+tie { + if leave == -1 || basis[i] < basis[leave] { + leave = i + } + } + } + inBasic[basis[leave]] = false + basis[leave] = enter + inBasic[enter] = true + } + return base.Errf("%s: %s: the pivot budget of %d ran out without pricing out", name, phase, budget) +} + +// artificialSum totals the basic artificials' values after phase 1: +// their sum is the total infeasibility phase 1 minimised. +func (s *standardForm) artificialSum(basis []int) (total, worst float64, worstRow int, err error) { + // The identical basis was just factored without error at the top + // of the pivot loop's final iteration; the guard keeps the + // invariant explicit rather than trusted. + f, err := s.refactor(basis) + if err != nil { + return 0, 0, -1, err + } + xb := growF(s.xb, len(s.rows)) + s.xb = xb + f.solve(s.b, xb) + total, worst, worstRow = 0, 0, -1 + for i, c := range basis { + if c >= s.nreal { + total += xb[i] + // Row 0 is a legal carrier of the worst infeasibility, so + // the unset sentinel is −1, not the zero the scan starts + // from: with 0 here any later, smaller artificial would + // overwrite the evidence through the disjunct. + if worstRow < 0 || xb[i] > worst { + worst, worstRow = xb[i], i + } + } + } + return total, worst, worstRow, nil +} + +// expelArtificials drives every artificial still basic after phase 1 +// out of the basis. A pivot on any real column with a non-zero entry +// in the artificial's row removes it directly (the pivot is +// degenerate: the artificial's value is zero at the phase-1 optimum). +// A row where no real column has such an entry is redundant, a linear +// combination of the others at the current vertex, so the row and its +// artificial leave the problem together and the reduced basis stays +// non-singular. +func (s *standardForm) expelArtificials(basis []int, inBasic []bool, tol float64) ([]int, error) { + // One unit vector and one solve target for the whole expulsion: each + // round clears the previous round's basis vector and the solve + // overwrites y whole. + for { + r := -1 + for i := range basis { + if basis[i] >= s.nreal { + r = i + break + } + } + if r == -1 { + return basis, nil + } + m := len(s.rows) + unit := growF(s.unit, m) + y := growF(s.y, m) + s.unit, s.y = unit, y + f, err := s.refactor(basis) + if err != nil { + return nil, base.Errf("phase 1: %w while expelling an artificial", err) + } + // Row r of B⁻¹: solve Bᵀ y = e_r, then the row is yᵀ. + clear(unit) + unit[r] = 1 + f.solveT(unit, y) + choice := -1 + for j := range s.nreal { + if inBasic[j] { + continue + } + dot := 0.0 + for i := range m { + dot += y[i] * s.rows[i][j] + } + if math.Abs(dot) > tol { + choice = j + break + } + } + if choice >= 0 { + inBasic[basis[r]] = false + basis[r] = choice + inBasic[choice] = true + continue + } + s.rows = append(s.rows[:r], s.rows[r+1:]...) + s.b = append(s.b[:r], s.b[r+1:]...) + inBasic[basis[r]] = false + basis = append(basis[:r], basis[r+1:]...) + } +} + +// solution reconstructs the point from the final basis and values the +// cost on it. Basic values that round a hair below zero are clamped: +// x ≥ 0 is the contract the caller sees. +func (s *standardForm) solution(basis []int, cost []float64) ([]float64, float64, error) { + m := len(s.rows) + f, err := s.refactor(basis) + if err != nil { + return nil, 0, err + } + xb := growF(s.xb, m) + s.xb = xb + f.solve(s.b, xb) + x := make([]float64, s.nreal) + value := 0.0 + for i, c := range basis { + if c < s.nreal { + v := math.Max(xb[i], 0) + x[c] = v + value += cost[c] * v + } + } + return x, value, nil +} + +// lu holds an LU factorisation with partial pivoting of a small dense +// square matrix: PA = LU with the swaps recorded in piv. The simplex +// refactorises it once per pivot and the active-set solver in qp.go +// factors a KKT system with it per iteration, so the type is shared +// machinery for both. +type lu struct { + n int + a []float64 // row-major, factored in place + piv []int // row swaps in application order +} + +// factorLU factorises the n×n row-major matrix mat into a fresh +// factorisation. A pivot vanishing against the matrix's scale is a +// singular matrix, reported as an error naming the column: for the +// simplex that is a basis no longer invertible, for the KKT system an +// active set that has lost rank. +func factorLU(mat []float64, n int) (*lu, error) { + f := &lu{} + if err := f.factor(mat, n); err != nil { + return nil, err + } + return f, nil +} + +// factor refactorises the receiver on the n×n row-major matrix mat, +// reusing the storage a previous factorisation left behind: the +// simplex's basis and the active-set solver's KKT system are both +// refactorised once per iteration, so one factor per solve replaces one +// per iteration. mat is left untouched; every entry of the workspace is +// overwritten from it, which is what makes the reuse invisible in the +// result. +func (f *lu) factor(mat []float64, n int) error { + if n == 0 { + f.n, f.a, f.piv = 0, f.a[:0], f.piv[:0] + return nil + } + if cap(f.a) < n*n { + f.a = make([]float64, n*n) + } + if cap(f.piv) < n { + f.piv = make([]int, n) + } + f.n, f.a, f.piv = n, f.a[:n*n], f.piv[:n] + a, piv := f.a, f.piv + // The copy and the scale scan are one fused pass: the scale is the + // maximum over the same values either way. + scale := 0.0 + for i, v := range mat[:n*n] { + a[i] = v + if x := math.Abs(v); x > scale { + scale = x + } + } + if scale == 0 { + return base.Errf("the matrix is singular (a zero matrix)") + } + for k := range n { + p, best := k, math.Abs(a[k*n+k]) + for i := k + 1; i < n; i++ { + if v := math.Abs(a[i*n+k]); v > best { + p, best = i, v + } + } + piv[k] = p + if best <= 1e-14*scale { + return base.Errf("the matrix is singular to working precision (pivot %g in column %d)", best, k+1) + } + if p != k { + for j := range n { + a[k*n+j], a[p*n+j] = a[p*n+j], a[k*n+j] + } + } + inv := 1 / a[k*n+k] + for i := k + 1; i < n; i++ { + e := a[i*n+k] * inv + a[i*n+k] = e + if e != 0 { + for j := k + 1; j < n; j++ { + a[i*n+j] -= e * a[k*n+j] + } + } + } + } + return nil +} + +// solve writes A⁻¹ b into x: the recorded swaps forward, then the unit +// lower triangle forward, then the upper triangle back. b is left +// untouched. +func (f *lu) solve(b, x []float64) { + n := f.n + copy(x, b) + for k := range n { + x[k], x[f.piv[k]] = x[f.piv[k]], x[k] + } + for i := 1; i < n; i++ { + s := x[i] + for k := range i { + s -= f.a[i*n+k] * x[k] + } + x[i] = s + } + for i := n - 1; i >= 0; i-- { + s := x[i] + for k := i + 1; k < n; k++ { + s -= f.a[i*n+k] * x[k] + } + x[i] = s / f.a[i*n+i] + } +} + +// solveT writes Aᵀ⁻¹ b into x. With PA = LU the transpose factors as +// Aᵀ = UᵀLᵀP, so the solve runs Uᵀ forward, Lᵀ back and undoes the +// swaps in reverse. The dual prices of the simplex and the redundant +// row scan of the phase transition both come through here. +func (f *lu) solveT(b, x []float64) { + n := f.n + copy(x, b) + for i := range n { // Uᵀ w = b, forward, diagonal uᵢᵢ + s := x[i] + for k := range i { + s -= f.a[k*n+i] * x[k] + } + x[i] = s / f.a[i*n+i] + } + for i := n - 1; i >= 0; i-- { // Lᵀ v = w, back, unit diagonal + s := x[i] + for k := i + 1; k < n; k++ { + s -= f.a[k*n+i] * x[k] + } + x[i] = s + } + for k := n - 1; k >= 0; k-- { // x = Pᵀ v: the swaps in reverse + x[k], x[f.piv[k]] = x[f.piv[k]], x[k] + } +} diff --git a/optim/simplex_test.go b/optim/simplex_test.go new file mode 100644 index 0000000..50ecf00 --- /dev/null +++ b/optim/simplex_test.go @@ -0,0 +1,450 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package optim + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestLUFactorSolve pins the shared factorisation machinery on a +// non-symmetric matrix: both triangle directions are load-bearing, +// solve for the primal ratios and solveT for the dual prices and the +// redundant-row scan. +func TestLUFactorSolve(t *testing.T) { + // A = [[2, 1], [4, 3]], det = 2: A(1, 3) = (5, 13) and + // Aᵀ(2, 1) = (8, 5). + f, err := factorLU([]float64{2, 1, 4, 3}, 2) + if err != nil { + t.Fatalf("factorLU: %v", err) + } + x := make([]float64, 2) + f.solve([]float64{5, 13}, x) + if math.Abs(x[0]-1) > 1e-12 || math.Abs(x[1]-3) > 1e-12 { + t.Fatalf("solve = (%g, %g), want (1, 3)", x[0], x[1]) + } + f.solveT([]float64{8, 5}, x) + if math.Abs(x[0]-2) > 1e-12 || math.Abs(x[1]-1) > 1e-12 { + t.Fatalf("solveT = (%g, %g), want (2, 1)", x[0], x[1]) + } + // A singular matrix is refused, not divided through. + if _, err := factorLU([]float64{1, 2, 2, 4}, 2); err == nil { + t.Fatal("a singular matrix was factored") + } + // So is a zero matrix, named as such rather than a pivot. + if _, err := factorLU(make([]float64, 4), 2); err == nil { + t.Fatal("a zero matrix was factored") + } + if _, err := factorLU(nil, 0); err != nil { + t.Fatalf("a zero-sized factorisation was refused: %v", err) + } +} + +// TestMinimiseLinearStandardVertex pins a hand-built optimum at a +// known vertex: min −(2x + 3y) over x + y ≤ 4, 2x + y ≤ 6 with the +// slacks carried explicitly. The vertices price out at 0, 6, 10 and +// 12, so the answer is the vertex (0, 4, 0, 2) with value −12. +func TestMinimiseLinearStandardVertex(t *testing.T) { + c := mustFloats(t, []float64{-2, -3, 0, 0}) + a := mustFloats(t, []float64{1, 1, 1, 0, 2, 1, 0, 1}, 2, 4) + b := mustFloats(t, []float64{4, 6}) + x, value, err := MinimiseLinear(c, a, b, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinear: %v", err) + } + want := []float64{0, 4, 0, 2} + for i, w := range want { + if math.Abs(x.FloatAt(i)-w) > 1e-9 { + t.Fatalf("x[%d] = %.12g, want %g", i, x.FloatAt(i), w) + } + } + if math.Abs(value+12) > 1e-9 { + t.Fatalf("value = %.12g, want −12", value) + } +} + +// TestMinimiseLinearRowsVertex drives the two-sided wrapper over the +// same polytope expressed as house rows with free variables: max +// x + 2y on x + y ≤ 4, x + 3y ≤ 6, x, y ≥ 0. The vertex prices are 0, +// 4, 5 and 4, so the optimum is the vertex (3, 1) with value 5. +func TestMinimiseLinearRowsVertex(t *testing.T) { + c := mustFloats(t, []float64{-1, -2}) + cons := LinearConstraints{ + A: mustFloats(t, []float64{1, 1, 1, 3, 1, 0, 0, 1}, 4, 2), + Lower: []float64{math.Inf(-1), math.Inf(-1), 0, 0}, + Upper: []float64{4, 6, math.Inf(1), math.Inf(1)}, + } + x, value, err := MinimiseLinearRows(c, cons, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinearRows: %v", err) + } + if math.Abs(x.FloatAt(0)-3) > 1e-9 || math.Abs(x.FloatAt(1)-1) > 1e-9 { + t.Fatalf("x = (%.12g, %.12g), want the vertex (3, 1)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(value+5) > 1e-9 { + t.Fatalf("value = %.12g, want −5", value) + } +} + +// TestMinimiseLinearNegativeRightHandSide pins the row negation: a +// standard-form row arrives with b < 0 and must come out of the +// artificial phase feasible all the same. +func TestMinimiseLinearNegativeRightHandSide(t *testing.T) { + c := mustFloats(t, []float64{1, 1}) + a := mustFloats(t, []float64{-1, -1}, 1, 2) + b := mustFloats(t, []float64{-2}) + x, value, err := MinimiseLinear(c, a, b, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinear: %v", err) + } + if math.Abs(value-2) > 1e-9 { + t.Fatalf("value = %.12g, want 2", value) + } + if math.Abs(x.FloatAt(0)+x.FloatAt(1)-2) > 1e-9 { + t.Fatalf("x + y = %.12g, want 2", x.FloatAt(0)+x.FloatAt(1)) + } +} + +// TestMinimiseLinearBudget pins the pivot-budget refusal: one pivot +// cannot carry the vertex problem to its optimum, and an exhausted +// budget is an error, never a silent answer. +func TestMinimiseLinearBudget(t *testing.T) { + c := mustFloats(t, []float64{-2, -3, 0, 0}) + a := mustFloats(t, []float64{1, 1, 1, 0, 2, 1, 0, 1}, 2, 4) + b := mustFloats(t, []float64{4, 6}) + _, _, err := MinimiseLinear(c, a, b, LinearProgramOptions{MaxIterations: 1}) + if err == nil || !strings.Contains(err.Error(), "budget") { + t.Fatalf("error = %v, want the pivot-budget refusal", err) + } +} + +// TestMinimiseLinearDegenerateTerminates pins the anti-cycling +// guarantee where it earns its keep: two identical equality rows plus +// a third at twice the scale leave the problem degenerate and +// redundant at once, the classic rules can pivot forever on such a +// vertex, and Bland's rule must terminate, dropping the redundant rows +// and their artificials, with a point on the feasible segment. +func TestMinimiseLinearDegenerateTerminates(t *testing.T) { + c := mustFloats(t, []float64{-1, -1}) + cons := LinearConstraints{ + A: mustFloats(t, []float64{1, 1, 1, 1, 2, 2}, 3, 2), + Lower: []float64{1, 1, 2}, + Upper: []float64{1, 1, 2}, + } + x, value, err := MinimiseLinearRows(c, cons, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinearRows on a degenerate problem: %v", err) + } + if math.Abs(value+1) > 1e-9 { + t.Fatalf("value = %.12g, want −1", value) + } + if x.FloatAt(0) < -1e-9 || x.FloatAt(1) < -1e-9 { + t.Fatalf("x = (%g, %g) left the non-negative quadrant", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(x.FloatAt(0)+x.FloatAt(1)-1) > 1e-9 { + t.Fatalf("x + y = %.12g, want 1", x.FloatAt(0)+x.FloatAt(1)) + } +} + +// TestMinimiseLinearDegenerateOrigin pins the other degenerate shape: +// a zero right-hand side keeps the artificials basic at the phase-1 +// optimum, so the phase transition must pivot them out with degenerate +// steps before phase 2. The feasible set of x + y = 0, x − y = 0 over +// x, y ≥ 0 is the origin alone. +func TestMinimiseLinearDegenerateOrigin(t *testing.T) { + c := mustFloats(t, []float64{-1, -1}) + a := mustFloats(t, []float64{1, 1, 1, -1}, 2, 2) + b := mustFloats(t, []float64{0, 0}) + x, value, err := MinimiseLinear(c, a, b, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinear: %v", err) + } + if value != 0 { + t.Fatalf("value = %g, want 0", value) + } + for i := range 2 { + if x.FloatAt(i) != 0 { + t.Fatalf("x[%d] = %g, want 0", i, x.FloatAt(i)) + } + } +} + +// TestMinimiseLinearBeale pins Beale's cycling example, the classic +// demonstration that Dantzig's rule can pivot forever (Beale, 1955; +// the form in Chvátal's Linear Programming, chapter 3). The optimum is +// 0.05 at (1/25, 0, 1, 0) with the first slack carrying the 0.03 the +// first row leaves loose, so minimising the negated objective returns +// −0.05 there. +func TestMinimiseLinearBeale(t *testing.T) { + c := mustFloats(t, []float64{-0.75, 150, -0.02, 6, 0, 0, 0}) + a := mustFloats(t, []float64{ + 0.25, -60, -0.04, 9, 1, 0, 0, + 0.5, -90, -0.02, 3, 0, 1, 0, + 0, 0, 1, 0, 0, 0, 1, + }, 3, 7) + b := mustFloats(t, []float64{0, 0, 1}) + x, value, err := MinimiseLinear(c, a, b, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinear on Beale's example: %v", err) + } + want := []float64{1.0 / 25.0, 0, 1, 0, 0.03, 0, 0} + for i, w := range want { + if math.Abs(x.FloatAt(i)-w) > 1e-9 { + t.Fatalf("x[%d] = %.12g, want %.12g", i, x.FloatAt(i), w) + } + } + if math.Abs(value+0.05) > 1e-9 { + t.Fatalf("value = %.12g, want −0.05", value) + } +} + +// TestMinimiseLinearInfeasible pins the phase-1 refusal: x ≥ 1 and +// x ≤ 0 share no feasible point, and the error must carry the +// infeasibility the artificial phase ended with. +func TestMinimiseLinearInfeasible(t *testing.T) { + c := mustFloats(t, []float64{1}) + cons := LinearConstraints{ + A: mustFloats(t, []float64{1, 1}, 2, 1), + Lower: []float64{math.Inf(-1), 1}, + Upper: []float64{0, math.Inf(1)}, + } + _, _, err := MinimiseLinearRows(c, cons, LinearProgramOptions{}) + if err == nil { + t.Fatal("an infeasible problem returned a solution") + } + if !strings.Contains(err.Error(), "infeasible") { + t.Fatalf("error = %v, want the phase-1 infeasibility evidence", err) + } + // The named row must be the one that carries the worst figure, not + // whichever artificial the scan touched after row 1: two empty rows + // leave their artificials at 1 and 0.5, and the zero-initialised + // worst-row sentinel once let the smaller overwrite the evidence. + _, _, err = MinimiseLinear(mustFloats(t, []float64{0}), + mustFloats(t, []float64{0, 0}, 2, 1), mustFloats(t, []float64{1, 0.5}, 2), LinearProgramOptions{}) + if err == nil || !strings.Contains(err.Error(), "row 1 still carries 1") { + t.Fatalf("worst row misreported: %v", err) + } +} + +// TestMinimiseLinearUnbounded pins the ray refusal in both entries. +func TestMinimiseLinearUnbounded(t *testing.T) { + // Standard form: x = y with min −x runs along the ray (t, t). + _, _, err := MinimiseLinear(mustFloats(t, []float64{-1, 0}), + mustFloats(t, []float64{1, -1}, 1, 2), mustFloats(t, []float64{0}), LinearProgramOptions{}) + if err == nil || !strings.Contains(err.Error(), "unbounded") { + t.Fatalf("standard form: error = %v, want an unbounded refusal", err) + } + // Rows: x ≥ 0 with min −x. + _, _, err = MinimiseLinearRows(mustFloats(t, []float64{-1}), + LinearConstraints{A: mustFloats(t, []float64{1}, 1, 1), + Lower: []float64{0}, Upper: []float64{math.Inf(1)}}, LinearProgramOptions{}) + if err == nil || !strings.Contains(err.Error(), "unbounded") { + t.Fatalf("rows: error = %v, want an unbounded refusal", err) + } +} + +// TestMinimiseLinearNoRows pins the row-free standard form: over x ≥ 0 +// alone a non-negative cost bottoms out at the origin and a negative +// cost is unbounded. +func TestMinimiseLinearNoRows(t *testing.T) { + x, value, err := MinimiseLinear(mustFloats(t, []float64{1, 2}), core.New(core.Float, 0, 2), + core.New(core.Float, 0), LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinear: %v", err) + } + if value != 0 { + t.Fatalf("value = %g, want 0", value) + } + for i := range 2 { + if x.FloatAt(i) != 0 { + t.Fatalf("x[%d] = %g, want 0", i, x.FloatAt(i)) + } + } + if _, _, err := MinimiseLinear(mustFloats(t, []float64{-1, 2}), core.New(core.Float, 0, 2), + core.New(core.Float, 0), LinearProgramOptions{}); err == nil { + t.Fatal("a negative cost without rows was not refused as unbounded") + } +} + +// TestMinimiseLinearRefusals checks every input gate of both entries. +func TestMinimiseLinearRefusals(t *testing.T) { + empty := core.New(core.Float, 0) + c2 := mustFloats(t, []float64{1, 2}) + a22 := mustFloats(t, []float64{1, 1, 0, 1}, 2, 2) + b2 := mustFloats(t, []float64{1, 1}) + cases := []struct { + name string + run func() error + }{ + {"empty cost", func() error { + _, _, err := MinimiseLinear(empty, a22, b2, LinearProgramOptions{}) + return err + }}, + {"complex cost", func() error { + _, _, err := MinimiseLinear(mustComplexPoint(t), a22, b2, LinearProgramOptions{}) + return err + }}, + {"nil matrix", func() error { + _, _, err := MinimiseLinear(c2, nil, b2, LinearProgramOptions{}) + return err + }}, + {"matrix shape", func() error { + _, _, err := MinimiseLinear(c2, mustFloats(t, []float64{1, 1}), b2, LinearProgramOptions{}) + return err + }}, + {"right-hand side length", func() error { + _, _, err := MinimiseLinear(c2, a22, mustFloats(t, []float64{1}), LinearProgramOptions{}) + return err + }}, + {"NaN coefficient", func() error { + _, _, err := MinimiseLinear(c2, mustFloats(t, []float64{math.NaN(), 1, 0, 1}, 2, 2), b2, LinearProgramOptions{}) + return err + }}, + {"NaN cost", func() error { + _, _, err := MinimiseLinear(mustFloats(t, []float64{math.NaN(), 1}), a22, b2, LinearProgramOptions{}) + return err + }}, + {"complex constraint matrix", func() error { + ca, _ := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + _, _, err := MinimiseLinear(c2, ca, b2, LinearProgramOptions{}) + return err + }}, + {"complex right-hand side", func() error { + cb, _ := core.FromComplexes([]complex128{1, 1}, 2) + _, _, err := MinimiseLinear(c2, a22, cb, LinearProgramOptions{}) + return err + }}, + {"NaN right-hand side", func() error { + _, _, err := MinimiseLinear(c2, a22, mustFloats(t, []float64{math.NaN(), 1}), LinearProgramOptions{}) + return err + }}, + } + for _, c := range cases { + if err := c.run(); err == nil { + t.Fatalf("MinimiseLinear: %s accepted", c.name) + } + } + rows := []struct { + name string + cons LinearConstraints + }{ + {"nil matrix", LinearConstraints{Lower: []float64{0}, Upper: []float64{1}}}, + {"matrix shape", LinearConstraints{A: mustFloats(t, []float64{1, 1}), Lower: []float64{0}, Upper: []float64{1}}}, + {"no rows", LinearConstraints{A: core.New(core.Float, 0, 2), Lower: nil, Upper: nil}}, + {"short bounds", LinearConstraints{A: a22, Lower: []float64{0}, Upper: []float64{}}}, + {"crossed bounds", LinearConstraints{A: a22, Lower: []float64{1, 0}, Upper: []float64{0, 1}}}, + {"equality at infinity", LinearConstraints{A: a22, + Lower: []float64{math.Inf(1), 0}, Upper: []float64{math.Inf(1), 1}}}, + {"NaN bound", LinearConstraints{A: a22, Lower: []float64{math.NaN(), 0}, Upper: []float64{1, 1}}}, + {"NaN coefficient", LinearConstraints{A: mustFloats(t, []float64{math.NaN(), 1, 0, 1}, 2, 2), + Lower: []float64{0, 0}, Upper: []float64{1, 1}}}, + } + for _, r := range rows { + if _, _, err := MinimiseLinearRows(c2, r.cons, LinearProgramOptions{}); err == nil { + t.Fatalf("MinimiseLinearRows: %s accepted", r.name) + } + } + // The rows entry's own cost gates. + rowCosts := []struct { + name string + c *core.Array + }{ + {"empty cost", core.New(core.Float, 0)}, + {"complex cost", mustComplexPoint(t)}, + {"NaN cost", mustFloats(t, []float64{math.NaN(), 1})}, + } + for _, rc := range rowCosts { + if _, _, err := MinimiseLinearRows(rc.c, LinearConstraints{A: a22, Lower: []float64{0, 0}, Upper: []float64{1, 1}}, + LinearProgramOptions{}); err == nil { + t.Fatalf("MinimiseLinearRows: %s accepted", rc.name) + } + } + // Complex constraint matrices are refused here too. + ca, _ := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + if _, _, err := MinimiseLinearRows(c2, LinearConstraints{A: ca, Lower: []float64{0, 0}, Upper: []float64{1, 1}}, + LinearProgramOptions{}); err == nil { + t.Fatal("MinimiseLinearRows: a complex constraint matrix was accepted") + } +} + +// TestMinimiseLinearRowsNegativeBounds pins the negation of a +// right-hand side the artificial basis needs: a two-sided model whose +// bound rows come out negative (a lower bound below zero, a bare +// negative equality) builds standard rows with b >= 0 and solves, +// where an unnegated row lost primal feasibility at pivot 0. +func TestMinimiseLinearRowsNegativeBounds(t *testing.T) { + a, err := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2) + if err != nil { + t.Fatal(err) + } + cons := LinearConstraints{ + A: a, + Lower: []float64{-1, 2}, + Upper: []float64{math.Inf(1), math.Inf(1)}, + } + x, v, err := MinimiseLinearRows(mustFloats(t, []float64{1, 1}), cons, LinearProgramOptions{}) + if err != nil { + t.Fatalf("negative lower bound refused: %v", err) + } + if v != 1 || x.FloatAt(0) != -1 || x.FloatAt(1) != 2 { + t.Fatalf("x = (%g, %g) value %g, want (-1, 2) at 1", x.FloatAt(0), x.FloatAt(1), v) + } + // The bare negative equality rides the same negation through the + // QP entry's feasibility path. + h, herr := core.FromFloats([]float64{2, 0, 0, 2}, 2, 2) + if herr != nil { + t.Fatal(herr) + } + eq, eerr := core.FromFloats([]float64{1, 1}, 1, 2) + if eerr != nil { + t.Fatal(eerr) + } + eqCons := LinearConstraints{A: eq, Lower: []float64{-1}, Upper: []float64{-1}} + q, qv, _, qerr := MinimiseQP(h, mustFloats(t, []float64{4, 0}), eqCons, nil, QPOptions{}) + if qerr != nil { + t.Fatalf("negative equality refused: %v", qerr) + } + if math.Abs(q.FloatAt(0)-(-1.5)) > 1e-9 || math.Abs(q.FloatAt(1)-0.5) > 1e-9 || math.Abs(qv+3.5) > 1e-9 { + t.Fatalf("x = (%g, %g) value %g, want (-1.5, 0.5) at -3.5", q.FloatAt(0), q.FloatAt(1), qv) + } +} + +// TestMinimiseLinearExpelsTwoArtificials pins the reused unit vector in +// the artificial expulsion. The two rows below are exact negatives, so +// the phase-1 reduced costs cancel for every column: the artificial +// phase ends with both artificials still basic at zero and the +// expulsion runs twice, the second round reading a freshly cleared unit +// vector. A stale one scores the columns through the sum of two rows of +// B⁻¹ and swaps in a column that leaves the basis singular, so a solve +// that must answer the origin fails instead. +func TestMinimiseLinearExpelsTwoArtificials(t *testing.T) { + a := mustFloats(t, []float64{1, 2, -1, -2}, 2, 2) + b := mustFloats(t, []float64{0, 0}) + c := mustFloats(t, []float64{1, 1}) + x, value, err := MinimiseLinear(c, a, b, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinear: %v", err) + } + if value != 0 { + t.Fatalf("value = %.12g, want 0", value) + } + for i := range x.Len() { + if x.FloatAt(i) != 0 { + t.Fatalf("x[%d] = %.12g, want 0", i, x.FloatAt(i)) + } + } + // The same two-round expulsion over a scaled right-hand side: the + // origin is the only feasible point whatever the rows' scale. + scaled := mustFloats(t, []float64{3, 6, -3, -6}, 2, 2) + x, value, err = MinimiseLinear(c, scaled, b, LinearProgramOptions{}) + if err != nil { + t.Fatalf("MinimiseLinear over the scaled rows: %v", err) + } + if value != 0 || x.FloatAt(0) != 0 || x.FloatAt(1) != 0 { + t.Fatalf("scaled rows: x = (%g, %g), value %g, want the origin", x.FloatAt(0), x.FloatAt(1), value) + } +} diff --git a/optim/solvercoverage_test.go b/optim/solvercoverage_test.go new file mode 100644 index 0000000..de56eac --- /dev/null +++ b/optim/solvercoverage_test.go @@ -0,0 +1,321 @@ +package optim + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Solver-coverage pins: every test drives an optimiser on a problem +// class with a closed-form answer the existing fixtures do not build, +// from corner optima over degenerate active sets to mixed nonlinear +// constraints. + +func mustMatrix(t *testing.T, vals []float64, r, c int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, r, c) + if err != nil { + t.Fatal(err) + } + return a +} + +// TestLBFGSCornerOptimumWithEqualityBound minimises +// f = (x-3)^2 + (y+1)^2 + (z-2)^2 over x <= 1 with z frozen at 2: the +// constrained optimum (1, -1, 2) keeps the wall term at 4, so 4 is the +// answer the projected-gradient machinery must report. +func TestLBFGSCornerOptimumWithEqualityBound(t *testing.T) { + f := func(a *core.Array) (float64, error) { + x, y, z := a.FloatAt(0), a.FloatAt(1), a.FloatAt(2) + return (x-3)*(x-3) + (y+1)*(y+1) + (z-2)*(z-2), nil + } + grad := func(a *core.Array) (*core.Array, error) { + x, y, z := a.FloatAt(0), a.FloatAt(1), a.FloatAt(2) + return core.FromFloats([]float64{2 * (x - 3), 2 * (y + 1), 2 * (z - 2)}, 3) + } + x0, _ := core.FromFloats([]float64{0, 0, 2}, 3) + opts := LBFGSOptions{ + Tolerance: 1e-12, + Lower: []float64{math.Inf(-1), math.Inf(-1), 2}, + Upper: []float64{1, math.Inf(1), 2}, + } + p, fv, err := MinimiseLBFGS(f, grad, x0, opts) + if err != nil { + t.Fatal(err) + } + if math.Abs(fv-4) > 1e-10 { + t.Fatalf("f = %g, want 4", fv) + } + if math.Abs(p.FloatAt(0)-1) > 1e-9 || math.Abs(p.FloatAt(1)+1) > 1e-9 || p.FloatAt(2) != 2 { + t.Fatalf("point (%g, %g, %g), want (1, -1, 2)", p.FloatAt(0), p.FloatAt(1), p.FloatAt(2)) + } +} + +// TestLBFGSCornerOptimumFiniteDifference repeats the corner problem +// without a gradient callback, so the one-sided difference stencils at +// the walls carry the run. +func TestLBFGSCornerOptimumFiniteDifference(t *testing.T) { + f := func(a *core.Array) (float64, error) { + x, y := a.FloatAt(0), a.FloatAt(1) + return (x-3)*(x-3) + (y+1)*(y+1), nil + } + x0, _ := core.FromFloats([]float64{0, 0}, 2) + opts := LBFGSOptions{ + Tolerance: 1e-10, + Lower: []float64{-2, math.Inf(-1)}, + Upper: []float64{1, math.Inf(1)}, + MaxIterations: 500, + } + p, fv, err := MinimiseLBFGS(f, nil, x0, opts) + if err != nil { + t.Fatal(err) + } + if math.Abs(fv-4) > 1e-8 { + t.Fatalf("f = %g, want 4", fv) + } + if math.Abs(p.FloatAt(0)-1) > 1e-5 || math.Abs(p.FloatAt(1)+1) > 1e-5 { + t.Fatalf("point (%g, %g), want (1, -1)", p.FloatAt(0), p.FloatAt(1)) + } +} + +// TestQPDegenerateRatioTies builds a problem whose ratio test ties +// three rows at one point and whose release cycle must then free a +// slack row: min 1/2(x^2+y^2) - x - y over x+y >= 1, x >= 1/2, +// y >= 1/2. All rows block at (1/2, 1/2); the optimum is (1, 1) with +// the first row slack and a zero multiplier. +func TestQPDegenerateRatioTies(t *testing.T) { + h, _ := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2) + c, _ := core.FromFloats([]float64{-1, -1}, 2) + cons := LinearConstraints{ + A: mustMatrix(t, []float64{1, 1, 1, 0, 0, 1}, 3, 2), + Lower: []float64{1, 0.5, 0.5}, + Upper: []float64{math.Inf(1), math.Inf(1), math.Inf(1)}, + } + x, fv, multipliers, err := MinimiseQP(h, c, cons, nil, QPOptions{}) + if err != nil { + t.Fatal(err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-8 || math.Abs(x.FloatAt(1)-1) > 1e-8 { + t.Fatalf("point (%g, %g), want (1, 1)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(fv+1) > 1e-10 { + t.Fatalf("f = %g, want -1", fv) + } + if len(multipliers) != 3 { + t.Fatalf("multipliers %v", multipliers) + } + if multipliers[0] > 1e-10 { + t.Fatalf("slack row multiplier %g, want 0", multipliers[0]) + } +} + +// TestCMAESRotatedIllConditionedQuadratic runs the strategy on a valley +// rotated 45 degrees with a conditioning of 100, where a diagonal +// sampler cannot turn and the full covariance must. +func TestCMAESRotatedIllConditionedQuadratic(t *testing.T) { + f := func(a *core.Array) (float64, error) { + x, y := a.FloatAt(0), a.FloatAt(1) + u, v := (x+y)/math.Sqrt2, (x-y)/math.Sqrt2 + return u*u + 100*v*v, nil + } + x0, _ := core.FromFloats([]float64{5, -5}, 2) + p, fv, err := MinimiseCMAES(f, x0, CMAESOptions{Generations: 2000, Tolerance: 1e-10}) + if err != nil { + t.Fatal(err) + } + if fv > 1e-8 { + t.Fatalf("f = %g, want ~0", fv) + } + if math.Abs(p.FloatAt(0)) > 1e-3 || math.Abs(p.FloatAt(1)) > 1e-3 { + t.Fatalf("point (%g, %g), want ~(0, 0)", p.FloatAt(0), p.FloatAt(1)) + } +} + +// TestDifferentialEvolutionBowl drives the rand/1/bin scheme on a +// separable bowl to full precision against the seeded bounds. +func TestDifferentialEvolutionBowl(t *testing.T) { + f := func(a *core.Array) (float64, error) { + s := 0.0 + for i := range a.Len() { + d := a.FloatAt(i) - float64(i+1) + s += d * d + } + return s, nil + } + lower, _ := core.FromFloats([]float64{-10, -10, -10}, 3) + upper, _ := core.FromFloats([]float64{10, 10, 10}, 3) + p, fv, err := MinimiseDifferentialEvolution(f, lower, upper, DifferentialEvolutionOptions{ + Generations: 600, Seed: 7, + }) + if err != nil { + t.Fatal(err) + } + if fv > 1e-10 { + t.Fatalf("f = %g, want ~0", fv) + } + for i := range 3 { + if math.Abs(p.FloatAt(i)-float64(i+1)) > 1e-4 { + t.Fatalf("point[%d] = %g, want %g", i, p.FloatAt(i), float64(i+1)) + } + } +} + +// TestLinearRowsRedundantEqualityExpelled adds a third equality that is +// the sum of the first two, so phase 1 must expel a redundant +// artificial through the row-drop path before phase 2 prices. +func TestLinearRowsRedundantEqualityExpelled(t *testing.T) { + c, _ := core.FromFloats([]float64{1, 1}, 2) + cons := LinearConstraints{ + A: mustMatrix(t, []float64{1, 1, 1, -1, 2, 0}, 3, 2), + Lower: []float64{2, 0, 2}, + Upper: []float64{2, 0, 2}, + } + x, fv, err := MinimiseLinearRows(c, cons, LinearProgramOptions{}) + if err != nil { + t.Fatal(err) + } + if math.Abs(x.FloatAt(0)-1) > 1e-7 || math.Abs(x.FloatAt(1)-1) > 1e-7 { + t.Fatalf("point (%g, %g), want (1, 1)", x.FloatAt(0), x.FloatAt(1)) + } + if math.Abs(fv-2) > 1e-7 { + t.Fatalf("value %g, want 2", fv) + } +} + +// TestSimplexRosenbrock walks the derivative-free simplex down the +// banana valley to its floor at (1, 1). +func TestSimplexRosenbrock(t *testing.T) { + f := func(a *core.Array) (float64, error) { + x, y := a.FloatAt(0), a.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + } + x0, _ := core.FromFloats([]float64{-1.2, 1}, 2) + p, fv, err := Minimise(f, x0, MinimiseOptions{MaxIterations: 20000, Tolerance: 1e-12}) + if err != nil { + t.Fatal(err) + } + if fv > 1e-12 { + t.Fatalf("f = %g, want ~0", fv) + } + if math.Abs(p.FloatAt(0)-1) > 1e-4 || math.Abs(p.FloatAt(1)-1) > 1e-4 { + t.Fatalf("point (%g, %g), want (1, 1)", p.FloatAt(0), p.FloatAt(1)) + } +} + +// TestRootSystemBroyden solves a mildly nonlinear system with the +// maintained inverse and verifies both equations at the answer. +func TestRootSystemBroyden(t *testing.T) { + f := func(x *core.Array) (*core.Array, error) { + a, b := x.FloatAt(0), x.FloatAt(1) + return core.FromFloats([]float64{ + a*a + b*b - 4, + math.Exp(a) + b - 3, + }, 2) + } + x0, _ := core.FromFloats([]float64{1, 1}, 2) + x, res, err := FindRootSystem(f, x0, RootSystemOptions{Tolerance: 1e-11, UseBroyden: true, MaxIterations: 200}) + if err != nil { + t.Fatal(err) + } + if res > 1e-11 { + t.Fatalf("residual %g", res) + } + a, b := x.FloatAt(0), x.FloatAt(1) + if math.Abs(a*a+b*b-4) > 1e-9 || math.Abs(math.Exp(a)+b-3) > 1e-9 { + t.Fatalf("solution (%g, %g) does not satisfy the system", a, b) + } +} + +// TestLevenbergMarquardtWeightedCovariance fits a line under per-point +// variances and checks the parameters and the reported covariance. +func TestLevenbergMarquardtSigmaWeightsAndCovariance(t *testing.T) { + xs := []float64{0, 1, 2, 3, 4} + ys := []float64{0.5, 2.49, 4.52, 6.48, 8.51} + residual := func(p *core.Array) (*core.Array, error) { + out := core.New(core.Float, len(xs)) + v := out.RawFloats() + for i, x := range xs { + v[i] = p.FloatAt(0)*x + p.FloatAt(1) - ys[i] + } + return out, nil + } + sigma, _ := core.FromFloats([]float64{1, 1, 4, 1, 4}, 5) + p0, _ := core.FromFloats([]float64{0, 0}, 2) + res, err := LevenbergMarquardtFit(residual, p0, LMOptions{Sigma: sigma, RequestCovariance: true}) + if err != nil { + t.Fatal(err) + } + if res.Status != FitConverged { + t.Fatalf("status %v", res.Status) + } + if math.Abs(res.Parameters.FloatAt(0)-2) > 0.05 || math.Abs(res.Parameters.FloatAt(1)-0.5) > 0.05 { + t.Fatalf("parameters (%g, %g), want ~(2, 0.5)", + res.Parameters.FloatAt(0), res.Parameters.FloatAt(1)) + } + if res.Covariance == nil { + t.Fatal("covariance missing") + } +} + +// TestNonlinearConstrainedMixedRows minimises a bowl over an equality +// and an inequality at once, with the inequality active at the answer: +// min (x-2)^2 + (y-2)^2 over x + y = 1 and x <= 1/2 sits at +// (1/2, 1/2) with f = 4.5 and a non-negative inequality multiplier. +func TestNonlinearConstrainedMixedRows(t *testing.T) { + f := func(a *core.Array) (float64, error) { + x, y := a.FloatAt(0), a.FloatAt(1) + return (x-2)*(x-2) + (y-2)*(y-2), nil + } + grad := func(a *core.Array) (*core.Array, error) { + x, y := a.FloatAt(0), a.FloatAt(1) + return core.FromFloats([]float64{2 * (x - 2), 2 * (y - 2)}, 2) + } + cons := NonlinearConstraints{ + Equalities: []func(*core.Array) (float64, error){func(a *core.Array) (float64, error) { + return a.FloatAt(0) + a.FloatAt(1) - 1, nil + }}, + Inequalities: []func(*core.Array) (float64, error){func(a *core.Array) (float64, error) { + return a.FloatAt(0) - 0.5, nil + }}, + } + x0, _ := core.FromFloats([]float64{0, 0}, 2) + p, fv, mult, err := MinimiseNonlinearConstrained(f, grad, x0, cons, LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatal(err) + } + if math.Abs(p.FloatAt(0)-0.5) > 1e-4 || math.Abs(p.FloatAt(1)-0.5) > 1e-4 { + t.Fatalf("point (%g, %g), want (0.5, 0.5)", p.FloatAt(0), p.FloatAt(1)) + } + if fv < 4.49 || fv > 4.51 { + t.Fatalf("f = %g, want 4.5", fv) + } + if len(mult) != 2 || mult[1] < 0 { + t.Fatalf("multipliers %v", mult) + } +} + +// TestSimulatedAnnealingTwoWell puts a deep left well against a +// shallower right one, starting in the right: the Metropolis walk must +// cross the barrier and report the global basin. +func TestSimulatedAnnealingTwoWell(t *testing.T) { + f := func(a *core.Array) (float64, error) { + x := a.FloatAt(0) + w1 := (x + 2) * (x + 2) + w2 := (x-3)*(x-3) + 0.5 + return math.Min(w1, w2), nil + } + x0, _ := core.FromFloats([]float64{3}, 1) + p, fv, err := MinimiseSimulatedAnnealing(f, x0, SimulatedAnnealingOptions{ + Steps: 40000, Seed: 5, StepScale: 0.5, AllowBudgetExit: true, + }) + if err != nil { + t.Fatal(err) + } + if fv > 0.1 { + t.Fatalf("f = %g, want the left well (~0)", fv) + } + if p.FloatAt(0) > 0 { + t.Fatalf("point %g, want the left well", p.FloatAt(0)) + } +} diff --git a/oracle_norace_test.go b/oracle_norace_test.go new file mode 100644 index 0000000..e193746 --- /dev/null +++ b/oracle_norace_test.go @@ -0,0 +1,10 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +//go:build !race + +package tensor + +// oracleRaceSuffix is empty without the race detector; see +// oracle_race_test.go for why a race build carries its own digests. +const oracleRaceSuffix = "" diff --git a/oracle_race_test.go b/oracle_race_test.go new file mode 100644 index 0000000..39f9d7d --- /dev/null +++ b/oracle_race_test.go @@ -0,0 +1,13 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +//go:build race + +package tensor + +// oracleRaceSuffix marks the digest block of a race-instrumented build. +// The instrumentation changes what the compiler contracts into a fused +// multiply-add, so a race build rounds differently from an otherwise +// identical plain build: it is a configuration of its own, recorded +// beside the GOAMD64 level rather than mixed into it. +const oracleRaceSuffix = "+race" diff --git a/oracle_test.go b/oracle_test.go new file mode 100644 index 0000000..7d82a9e --- /dev/null +++ b/oracle_test.go @@ -0,0 +1,2852 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package tensor + +import ( + "crypto/sha256" + "encoding/binary" + "encoding/hex" + "fmt" + "hash" + "maps" + "math" + "os" + "path/filepath" + "runtime" + "runtime/debug" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/spmd" +) + +// The oracle harness. Every domain runs one fixed workload through the +// re-exported facade only, the raw bits of every output are hashed, and +// the per-case digest is pinned in the table below. A silent behaviour +// change, a swapped reduction order, a regression in a kernel: any of +// them moves the digest and fails the test. A deliberate change to an +// algorithm re-records the table in the same commit, so the diff shows +// exactly which case moved. +// +// Inputs are fixed literals or the seeded generator, which is +// bit-stable across Go releases by construction, and the parallel +// kernels promise a fixed reduction order, so the digest is +// deterministic on a given architecture. Run with +// TENSOR_ORACLE_RECORD=1 to print the current digests instead of +// comparing them. + +// oracleDigest accumulates the raw bits of the given values. Supported +// carriers: *Array, Scalar, float64, []float64, string, bool and int; +// anything else is a harness bug and panics. +type oracleDigest struct { + h hash.Hash +} + +func newOracleDigest() *oracleDigest { return &oracleDigest{h: sha256.New()} } + +func (d *oracleDigest) array(t *testing.T, a *core.Array) { + d.h.Write([]byte{byte(a.Dtype())}) + for _, s := range a.Shape() { + var buf [8]byte + binary.LittleEndian.PutUint64(buf[:], uint64(s)) + d.h.Write(buf[:]) + } + switch a.Dtype() { + case core.Int: + for _, v := range a.RawInts()[:a.Len()] { + var buf [8]byte + binary.LittleEndian.PutUint64(buf[:], uint64(v)) + d.h.Write(buf[:]) + } + case core.Float32: + for _, v := range a.RawFloat32s()[:a.Len()] { + var buf [4]byte + binary.LittleEndian.PutUint32(buf[:], math.Float32bits(v)) + d.h.Write(buf[:]) + } + case core.Float16: + for _, v := range a.RawHalves()[:a.Len()] { + var buf [2]byte + binary.LittleEndian.PutUint16(buf[:], v) + d.h.Write(buf[:]) + } + case core.Float: + for _, v := range a.RawFloats()[:a.Len()] { + d.f64(v) + } + case core.Complex: + for _, v := range a.RawComplexes()[:a.Len()] { + d.f64(real(v)) + d.f64(imag(v)) + } + case core.Bool: + // The boolean payload hashes as its own bytes, one 0 or 1 per + // element, the layout every mask round-trip is read back in. + for _, v := range a.RawBools()[:a.Len()] { + if v { + d.h.Write([]byte{1}) + } else { + d.h.Write([]byte{0}) + } + } + case core.Int8: + for _, v := range a.RawInt8s()[:a.Len()] { + d.h.Write([]byte{byte(v)}) + } + case core.Uint8: + for _, v := range a.RawUint8s()[:a.Len()] { + d.h.Write([]byte{byte(v)}) + } + case core.Int16: + for _, v := range a.RawInt16s()[:a.Len()] { + var buf [2]byte + binary.LittleEndian.PutUint16(buf[:], uint16(v)) + d.h.Write(buf[:]) + } + case core.Uint16: + for _, v := range a.RawUint16s()[:a.Len()] { + var buf [2]byte + binary.LittleEndian.PutUint16(buf[:], uint16(v)) + d.h.Write(buf[:]) + } + case core.Int32: + for _, v := range a.RawInt32s()[:a.Len()] { + var buf [4]byte + binary.LittleEndian.PutUint32(buf[:], uint32(v)) + d.h.Write(buf[:]) + } + case core.Uint32: + for _, v := range a.RawUint32s()[:a.Len()] { + var buf [4]byte + binary.LittleEndian.PutUint32(buf[:], uint32(v)) + d.h.Write(buf[:]) + } + default: + // An unhashed payload would pin nothing: fail loudly in the + // test rather than panicking or hashing the wrong bytes. + t.Fatalf("oracle harness: array payload of dtype %s (%d) has no hash rule", a.Dtype(), a.Dtype()) + } +} + +func (d *oracleDigest) f64(v float64) { + var buf [8]byte + binary.LittleEndian.PutUint64(buf[:], math.Float64bits(v)) + d.h.Write(buf[:]) +} + +func (d *oracleDigest) floats(vals []float64) { + var n [8]byte + binary.LittleEndian.PutUint64(n[:], uint64(len(vals))) + d.h.Write(n[:]) + for _, v := range vals { + d.f64(v) + } +} + +func (d *oracleDigest) int(v int) { + var buf [8]byte + binary.LittleEndian.PutUint64(buf[:], uint64(v)) + d.h.Write(buf[:]) +} + +// result carries one output of an oracle case into the digest. +func (d *oracleDigest) result(t *testing.T, v any) { + switch x := v.(type) { + case *Array: + d.array(t, x) + case Scalar: + d.f64(x.Float()) + case float64: + d.f64(x) + case []float64: + d.floats(x) + case string: + d.h.Write([]byte(x)) + case bool: + // Error paths carry the fact of the refusal, not its wording: + // the text may be reworded, the refusal must not disappear. + if x { + d.h.Write([]byte{1}) + } else { + d.h.Write([]byte{0}) + } + case int: + d.int(x) + default: + panic("oracle harness: unsupported carrier type") + } +} + +func (d *oracleDigest) sum() string { return hex.EncodeToString(d.h.Sum(nil)) } + +// mustA builds an array from fixed values. +func mustA(t *testing.T, vals []float64, shape ...int) *Array { + t.Helper() + a, err := FromFloats(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +// randA draws n deterministic samples through the seeded generator. +func randA(t *testing.T, seed int64, n int) *Array { + t.Helper() + g := NewGenerator(seed) + a, err := Floats(g, n) + if err != nil { + t.Fatal(err) + } + return a +} + +// wobble is the bounded deterministic wiggle the formula-driven +// workloads use in place of a generator: a fixed function of the +// index, so a case's inputs are its own, not the generator state's. +func wobble(i int) float64 { + return 0.3*math.Sin(7.3*float64(i)+1.1)*math.Cos(2.1*float64(i)) + + 0.1*math.Sin(0.7*float64(i)) +} + +// oracleCases is the harness workload: one entry per domain feature, +// every output it produces fed into the digest. +var oracleCases = []struct { + name string + run func(t *testing.T) []any +}{ + {"spmd-shards", func(t *testing.T) []any { + // The distributed reduction answers the single-array + // reduction's exact bits at any world size: the digest pins + // the shards' answers beside the single-array ones they must + // equal, over a length that cuts the fold's partition into + // several blocks. + const gn = 200001 + vals := make([]float64, gn) + for i := range vals { + vals[i] = float64((i*6559)%2001-1000) / 7.0 + if i%401 == 3 { + vals[i] = math.NaN() + } + } + yVals := make([]float64, gn) + for i := range yVals { + yVals[i] = float64((i*7919)%901-450) / 5.0 + } + whole, err := FromFloats(vals, gn) + if err != nil { + t.Fatal(err) + } + wantSum := Sum(whole) + wantMax, err := Max(whole) + if err != nil { + t.Fatal(err) + } + out := []any{wantSum, wantMax} + for _, size := range []int{1, 3, 5} { + // Each rank writes its own slots; the digest takes them in + // rank order after the world has joined up again. + answers := make([]core.Scalar, size*5) + err := spmd.Launch(size, func(w *spmd.World) error { + span, err := spmd.Partition(gn, size, w.Rank()) + if err != nil { + return err + } + local, err := Slice(whole, 0, span.Lo, span.Hi) + if err != nil { + return err + } + gotSum, err := w.AllReduceShards(local, span, spmd.Sum) + if err != nil { + return err + } + gotMax, err := w.AllReduceShards(local, span, spmd.Max) + if err != nil { + return err + } + gotProd, err := w.AllReduceShards(local, span, spmd.Prod) + if err != nil { + return err + } + gotNorm, err := w.AllReduceNormShards(local, span, 2) + if err != nil { + return err + } + second, err := FromFloats(yVals, gn) + if err != nil { + return err + } + secondLocal, err := Slice(second, 0, span.Lo, span.Hi) + if err != nil { + return err + } + gotDot, err := w.AllReduceDotShards(local, secondLocal, span) + if err != nil { + return err + } + wantProd, err := Prod(whole, 0, false) + if err != nil { + return err + } + wantNorm, err := Norm(whole, 2, 0, false) + if err != nil { + return err + } + wantDot, err := Dot(whole, second) + if err != nil { + return err + } + if math.Float64bits(gotSum.Float()) != math.Float64bits(wantSum.Float()) || + math.Float64bits(gotMax.Float()) != math.Float64bits(wantMax.Float()) || + math.Float64bits(gotProd.Float()) != math.Float64bits(wantProd.FloatAt(0)) || + math.Float64bits(gotNorm.Float()) != math.Float64bits(wantNorm.FloatAt(0)) || + math.Float64bits(gotDot.Float()) != math.Float64bits(wantDot.Float()) { + t.Fatalf("size %d: the sharded answers moved against the single-array reduction", size) + } + answers[w.Rank()*5] = gotSum + answers[w.Rank()*5+1] = gotMax + answers[w.Rank()*5+2] = gotProd + answers[w.Rank()*5+3] = gotNorm + answers[w.Rank()*5+4] = gotDot + return nil + }) + if err != nil { + t.Fatal(err) + } + for _, s := range answers { + out = append(out, s) + } + out = append(out, fmt.Sprintf("size %d agreed", size)) + } + return out + }}, + {"core-elementwise", func(t *testing.T) []any { + a := randA(t, 1, 64) + b := randA(t, 2, 64) + sum, _ := Add(a, b) + prod, _ := Mul(a, b) + return []any{sum, prod, MulF(a, 1.5)} + }}, + {"core-matmul-einsum", func(t *testing.T) []any { + a := randA(t, 3, 64) + b := randA(t, 4, 64) + mm, err := MatMul2D(mustA(t, a.RawFloats(), 8, 8), mustA(t, b.RawFloats(), 8, 8)) + if err != nil { + t.Fatal(err) + } + e, err := Einsum("ij,jk->ik", mustA(t, a.RawFloats(), 8, 8), mustA(t, b.RawFloats(), 8, 8)) + if err != nil { + t.Fatal(err) + } + return []any{mm, e} + }}, + {"core-sort-argsort", func(t *testing.T) []any { + vals := randA(t, 5, 200).RawFloats() + vals[3] = math.NaN() + vals[17] = math.Copysign(0, -1) + vals[42] = math.Inf(1) + vals[43] = math.Inf(-1) + vals[7] = vals[8] // a tie + src := mustA(t, vals, len(vals)) + s, err := Sort(src) + if err != nil { + t.Fatal(err) + } + idx, err := ArgSort(src) + if err != nil { + t.Fatal(err) + } + return []any{s, idx} + }}, + {"core-reductions", func(t *testing.T) []any { + a := randA(t, 6, 1000) + b := randA(t, 7, 1000) + dot, err := Dot(a, b) + if err != nil { + t.Fatal(err) + } + mean, err := Mean(a) + if err != nil { + t.Fatal(err) + } + mn, err := Min(a) + if err != nil { + t.Fatal(err) + } + mx, err := Max(a) + if err != nil { + t.Fatal(err) + } + cs, err := CumSum(a, 0) + if err != nil { + t.Fatal(err) + } + return []any{Sum(a), mean, mn, mx, dot, cs} + }}, + {"core-shape", func(t *testing.T) []any { + a := randA(t, 8, 24) + m := mustA(t, a.RawFloats(), 4, 6) + tr, err := TransposeAxes(m, 1, 0) + if err != nil { + t.Fatal(err) + } + r, err := Reshape(m, 3, 8) + if err != nil { + t.Fatal(err) + } + sl, err := Slice(m, 0, 1, 3) + if err != nil { + t.Fatal(err) + } + return []any{tr, r, sl} + }}, + {"core-fft", func(t *testing.T) []any { + a := randA(t, 9, 256) + spec, err := FFT(a) + if err != nil { + t.Fatal(err) + } + back, err := IFFT(spec) + if err != nil { + t.Fatal(err) + } + return []any{spec, back} + }}, + {"core-quasirandom", func(t *testing.T) []any { + sob, err := SobolPoints(64, 4, 0) + if err != nil { + t.Fatal(err) + } + sob2, err := SobolPoints(32, 4, 1000) + if err != nil { + t.Fatal(err) + } + hal, err := HaltonPoints(64, 4, 0) + if err != nil { + t.Fatal(err) + } + return []any{sob, sob2, hal} + }}, + {"linalg-solve-inv-det", func(t *testing.T) []any { + a := randA(t, 10, 25) + // Diagonal dominance keeps the solve well conditioned. + m := mustA(t, a.RawFloats(), 5, 5) + md := m.RawFloats() + for i := range 5 { + md[i*5+i] += 10 + } + b := randA(t, 11, 5) + x, err := Solve(m, b) + if err != nil { + t.Fatal(err) + } + iv, err := Inv(m) + if err != nil { + t.Fatal(err) + } + det, err := Det(m) + if err != nil { + t.Fatal(err) + } + return []any{x, iv, det} + }}, + {"linalg-factorisations", func(t *testing.T) []any { + a := randA(t, 12, 36) + m := mustA(t, a.RawFloats(), 6, 6) + // Symmetrise, then shift to positive definite. + sym := make([]float64, 36) + mm := m.RawFloats() + for i := range 6 { + for j := range 6 { + sym[i*6+j] = (mm[i*6+j] + mm[j*6+i]) / 2 + } + sym[i*6+i] += 8 + } + pd := mustA(t, sym, 6, 6) + q, r, err := QR(pd) + if err != nil { + t.Fatal(err) + } + chol, err := Cholesky(pd) + if err != nil { + t.Fatal(err) + } + u, sigma, vt, err := SVD(pd) + if err != nil { + t.Fatal(err) + } + vals, vecs, err := Eigen(pd) + if err != nil { + t.Fatal(err) + } + return []any{q, r, chol, u, sigma, vt, vals, vecs} + }}, + {"linalg-sparse-complex", func(t *testing.T) []any { + const n = 32 + // Hermitian tridiagonal: real diagonal, imaginary couplings + // conjugated across the diagonal. + var hIdx []int64 + var hVal []complex128 + for i := range n { + hIdx = append(hIdx, int64(i), int64(i)) + hVal = append(hVal, 2+0i) + if i+1 < n { + hIdx = append(hIdx, int64(i), int64(i+1)) + hVal = append(hVal, 0.5i) + hIdx = append(hIdx, int64(i+1), int64(i)) + hVal = append(hVal, -0.5i) + } + } + hIndices, err := FromInts(hIdx, len(hVal), 2) + if err != nil { + t.Fatal(err) + } + hValues, err := FromComplexes(hVal, len(hVal)) + if err != nil { + t.Fatal(err) + } + h, err := NewSparseCOO(hIndices, hValues, []int{n, n}) + if err != nil { + t.Fatal(err) + } + ones := make([]complex128, n) + for i := range ones { + ones[i] = 1 + } + rhs, err := FromComplexes(ones, n) + if err != nil { + t.Fatal(err) + } + xCG, err := SpSolveComplexCG(h, rhs, 1e-12, 200) + if err != nil { + t.Fatal(err) + } + // Genuinely non-Hermitian: different off-diagonal values. + var gIdx []int64 + var gVal []complex128 + for i := range n { + gIdx = append(gIdx, int64(i), int64(i)) + gVal = append(gVal, 2+1i) + if i+1 < n { + gIdx = append(gIdx, int64(i), int64(i+1)) + gVal = append(gVal, 1) + gIdx = append(gIdx, int64(i+1), int64(i)) + gVal = append(gVal, 0.25i) + } + } + gIndices, err := FromInts(gIdx, len(gVal), 2) + if err != nil { + t.Fatal(err) + } + gValues, err := FromComplexes(gVal, len(gVal)) + if err != nil { + t.Fatal(err) + } + g, err := NewSparseCOO(gIndices, gValues, []int{n, n}) + if err != nil { + t.Fatal(err) + } + xBiCG, err := SpSolveComplexBiCGSTAB(g, rhs, 1e-12, 400) + if err != nil { + t.Fatal(err) + } + vals, vecs, err := SpEigenComplex(h, 4, NewGenerator(11)) + if err != nil { + t.Fatal(err) + } + return []any{xCG, xBiCG, vals, vecs} + }}, + {"linalg-sparse-cholesky", func(t *testing.T) []any { + // A shuffled 5-point Laplacian: the ordering does the work, + // the solve must answer A⁻¹·(A·v) = v for both orderings. + const n = 100 + g := NewGenerator(31) + perm := make([]int, n) + for i := range perm { + perm[i] = i + } + for i := n - 1; i > 0; i-- { + j := int(g.Next() % uint64(i+1)) + perm[i], perm[j] = perm[j], perm[i] + } + label := func(x, y, w int) int { return perm[y*w+x] } + var idx []int64 + var vals []float64 + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + const w, h = 10, 10 + for y := range h { + for x := range w { + add(label(x, y, w), label(x, y, w), 4) + if x+1 < w { + add(label(x, y, w), label(x+1, y, w), -1) + add(label(x+1, y, w), label(x, y, w), -1) + } + if y+1 < h { + add(label(x, y, w), label(x, y+1, w), -1) + add(label(x, y+1, w), label(x, y, w), -1) + } + } + } + indices, err := FromInts(idx, len(vals), 2) + if err != nil { + t.Fatal(err) + } + valArray, err := FloatsFromArray(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + coo, err := NewSparseCOO(indices, valArray, []int{n, n}) + if err != nil { + t.Fatal(err) + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatal(err) + } + v := make([]float64, n) + for i := range v { + v[i] = float64(i%13) - 6 + 0.5*float64(i%7) + } + xTrue := New(Float, n) + for i := range v { + xTrue.RawFloats()[i] = v[i] + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatal(err) + } + nat, err := NewSparseCholesky(coo, SparseOrderingNatural) + if err != nil { + t.Fatal(err) + } + rcm, err := NewSparseCholesky(coo, SparseOrderingReverseCuthillMcKee) + if err != nil { + t.Fatal(err) + } + xn, err := nat.Solve(b) + if err != nil { + t.Fatal(err) + } + xr, err := rcm.Solve(b) + if err != nil { + t.Fatal(err) + } + pn := nat.Permutation() + pr := rcm.Permutation() + pnf := make([]float64, len(pn)) + prf := make([]float64, len(pr)) + for i := range pn { + pnf[i] = float64(pn[i]) + prf[i] = float64(pr[i]) + } + return []any{pnf, prf, float64(nat.NNZ()), float64(rcm.NNZ()), xn, xr} + }}, + {"linalg-sparse-lu", func(t *testing.T) []any { + // A nonsymmetric banded system with real fill: the factor must + // invert the map the CSR MatVec builds, pivoting included. + const n = 40 + g := NewGenerator(41) + var idx []int64 + var vals []float64 + add := func(r, c int, v float64) { + idx = append(idx, int64(r), int64(c)) + vals = append(vals, v) + } + for i := range n { + add(i, i, 5+2*g.Unit()) + if i+1 < n { + add(i, i+1, -1-g.Unit()) + add(i+1, i, -1-g.Unit()) + } + if i+3 < n { + add(i, i+3, -0.6*g.Unit()) + } + if i+4 < n { + add(i+4, i, -0.4*g.Unit()) + } + } + indices, err := FromInts(idx, len(vals), 2) + if err != nil { + t.Fatal(err) + } + valArray, err := FloatsFromArray(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + coo, err := NewSparseCOO(indices, valArray, []int{n, n}) + if err != nil { + t.Fatal(err) + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatal(err) + } + v := make([]float64, n) + for i := range v { + v[i] = float64(i%9) - 4 + 0.25*float64(i%5) + } + xTrue := New(Float, n) + for i := range v { + xTrue.RawFloats()[i] = v[i] + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatal(err) + } + f, err := NewSparseLU(coo) + if err != nil { + t.Fatal(err) + } + x, err := f.Solve(b) + if err != nil { + t.Fatal(err) + } + p := f.Permutation() + pf := make([]float64, len(p)) + for i := range p { + pf[i] = float64(p[i]) + } + return []any{pf, float64(f.NNZ()), x} + }}, + {"integrate-fem-poisson", func(t *testing.T) []any { + // The manufactured solution u = sin(πx)·sin(πy) on a 10×10 + // structured mesh of the unit square: assemble, lift the + // boundary, solve through the sparse Cholesky. + const m = 10 + n := (m + 1) * (m + 1) + vertices := make([]float64, 2*n) + for j := range m + 1 { + for i := range m + 1 { + vertices[2*(j*(m+1)+i)] = float64(i) / float64(m) + vertices[2*(j*(m+1)+i)+1] = float64(j) / float64(m) + } + } + at := func(i, j int) int64 { return int64(j*(m+1) + i) } + idx := make([]int64, 0, 6*m*m) + for j := range m { + for i := range m { + idx = append(idx, at(i, j), at(i+1, j), at(i+1, j+1), at(i, j), at(i+1, j+1), at(i, j+1)) + } + } + vArr, err := FromFloats(vertices, n, 2) + if err != nil { + t.Fatal(err) + } + tArr, err := FromInts(idx, m*m*2, 3) + if err != nil { + t.Fatal(err) + } + mesh, err := NewTriangleMesh2D(vArr, tArr) + if err != nil { + t.Fatal(err) + } + sol := func(x, y float64) float64 { return math.Sin(math.Pi*x) * math.Sin(math.Pi*y) } + src := func(x, y float64) float64 { return 2 * math.Pi * math.Pi * sol(x, y) } + var bound []int + var vals []float64 + for j := range m + 1 { + for i := range m + 1 { + if i == 0 || i == m || j == 0 || j == m { + bound = append(bound, j*(m+1)+i) + vals = append(vals, sol(float64(i)/float64(m), float64(j)/float64(m))) + } + } + } + u, err := SolvePoissonFEM2D(mesh, src, FEMPoissonOptions{Kappa: 1, DirichletNodes: bound, DirichletValues: vals, Ordering: SparseOrderingReverseCuthillMcKee}) + if err != nil { + t.Fatal(err) + } + first := make([]float64, 0, 12) + for i := range 12 { + first = append(first, u.FloatAt(i)) + } + return []any{first, u.FloatAt(n / 2), u.FloatAt(n - 1)} + }}, + {"signal-welch-savgol", func(t *testing.T) []any { + x := randA(t, 13, 512) + freqs, psd, err := WelchPSD(x, 1000, 128, 64, "hann") + if err != nil { + t.Fatal(err) + } + sm, err := SavitzkyGolay(x, 11, 3) + if err != nil { + t.Fatal(err) + } + return []any{freqs, psd, sm} + }}, + {"signal-wavelets", func(t *testing.T) []any { + x := randA(t, 14, 128) + coef, err := DWT(x, 3) + if err != nil { + t.Fatal(err) + } + back, err := IDWT(coef, 3) + if err != nil { + t.Fatal(err) + } + cw, err := CWT(x, Morlet, []float64{1, 2, 4, 8}, 1) + if err != nil { + t.Fatal(err) + } + return []any{coef, back, cw} + }}, + {"signal-lombscargle-conv", func(t *testing.T) []any { + g := NewGenerator(15) + times, err := Floats(g, 100) + if err != nil { + t.Fatal(err) + } + acc := 0.0 + for i, v := range times.RawFloats() { + acc += v + 0.1 + times.RawFloats()[i] = acc // irregular but increasing + } + vals := randA(t, 16, 100) + freqs, power, err := LombScargle(times, vals, 0.1, 5, 32) + if err != nil { + t.Fatal(err) + } + input := mustA(t, randA(t, 17, 64).RawFloats(), 1, 1, 64) + kernel := mustA(t, randA(t, 18, 5).RawFloats(), 1, 1, 5) + conv, err := Conv1D(input, kernel, nil, 1, 0, 1) + if err != nil { + t.Fatal(err) + } + return []any{freqs, power, conv} + }}, + {"signal-filters-stencils", func(t *testing.T) []any { + x := randA(t, 19, 256) + b, a, err := ButterworthLowPass(4, 1000, 100) + if err != nil { + t.Fatal(err) + } + filt, err := FilterApply(b, a, x) + if err != nil { + t.Fatal(err) + } + grad, err := Gradient1D(x, 0.5) + if err != nil { + t.Fatal(err) + } + m := mustA(t, randA(t, 20, 64).RawFloats(), 8, 8) + lap, err := Laplacian(m, 1.0, 1.0) + if err != nil { + t.Fatal(err) + } + return []any{filt, grad, lap} + }}, + {"stats-moments-quantiles", func(t *testing.T) []any { + a := randA(t, 21, 1000) + v, err := Var(a) + if err != nil { + t.Fatal(err) + } + s, err := Std(a) + if err != nil { + t.Fatal(err) + } + med, err := Median(a) + if err != nil { + t.Fatal(err) + } + q, err := Quantile(a, []float64{0.05, 0.25, 0.5, 0.75, 0.95}) + if err != nil { + t.Fatal(err) + } + mean, err := Mean(a) + if err != nil { + t.Fatal(err) + } + return []any{mean, v, s, med, q} + }}, + {"stats-correlation-regression", func(t *testing.T) []any { + x := randA(t, 22, 200) + ac, err := Autocorrelate(x, 16) + if err != nil { + t.Fatal(err) + } + pac, err := PartialAutocorrelate(x, 8) + if err != nil { + t.Fatal(err) + } + // The design carries its own intercept column. + xd := x.RawFloats() + design := make([]float64, 400) + for i := range 200 { + design[i*2] = 1 + design[i*2+1] = xd[i] + } + lr, err := LinearRegression(mustA(t, design, 200, 2), randA(t, 23, 200)) + if err != nil { + t.Fatal(err) + } + return []any{ac, pac, lr.Coefficients, lr.StandardErrors, lr.PValues, + lr.ResidualVariance, lr.RSquared} + }}, + {"stats-poisson-regression", func(t *testing.T) []any { + // Counts from the seeded generator: one draw per row at the + // mean the true coefficients produce. + g := NewGenerator(24) + xv := randA(t, 23, 120) + design := make([]float64, 240) + y := make([]float64, 120) + for i := range 120 { + design[i*2] = 1 + design[i*2+1] = xv.FloatAt(i) + counts, err := PoissonDraws(g, 1, math.Exp(0.3+0.7*xv.FloatAt(i))) + if err != nil { + t.Fatal(err) + } + y[i] = counts.FloatAt(0) + } + res, err := PoissonRegression(mustA(t, design, 120, 2), mustA(t, y, 120)) + if err != nil { + t.Fatal(err) + } + return []any{res.Coefficients, res.StandardErrors, res.PValues, + res.Fitted, res.LogLikelihood, res.Iterations} + }}, + {"stats-distributions", func(t *testing.T) []any { + tt, df, p, err := WelchTTest(randA(t, 24, 50), randA(t, 25, 60)) + if err != nil { + t.Fatal(err) + } + gc, err := GammaCDF(2.5, 3, 1.5) + if err != nil { + t.Fatal(err) + } + ec, err := ExponentialCDF(1.25, 2) + if err != nil { + t.Fatal(err) + } + sc, err := StudentTCDF(1.5, 7) + if err != nil { + t.Fatal(err) + } + pc, err := PoissonCDF(3, 2.5) + if err != nil { + t.Fatal(err) + } + bc, err := BinomialCDF(4, 10, 0.35) + if err != nil { + t.Fatal(err) + } + nq, err := NormalQuantile(0.975) + if err != nil { + t.Fatal(err) + } + return []any{NormalCDF(1.96), tt, df, p, gc, ec, sc, pc, bc, nq} + }}, + {"integrate-quadrature-ode", func(t *testing.T) []any { + q, qerr, err := IntegrateFunction(func(x float64) (float64, error) { + return math.Exp(-x * x), nil + }, 0, 1, QuadratureOptions{}) + if err != nil { + t.Fatal(err) + } + decay := func(t float64, y *Array) (*Array, error) { return MulF(y, -1), nil } + y0 := mustA(t, []float64{1}, 1) + end, err := IntegrateODE(decay, 0, 2, y0, ODEOptions{MaxSteps: 10000}) + if err != nil { + t.Fatal(err) + } + nd, err := IntegrateND(func(x []float64) float64 { + return math.Exp(-x[0] - x[1]*x[1]) + }, []float64{0, 0}, []float64{1, 1}, CubatureOptions{}) + if err != nil { + t.Fatal(err) + } + heat, err := IntegrateHeat1D(mustA(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 8), + 0.1, 0.1, 0.05, 0.005, 3, 0, 0) + if err != nil { + t.Fatal(err) + } + return []any{q, qerr, end, nd, heat} + }}, + {"optim-roots-minima", func(t *testing.T) []any { + root, err := FindRoot(func(x float64) float64 { return math.Cos(x) - x }, 0, 1, 1e-12) + if err != nil { + t.Fatal(err) + } + rosen := func(p *Array) (float64, error) { + x := p.FloatAt(0) + y := p.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + } + pt, val, err := Minimise(rosen, mustA(t, []float64{-1.2, 1}, 2), + MinimiseOptions{MaxIterations: 500, Tolerance: 1e-12}) + if err != nil { + t.Fatal(err) + } + return []any{root, pt, val} + }}, + {"optim-bounded-minima", func(t *testing.T) []any { + // The bounded solver on the wall case: with x forced past 1.5 + // the Rosenbrock neck at (1, 1) is infeasible and the minimum + // sits at (1.5, 2.25) with value 0.25, the first coordinate + // pinned by the projection and released never. + rosen := func(p *Array) (float64, error) { + x := p.FloatAt(0) + y := p.FloatAt(1) + return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil + } + pt, val, err := MinimiseLBFGS(rosen, nil, mustA(t, []float64{-1.2, 1}, 2), + LBFGSOptions{Tolerance: 1e-10, Lower: []float64{1.5, math.Inf(-1)}}) + if err != nil { + t.Fatal(err) + } + refused := false + _, _, err = MinimiseLBFGS(rosen, nil, mustA(t, []float64{0, 0}, 2), + LBFGSOptions{Lower: []float64{2}, Upper: []float64{1}}) + if err != nil { + refused = true + } + return []any{pt, val, refused} + }}, + {"optim-constrained-minima", func(t *testing.T) []any { + // The augmented Lagrangian on an active inequality: the bowl + // around (2, -1) cut by x + y ≤ 0 bottoms out on the wall at + // (1.5, -1.5) with value 0.5. + bowl := func(p *Array) (float64, error) { + dx := p.FloatAt(0) - 2 + dy := p.FloatAt(1) + 1 + return dx*dx + dy*dy, nil + } + A := mustA(t, []float64{1, 1}, 1, 2) + pt, val, err := MinimiseConstrained(bowl, nil, mustA(t, []float64{4, 4}, 2), + LinearConstraints{A: A, Lower: []float64{math.Inf(-1)}, Upper: []float64{0}}, + LBFGSOptions{Tolerance: 1e-10}) + if err != nil { + t.Fatal(err) + } + refused := false + _, _, err = MinimiseConstrained(bowl, nil, mustA(t, []float64{0, 0}, 2), + LinearConstraints{Lower: []float64{0}, Upper: []float64{1}}, LBFGSOptions{}) + if err != nil { + refused = true + } + return []any{pt, val, refused} + }}, + {"grad-backward", func(t *testing.T) []any { + x := FromArray(mustA(t, []float64{1.5, -2, 3}, 3), true) + loss, err := x.Mul(x) + if err != nil { + t.Fatal(err) + } + l, err := loss.Sum() + if err != nil { + t.Fatal(err) + } + if err := l.Backward(); err != nil { + t.Fatal(err) + } + return []any{x.Grad()} + }}, + {"io-roundtrips", func(t *testing.T) []any { + dir := t.TempDir() + a := randA(t, 26, 24) + m := mustA(t, a.RawFloats(), 4, 6) + fitsPath := filepath.Join(dir, "o.fits") + if err := SaveFITS(fitsPath, m, map[string]string{"OBJECT": "oracle"}); err != nil { + t.Fatal(err) + } + backFits, hdr, err := LoadFITS(fitsPath) + if err != nil { + t.Fatal(err) + } + ncPath := filepath.Join(dir, "o.nc") + dims := []NetCDFDim{{Name: "row", Length: 4}, {Name: "col", Length: 6}} + vars := []NetCDFVar{{Name: "field", Dims: []string{"row", "col"}, Values: m}} + if err := SaveNetCDF(ncPath, dims, vars, map[string]string{"title": "oracle"}); err != nil { + t.Fatal(err) + } + _, backNc, ncAttrs, err := LoadNetCDF(ncPath) + if err != nil { + t.Fatal(err) + } + return []any{backFits, hdr["OBJECT"], backNc[0].Values, ncAttrs["title"]} + }}, + {"core-views-and-int-precision", func(t *testing.T) []any { + // Views, the integer range above 2^53 and a widened query + // array: the 2026-09 review found defects in all three. + x := randA(t, 41, 12) + view, err := Slice(x, 0, 3, 9) + if err != nil { + t.Fatal(err) + } + pow, err := PowI(view, 3) + if err != nil { + t.Fatal(err) + } + big, err := FromInts([]int64{1 << 53, 1<<53 + 1, -(1 << 53)}, 3) + if err != nil { + t.Fatal(err) + } + hi, err := ArgMax(big) + if err != nil { + t.Fatal(err) + } + lo, err := ArgMin(big) + if err != nil { + t.Fatal(err) + } + // A float32 view whose payload is longer than its extent: the + // widening inside Interpolate2D used to walk the payload. + xs, err := FromFloat32s(make([]float32, x.Len()), x.Len()) + if err != nil { + t.Fatal(err) + } + for i := range xs.Len() { + // Stay inside the (3, 4) grid's closed domain. + xs.SetFloatAt(i, float64(i%4)) + } + viewXs, err := Slice(xs, 0, 0, view.Len()) + if err != nil { + t.Fatal(err) + } + grid, err := FromFloats(make([]float64, 12), 3, 4) + if err != nil { + t.Fatal(err) + } + // The y queries need their own range: the grid's y domain is + // [0, rows−1], narrower than x's. + ys, err := FromFloat32s(make([]float32, viewXs.Len()), viewXs.Len()) + if err != nil { + t.Fatal(err) + } + for i := range ys.Len() { + ys.SetFloatAt(i, float64(i%3)) + } + interp, err := Interpolate2D(grid, viewXs, ys, 0, 0, 1, 1) + if err != nil { + t.Fatal(err) + } + return []any{pow, hi, lo, interp} + }}, + {"io-hdf5-fixture", func(t *testing.T) []any { + // The HDF5 reference-library fixture: contiguous ints, a chunked + // deflate+shuffle float vector, a float32 matrix in a group, + // and attributes inherited from the groups. + sets, err := LoadHDF5(filepath.Join("io", "testdata", "h5", "fixture.h5")) + if err != nil { + t.Fatal(err) + } + out := []any{} + for _, d := range sets { + out = append(out, d.Path, d.Values) + for _, k := range slices.Sorted(maps.Keys(d.Attrs)) { + out = append(out, k, d.Attrs[k]) + } + } + return out + }}, + {"io-hostile-inputs", func(t *testing.T) []any { + // A header that lies about its sizes must be refused, never + // allocated and never panicked over: both files here are a few + // dozen bytes and claim data no file could hold. + dir := t.TempDir() + ncPath := filepath.Join(dir, "hostile.nc") + if err := os.WriteFile(ncPath, hostileOracleHeader(), 0o644); err != nil { + t.Fatal(err) + } + _, _, _, ncErr := LoadNetCDF(ncPath) + fitsPath := filepath.Join(dir, "hostile.fits") + if err := os.WriteFile(fitsPath, hostileOracleTable(), 0o644); err != nil { + t.Fatal(err) + } + _, fitsErr := LoadFITSTable(fitsPath) + return []any{ncErr != nil, fitsErr != nil} + }}, + {"core-float16", func(t *testing.T) []any { + // The half dtype end to end: narrowing of values that straddle + // the format's corners, arithmetic through the promotion + // ladder, ordering that must widen before it compares. + vals := []float64{1, -2.5, 0.25, 65504, -65504, 1.0 / 16384, 1.0 / 16777216, 0.1, 2048.5} + h, err := FromFloat16s(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + s := Sum(h) + sorted, err := Sort(h) + if err != nil { + t.Fatal(err) + } + prod := MulF(h, 2) + widened, err := Add(h, mustA(t, make([]float64, len(vals)), len(vals))) + if err != nil { + t.Fatal(err) + } + ints, err := Astype(h, Int) + if err != nil { + t.Fatal(err) + } + return []any{h, s, sorted, prod, widened, ints, + HalfToFloat64(HalfFromFloat64(0.1)), HalfToFloat64(0x7BFF), + HalfToFloat64(0x0001), HalfFromFloat64(65520) == 0x7C00, + HalfFromFloat64(-1.0/33554432.0) == 0x8000} + }}, + {"signal-kalman-arma", func(t *testing.T) []any { + // The linear filter on a scalar random walk with noisy + // measurements from the seeded generator, the unscented filter + // over the same exactly linear model, and the AR fit with its + // theoretical spectrum. + g := NewGenerator(50) + z, err := Floats(g, 40) + if err != nil { + t.Fatal(err) + } + one := mustA(t, []float64{1}, 1, 1) + qv := mustA(t, []float64{0.01}, 1, 1) + rv := mustA(t, []float64{1}, 1, 1) + zero := mustA(t, []float64{0}, 1) + pv := mustA(t, []float64{1}, 1, 1) + res, err := KalmanFilter(z, one, one, KalmanOptions{ + InitialState: zero, InitialCovariance: pv, + ProcessNoise: qv, MeasurementNoise: rv}) + if err != nil { + t.Fatal(err) + } + uk, err := UnscentedKalmanFilter(z, + func(x *Array) (*Array, error) { return MulF(x, 1), nil }, + func(x *Array) (*Array, error) { return MulF(x, 1), nil }, + KalmanOptions{ + InitialState: zero, InitialCovariance: pv, + ProcessNoise: qv, MeasurementNoise: rv}) + if err != nil { + t.Fatal(err) + } + x, err := Floats(NewGenerator(51), 300) + if err != nil { + t.Fatal(err) + } + ar, err := EstimateAR(x, 2) + if err != nil { + t.Fatal(err) + } + freqs, psd, err := ARMASpectrum(ar, 32) + if err != nil { + t.Fatal(err) + } + return []any{res.States, res.Covariances, res.Innovations, res.LogLikelihood, + uk.LogLikelihood, uk.States, ar.AR, ar.InnovationVariance, ar.AIC, freqs, psd} + }}, + {"signal-windows-filtfilt-dwt", func(t *testing.T) []any { + // The window catalogue, the zero-phase filter, the 2-D median + // and the Daubechies transform over seeded inputs. + hann, err := WindowHann(16, false) + if err != nil { + t.Fatal(err) + } + kaiser, err := WindowKaiser(16, 4.5, true) + if err != nil { + t.Fatal(err) + } + x := randA(t, 52, 128) + b, a, err := ButterworthLowPass(3, 100, 20) + if err != nil { + t.Fatal(err) + } + ff, err := Filtfilt(b, a, x) + if err != nil { + t.Fatal(err) + } + img := mustA(t, randA(t, 53, 64).RawFloats(), 8, 8) + med, err := MedianFilter2D(img, 3) + if err != nil { + t.Fatal(err) + } + coef, err := DaubechiesDWT(x, DB4, 2, DWTPeriodic) + if err != nil { + t.Fatal(err) + } + back, err := DaubechiesIDWT(coef, DB4, 2, DWTPeriodic) + if err != nil { + t.Fatal(err) + } + return []any{hann, kaiser, ff, med, coef, back} + }}, + {"stats-distributions2-multipletest", func(t *testing.T) []any { + // The second-league distributions and the corrections, on + // fixed parameters and the classic Benjamini-Hochberg vector. + wc, err := WeibullCDF(1.5, 2, 3) + if err != nil { + t.Fatal(err) + } + wq, err := WeibullQuantile(0.4, 2, 3) + if err != nil { + t.Fatal(err) + } + lc, err := LognormalCDF(1, 0, 0.5) + if err != nil { + t.Fatal(err) + } + pq, err := ParetoQuantile(0.5, 1, 3) + if err != nil { + t.Fatal(err) + } + nb, err := NegativeBinomialCDF(5, 3, 0.4) + if err != nil { + t.Fatal(err) + } + ncx, err := NoncentralChiSquareCDF(9, 4, 2.5) + if err != nil { + t.Fatal(err) + } + nct, err := NoncentralTCDF(2, 8, 3) + if err != nil { + t.Fatal(err) + } + ncf, err := NoncentralFQuantile(0.5, 4, 10, 2) + if err != nil { + t.Fatal(err) + } + draws, err := DirichletDraws(NewGenerator(54), 12, []float64{2, 3, 4}) + if err != nil { + t.Fatal(err) + } + px := randA(t, 55, 100) + py := randA(t, 56, 100) + rho, err := SpearmanRho(px, py) + if err != nil { + t.Fatal(err) + } + tau, err := KendallTau(px, py) + if err != nil { + t.Fatal(err) + } + p := []float64{0.001, 0.008, 0.039, 0.041, 0.042, 0.06, 0.074, 0.205, + 0.212, 0.216, 0.222, 0.251, 0.269, 0.275, 0.34, 0.341, 0.384, + 0.456, 0.657, 0.876} + bonf, err := Bonferroni(p) + if err != nil { + t.Fatal(err) + } + holm, err := Holm(p) + if err != nil { + t.Fatal(err) + } + bh, err := BenjaminiHochberg(p) + if err != nil { + t.Fatal(err) + } + return []any{wc, wq, lc, pq, nb, ncx, nct, ncf, draws, rho, tau, bonf, holm, bh} + }}, + {"stats-regression2", func(t *testing.T) []any { + // The regularised, robust and quantile fits on one seeded + // design, with one gross outlier the robust fit must survive. + g := NewGenerator(57) + xv, err := Floats(g, 60) + if err != nil { + t.Fatal(err) + } + n := xv.Len() + design := make([]float64, 0, 2*n) + y := make([]float64, n) + for i := range n { + design = append(design, 1, xv.FloatAt(i)) + y[i] = 1 + 2*xv.FloatAt(i) + 0.3*(float64(i%7)-3) + } + y[7] += 100 + dArr := mustA(t, design, n, 2) + yArr := mustA(t, y, n) + xOnly := make([]float64, n) + for i := range n { + xOnly[i] = xv.FloatAt(i) + } + xArr := mustA(t, xOnly, n, 1) + lasso, err := Lasso(xArr, yArr, 0.05) + if err != nil { + t.Fatal(err) + } + en, err := ElasticNet(xArr, yArr, 0.05, 0.5) + if err != nil { + t.Fatal(err) + } + huber, err := HuberRegression(dArr, yArr) + if err != nil { + t.Fatal(err) + } + si, sl, err := TheilSenRegression(xv, yArr) + if err != nil { + t.Fatal(err) + } + qr, err := QuantileRegression(dArr, yArr, 0.75) + if err != nil { + t.Fatal(err) + } + return []any{lasso.Intercept, lasso.Coefficients, lasso.Iterations, + en.Coefficients, huber.Coefficients, huber.Scale, huber.Iterations, + si, sl, qr.Coefficients, qr.Objective} + }}, + {"stats-unsupervised", func(t *testing.T) []any { + // PCA over a seeded cloud, k-means over two blobs, the mixture + // over a univariate draw and the Gaussian-process posterior. + g := NewGenerator(58) + a1, err := Normal(g, 48, 0, 1) + if err != nil { + t.Fatal(err) + } + a2, err := Normal(g, 48, 0, 1) + if err != nil { + t.Fatal(err) + } + a3, err := Normal(g, 48, 0, 1) + if err != nil { + t.Fatal(err) + } + cloud := make([]float64, 144) + for i := range 48 { + cloud[3*i] = 3 * a1.FloatAt(i) + cloud[3*i+1] = a2.FloatAt(i) + cloud[3*i+2] = 0.1 * a3.FloatAt(i) + } + pca, err := PCA(mustA(t, cloud, 48, 3)) + if err != nil { + t.Fatal(err) + } + bg := NewGenerator(59) + b1, err := Normal(bg, 40, 0, 1) + if err != nil { + t.Fatal(err) + } + b2, err := Normal(bg, 40, 0, 1) + if err != nil { + t.Fatal(err) + } + b3, err := Normal(bg, 40, 8, 1) + if err != nil { + t.Fatal(err) + } + b4, err := Normal(bg, 40, 8, 1) + if err != nil { + t.Fatal(err) + } + blobs := make([]float64, 160) + for i := range 40 { + blobs[2*i] = b1.FloatAt(i) + blobs[2*i+1] = b2.FloatAt(i) + blobs[40+2*i] = b3.FloatAt(i) + blobs[40+2*i+1] = b4.FloatAt(i) + } + km, err := KMeans(NewGenerator(60), mustA(t, blobs, 80, 2), 2) + if err != nil { + t.Fatal(err) + } + centres := make([]float64, 0, 4) + for _, c := range km.Centres { + centres = append(centres, c...) + } + mg := NewGenerator(61) + hi, err := Normal(mg, 40, 6, 1) + if err != nil { + t.Fatal(err) + } + lo, err := Normal(mg, 40, 0, 1) + if err != nil { + t.Fatal(err) + } + mix := make([]float64, 80) + for i := range 40 { + mix[2*i] = hi.FloatAt(i) + mix[2*i+1] = lo.FloatAt(i) + } + gmm, err := GaussianMixture(NewGenerator(62), mustA(t, mix, 80, 1), 2) + if err != nil { + t.Fatal(err) + } + means := make([]float64, 0, len(gmm.Means)) + for _, m := range gmm.Means { + means = append(means, m...) + } + trainX := mustA(t, randA(t, 63, 16).RawFloats(), 16, 1) + sinVals, err := Sin(trainX) + if err != nil { + t.Fatal(err) + } + trainY, err := Reshape(sinVals, 16) + if err != nil { + t.Fatal(err) + } + testX := mustA(t, randA(t, 64, 8).RawFloats(), 8, 1) + kern, err := SquaredExponentialKernel(1) + if err != nil { + t.Fatal(err) + } + gp, err := GaussianProcessRegression(kern, trainX, trainY, 0.05, testX) + if err != nil { + t.Fatal(err) + } + mll, err := MarginalLogLikelihood(kern, trainX, trainY, 0.05) + if err != nil { + t.Fatal(err) + } + return []any{pca.Loadings, pca.ExplainedVarianceRatio, float64(km.Iterations), + centres, km.Inertia, gmm.Weights, means, gmm.BIC, gp.Mean, gp.Variance, mll} + }}, + {"stats-contingency", func(t *testing.T) []any { + // The exact and the asymptotic table tests over fixed tables, + // every alternative, and the effect sizes beside them. + out := []any{} + tables := [][]float64{ + {12, 5, 7, 10}, + {8, 2, 1, 5}, + {17, 0, 0, 23}, + {21, 8, 3, 9, 15, 11, 4, 12, 30}, + } + for ti, tab := range tables[:3] { + a, _, err := FisherExactTest(mustA(t, tab, 2, 2), TwoSided) + if err != nil { + t.Fatal(err) + } + less, _, err := FisherExactTest(mustA(t, tab, 2, 2), Less) + if err != nil { + t.Fatal(err) + } + greater, or, err := FisherExactTest(mustA(t, tab, 2, 2), Greater) + if err != nil { + t.Fatal(err) + } + out = append(out, a, less, greater, or) + p, err := McNemarTest(mustA(t, tab, 2, 2)) + if err != nil { + t.Fatal(err) + } + out = append(out, p) + if ti < 2 { + v, err := CramersV(mustA(t, tab, 2, 2)) + if err != nil { + t.Fatal(err) + } + out = append(out, v) + } + } + chi2, df, p, err := ChiSquareIndependence(mustA(t, tables[3], 3, 3)) + if err != nil { + t.Fatal(err) + } + v, err := CramersV(mustA(t, tables[3], 3, 3)) + if err != nil { + t.Fatal(err) + } + out = append(out, chi2, df, p, v) + return out + }}, + {"stats-mixedmodel", func(t *testing.T) []any { + // A balanced random-intercept fit and a random-slope fit, over + // formula-driven data no generator touches. + const groups8 = 8 + const per = 6 + y := make([]float64, 0, groups8*per) + labels := make([]int, 0, groups8*per) + for g := range groups8 { + effect := 2 * math.Sin(1.7*float64(g)+0.4) + for i := range per { + y = append(y, 5+effect+wobble(g*per+i)) + labels = append(labels, g) + } + } + ones := make([]float64, len(y)) + for i := range ones { + ones[i] = 1 + } + res, err := LinearMixedModel(mustA(t, y, len(y)), mustA(t, ones, len(y), 1), + mustA(t, ones, len(y), 1), labels) + if err != nil { + t.Fatal(err) + } + out := []any{res.Coefficients, res.StandardErrors, + res.RandomEffects[0], res.RandomEffects[7], res.RandomCovariance, + res.ResidualVariance, res.LogLikelihood, res.Converged, + fmt.Sprintf("%v", res.GroupLabels)} + // The slope fit: the random design carries the covariate + // alone, the identified configuration. + slopes := []float64{0.9, -1.1, 1.9, -0.3, -1.8, 0.4} + sy := make([]float64, 0, len(slopes)*per) + sx := make([]float64, 0, 2*len(slopes)*per) + sz := make([]float64, 0, len(slopes)*per) + slabels := make([]int, 0, len(slopes)*per) + jitter := make([]float64, 0, len(slopes)*per) + for g := range slopes { + for _, x := range []float64{-1, -0.6, -0.2, 0.2, 0.6, 1} { + jitter = append(jitter, wobble(len(jitter)+13)*0.1) + slabels = append(slabels, g) + sz = append(sz, x) + } + } + mean := 0.0 + for _, j := range jitter { + mean += j + } + mean /= float64(len(jitter)) + for g, s := range slopes { + for k := range per { + i := g*per + k + x := sz[i] + sy = append(sy, 1+2*x+0.5*s*x+0.1*(jitter[i]-mean)) + sx = append(sx, 1, x) + } + } + slope, err := LinearMixedModel(mustA(t, sy, len(sy)), mustA(t, sx, len(sy), 2), + mustA(t, sz, len(sz), 1), slabels) + if err != nil { + t.Fatal(err) + } + return append(out, slope.Coefficients, slope.RandomCovariance, + slope.ResidualVariance, slope.Converged) + }}, + {"stats-hmm", func(t *testing.T) []any { + // The scaled recursions, the decode and the Baum-Welch fit on + // a fixed sequence under a fixed model. + model, err := NewHiddenMarkovModel([]float64{0.6, 0.4}, + []float64{0.7, 0.3, 0.2, 0.8}, []float64{0.9, 0.1, 0.25, 0.75}) + if err != nil { + t.Fatal(err) + } + observations := []int{0, 0, 1, 0, 1, 1, 0, 1, 1, 1, 0, 0, 1, 0} + filtered, ll, err := model.Forward(observations) + if err != nil { + t.Fatal(err) + } + smoothed, sll, err := model.Smooth(observations) + if err != nil { + t.Fatal(err) + } + path, pathProb, err := model.Viterbi(observations) + if err != nil { + t.Fatal(err) + } + out := []any{filtered[len(observations)-1], smoothed[0], ll, sll, + fmt.Sprintf("%v", path), pathProb} + seq := make([]int, 240) + for i := range seq { + // The wiggle stays inside ±0.4, so the shifted scale is + // positive and the modulo lands inside the symbol set. + seq[i] = int(wobble(i)*100+200) % 3 + } + fit, err := FitHiddenMarkovModel(NewGenerator(77), seq, 2, 3) + if err != nil { + t.Fatal(err) + } + return append(out, fit.Model.Initial, fit.Model.Transition, fit.Model.Emission, + fit.LogLikelihood, fit.Converged) + }}, + {"stats-hierarchy", func(t *testing.T) []any { + // The five linkages over one fixed two-column sample, with the + // cuts the dendrogram answers. + sample := make([]float64, 24) + for i := range 12 { + sample[2*i] = wobble(2*i+5)*3 + float64(i%4) + sample[2*i+1] = wobble(2*i+31) * 2 + } + out := []any{} + for _, method := range []Linkage{SingleLinkage, CompleteLinkage, AverageLinkage, CentroidLinkage, WardLinkage} { + d, err := HierarchicalClustering(mustA(t, sample, 12, 2), method) + if err != nil { + t.Fatal(err) + } + labels, err := d.Cut(3) + if err != nil { + t.Fatal(err) + } + out = append(out, d.Heights, fmt.Sprintf("%v", labels), d.Sizes[10]) + } + first, err := HierarchicalClustering(mustA(t, sample, 12, 2), WardLinkage) + if err != nil { + t.Fatal(err) + } + byHeight, err := first.CutHeight(first.Heights[8]) + if err != nil { + t.Fatal(err) + } + return append(out, fmt.Sprintf("%v", byHeight)) + }}, + {"integrate-stiff-solvers", func(t *testing.T) []any { + // The stiff batch on one decay problem, the index-1 DAE on its + // circuit, the collocation on a linear two-point problem, the + // symplectic pair and the advection pair. + decay := func(t float64, y *Array) (*Array, error) { return MulF(y, -1), nil } + y0 := mustA(t, []float64{1}, 1) + stats := &BDFVarStats{} + bdf, err := IntegrateBDFVar(decay, 0, 2, y0, BDFVarOptions{MaxSteps: 10000, Stats: stats}) + if err != nil { + t.Fatal(err) + } + ros, err := IntegrateROS4(decay, 0, 2, y0, ODEOptions{MaxSteps: 10000}) + if err != nil { + t.Fatal(err) + } + daeRHS := func(t float64, y *Array) (*Array, error) { + out, err := Zeros(Float, 2) + if err != nil { + return nil, err + } + out.SetFloatAt(0, 1-y.FloatAt(0)) + out.SetFloatAt(1, y.FloatAt(1)-y.FloatAt(0)) + return out, nil + } + mass, err := FromFloats([]float64{1, 0, 0, 0}, 2, 2) + if err != nil { + t.Fatal(err) + } + dae, err := IntegrateDAE(daeRHS, mass, 0, 1, mustA(t, []float64{0, 0}, 2), 50, DAEOptions{}) + if err != nil { + t.Fatal(err) + } + colRHS := func(t float64, y *Array) (*Array, error) { + out, err := Zeros(Float, 2) + if err != nil { + return nil, err + } + out.SetFloatAt(0, y.FloatAt(1)) + return out, nil + } + col, err := SolveBoundaryCollocation(colRHS, 0, 1, mustA(t, []float64{0, 0}, 2), + BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}}, + CollocationOptions{}) + if err != nil { + t.Fatal(err) + } + accel := func(q *Array) (*Array, error) { return MulF(q, -1), nil } + positions, momenta, err := IntegrateYoshida4(accel, 0, 6.283185307179586, + mustA(t, []float64{1}, 1), mustA(t, []float64{0}, 1), 200) + if err != nil { + t.Fatal(err) + } + gradH := func(z *Array) (*Array, error) { + out, err := Zeros(Float, 2) + if err != nil { + return nil, err + } + out.SetFloatAt(0, z.FloatAt(1)) + out.SetFloatAt(1, z.FloatAt(0)) + return out, nil + } + mids, _, err := IntegrateMidpoint(gradH, 0, 1, mustA(t, []float64{1}, 1), + mustA(t, []float64{1}, 1), 32, MidpointOptions{}) + if err != nil { + t.Fatal(err) + } + n := 64 + dx := 1.0 / 65 + pulse := make([]float64, n) + for i := range n { + c := (float64(i) + 0.5) * dx + if c >= 0.3 && c <= 0.6 { + pulse[i] = 1 + } + } + upwind, err := IntegrateUpwindAdvection1D(mustA(t, pulse, n), 1, dx, 0.2, 0.9*dx, 2, 0, 0) + if err != nil { + t.Fatal(err) + } + koren, err := IntegrateAdvection1D(mustA(t, pulse, n), 1, dx, 0.2, 0.9*dx, 2, 0, 0) + if err != nil { + t.Fatal(err) + } + return []any{bdf, float64(stats.MaxOrder), ros, dae, + col.Values[len(col.Values)-1], float64(len(col.Mesh)), + positions[len(positions)-1], momenta[len(momenta)-1], mids[len(mids)-1], + upwind, koren} + }}, + {"integrate-fem3d", func(t *testing.T) []any { + // The tetrahedral Poisson assembly on the unit box with the + // manufactured sin(πx)·sin(πy)·sin(πz) solution. + mesh, err := BoxTetraMesh3D(0, 0, 0, 1, 1, 1, 4, 4, 4) + if err != nil { + t.Fatal(err) + } + 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) + } + var bound []int + var vals []float64 + 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 { + bound = append(bound, i) + vals = append(vals, 0) + } + } + u, err := SolvePoissonFEM3D(mesh, src, FEMPoisson3DOptions{ + Kappa: 1, DirichletNodes: bound, DirichletValues: vals, + Ordering: SparseOrderingReverseCuthillMcKee}) + if err != nil { + t.Fatal(err) + } + return []any{u.FloatAt(u.Len() / 2), u.FloatAt(u.Len() - 1), u.FloatAt(7)} + }}, + {"linalg-sparse-lsqr-rrqr-update", func(t *testing.T) []any { + // Sparse least squares against a consistent right-hand side, + // the pivoted QR on a seeded matrix and the sparse rank-one + // round trip on a banded factor. + g := NewGenerator(65) + const m, n = 12, 8 + var idx []int64 + var vals []float64 + for i := range m { + idx = append(idx, int64(i), int64(i%(n))) + vals = append(vals, 2+g.Unit()) + if i+1 < m && i%(n)+1 < n { + idx = append(idx, int64(i), int64(i%n+1)) + vals = append(vals, -0.5-0.5*g.Unit()) + } + } + indices, err := FromInts(idx, len(vals), 2) + if err != nil { + t.Fatal(err) + } + valArr, err := FloatsFromArray(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + coo, err := NewSparseCOO(indices, valArr, []int{m, n}) + if err != nil { + t.Fatal(err) + } + csr, err := CSRFromCOO(coo) + if err != nil { + t.Fatal(err) + } + xTrue := New(Float, n) + for i := range n { + xTrue.RawFloats()[i] = float64(i%5) - 2 + 0.25*float64(i%3) + } + b, err := csr.MatVec(xTrue) + if err != nil { + t.Fatal(err) + } + xl, infoL, err := SpLSQR(coo, b, 1e-12, 300, 0) + if err != nil { + t.Fatal(err) + } + xm, infoM, err := SpLSMR(coo, b, 1e-12, 300, 0) + if err != nil { + t.Fatal(err) + } + dm := mustA(t, randA(t, 66, 60).RawFloats(), 10, 6) + q, r, perm, rank, err := RRQR(dm) + if err != nil { + t.Fatal(err) + } + permF := make([]float64, len(perm)) + for i, p := range perm { + permF[i] = float64(p) + } + b6 := randA(t, 67, 10) + xr, err := SolveRRQR(dm, b6) + if err != nil { + t.Fatal(err) + } + const bn = 30 + var bIdx []int64 + var bVals []float64 + for i := range bn { + bIdx = append(bIdx, int64(i), int64(i)) + bVals = append(bVals, 4) + if i+1 < bn { + bIdx = append(bIdx, int64(i), int64(i+1)) + bVals = append(bVals, -1) + bIdx = append(bIdx, int64(i+1), int64(i)) + bVals = append(bVals, -1) + } + } + bIndices, err := FromInts(bIdx, len(bVals), 2) + if err != nil { + t.Fatal(err) + } + bValArr, err := FloatsFromArray(bVals, len(bVals)) + if err != nil { + t.Fatal(err) + } + bCoo, err := NewSparseCOO(bIndices, bValArr, []int{bn, bn}) + if err != nil { + t.Fatal(err) + } + factor, err := NewSparseCholesky(bCoo, SparseOrderingNatural) + if err != nil { + t.Fatal(err) + } + bv := New(Float, bn) + for i := range bn { + bv.RawFloats()[i] = float64(i%7) - 3 + } + before, err := factor.Solve(bv) + if err != nil { + t.Fatal(err) + } + upd := New(Float, bn) + upd.RawFloats()[7] = 2 + if err := factor.Update(upd); err != nil { + t.Fatal(err) + } + after, err := factor.Solve(bv) + if err != nil { + t.Fatal(err) + } + if err := factor.Downdate(upd); err != nil { + t.Fatal(err) + } + restored, err := factor.Solve(bv) + if err != nil { + t.Fatal(err) + } + return []any{xl, infoL.Criterion, infoL.ResidualNorm, float64(infoL.Iterations), + xm, infoM.Criterion, q, r, permF, float64(rank), xr, + before, after, restored} + }}, + {"optim-lp-qp-global2", func(t *testing.T) []any { + // The simplex through the two-sided rows, the active-set QP on + // an active wall, the seeded global pair on one rotated bowl, + // the nonlinear equality and the scalar bracket. + a := mustA(t, []float64{1, 1, 1, 0}, 2, 2) + xLP, vLP, err := MinimiseLinearRows(mustA(t, []float64{1, 1}, 2), + LinearConstraints{A: a, Lower: []float64{2, math.Inf(-1)}, Upper: []float64{math.Inf(1), 3}}, + LinearProgramOptions{}) + if err != nil { + t.Fatal(err) + } + h := mustA(t, []float64{2, 0, 0, 2}, 2, 2) + xQP, vQP, multQP, err := MinimiseQP(h, mustA(t, []float64{-2, -4}, 2), + LinearConstraints{A: mustA(t, []float64{1, 0}, 1, 2), + Lower: []float64{math.Inf(-1)}, Upper: []float64{1}}, + nil, QPOptions{}) + if err != nil { + t.Fatal(err) + } + cmaF := func(p *Array) (float64, error) { + var s float64 + w := []float64{4, 3, 2, 1} + for i := range 4 { + di := p.FloatAt(i) - float64(i+1) + s += w[i] * di * di + if i+1 < 4 { + s += 0.5 * di * (p.FloatAt(i+1) - float64(i+2)) + } + } + return s, nil + } + xCMA, vCMA, err := MinimiseCMAES(cmaF, mustA(t, []float64{2, -2, 1, -1}, 4), + CMAESOptions{Sigma0: 0.5, Generations: 300, Tolerance: 1e-10, Seed: 7}) + if err != nil { + t.Fatal(err) + } + xSA, vSA, err := MinimiseSimulatedAnnealing(cmaF, mustA(t, []float64{2, -2, 1, -1}, 4), + SimulatedAnnealingOptions{Steps: 20000, Tolerance: 1e-6, Seed: 42, AllowBudgetExit: true}) + if err != nil { + t.Fatal(err) + } + xNL, vNL, multNL, err := MinimiseNonlinearConstrained( + func(p *Array) (float64, error) { return p.FloatAt(0), nil }, + nil, mustA(t, []float64{0.5, 0.5}, 2), + NonlinearConstraints{Equalities: []func(*Array) (float64, error){ + func(p *Array) (float64, error) { + x, y := p.FloatAt(0), p.FloatAt(1) + return x*x + y*y - 1, nil + }, + }}, LBFGSOptions{}) + if err != nil { + t.Fatal(err) + } + root, err := FindRootBrent(func(x float64) float64 { return math.Cos(x) - x }, -10, 10, BrentOptions{}) + if err != nil { + t.Fatal(err) + } + broydenRHS := func(x *Array) (*Array, error) { + out, err := Zeros(Float, 2) + if err != nil { + return nil, err + } + xf, yf := x.FloatAt(0), x.FloatAt(1) + out.SetFloatAt(0, xf*xf+yf*yf-1) + out.SetFloatAt(1, xf-yf) + return out, nil + } + xBr, resBr, err := FindRootSystem(broydenRHS, mustA(t, []float64{2, 0.5}, 2), + RootSystemOptions{UseBroyden: true}) + if err != nil { + t.Fatal(err) + } + return []any{xLP, vLP, xQP, vQP, multQP, xCMA, vCMA, xSA, vSA, + xNL, vNL, multNL, root, xBr, resBr} + }}, + {"io-hdf5-write", func(t *testing.T) []any { + // The writer round trip: chunked, shuffled, deflated floats, a + // float32 matrix, wide integers, a nested group and attributes, + // read back through the unchanged reader in both superblocks. + dir := t.TempDir() + f64 := randA(t, 68, 32) + f32, err := FromFloat32s([]float32{1.5, -2.5, 3.25, 4, 5, 6}, 2, 3) + if err != nil { + t.Fatal(err) + } + wide, err := FromInts([]int64{5, -3, 1 << 40}, 3) + if err != nil { + t.Fatal(err) + } + nested, err := FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatal(err) + } + sets := []HDF5Dataset{ + {Path: "/signals/main", Shape: f64.Shape(), Values: f64, Attrs: map[string]string{"units": "m"}}, + {Path: "/image", Shape: f32.Shape(), Values: f32}, + {Path: "/wide", Shape: wide.Shape(), Values: wide}, + {Path: "/g/sub/deep", Shape: nested.Shape(), Values: nested}, + } + attrs := map[string]map[string]string{"/": {"title": "oracle"}, "/g": {"note": "middle"}} + classic := filepath.Join(dir, "classic.h5") + if err := SaveHDF5(classic, sets, attrs, HDF5WriteOptions{Gzip: 6, Shuffle: true}); err != nil { + t.Fatal(err) + } + back, err := LoadHDF5(classic) + if err != nil { + t.Fatal(err) + } + out := []any{} + for _, d := range back { + out = append(out, d.Path, d.Values) + for _, k := range slices.Sorted(maps.Keys(d.Attrs)) { + out = append(out, k, d.Attrs[k]) + } + } + latest := filepath.Join(dir, "latest.h5") + if err := SaveHDF5(latest, sets, attrs, HDF5WriteOptions{Latest: true}); err != nil { + t.Fatal(err) + } + backLatest, err := LoadHDF5(latest) + if err != nil { + t.Fatal(err) + } + out = append(out, len(backLatest) == len(back)) + for _, d := range backLatest { + if d.Path == "/signals/main" { + out = append(out, d.Values) + } + } + return out + }}, + {"io-hdf5-narrow-write-roundtrip", func(t *testing.T) []any { + // The seven narrow dtypes through the writer and back, the + // values at each dtype's extremes: bool through the boolean + // enumeration convention, every integer at its stored width and + // signedness. The same sets are written contiguous, in the + // latest layout and chunked with shuffle and deflate, and every + // read-back must land the written dtype and values. + dir := t.TempDir() + sets := []HDF5Dataset{} + must := func(values *Array, err error) *Array { + if err != nil { + t.Fatal(err) + } + return values + } + put := func(path string, values *Array) { + sets = append(sets, HDF5Dataset{Path: path, Shape: values.Shape(), Values: values}) + } + put("/bool", must(FromBools([]bool{false, true, true, false}, 2, 2))) + put("/i8", must(FromInt8s([]int8{-128, 127, 0, -1}, 2, 2))) + put("/u8", must(FromUint8s([]uint8{0, 255, 7, 128}, 2, 2))) + put("/i16", must(FromInt16s([]int16{-32768, 32767, 0, -2}, 2, 2))) + put("/u16", must(FromUint16s([]uint16{0, 65535, 5, 32768}, 2, 2))) + put("/i32", must(FromInt32s([]int32{-2147483648, 2147483647, 0, -3}, 2, 2))) + put("/u32", must(FromUint32s([]uint32{0, 4294967295, 9, 2147483648}, 2, 2))) + out := []any{} + for _, c := range []struct { + file string + opts HDF5WriteOptions + }{ + {"contiguous.h5", HDF5WriteOptions{}}, + {"latest.h5", HDF5WriteOptions{Latest: true}}, + {"chunked.h5", HDF5WriteOptions{Gzip: 1, Shuffle: true, ChunkBytes: 8}}, + } { + path := filepath.Join(dir, c.file) + if err := SaveHDF5(path, sets, nil, c.opts); err != nil { + t.Fatal(err) + } + back, err := LoadHDF5(path) + if err != nil { + t.Fatal(err) + } + out = append(out, c.file) + for _, d := range back { + out = append(out, d.Path, d.Values) + } + } + return out + }}, + {"io-hdf5-uint8-roundtrip", func(t *testing.T) []any { + // A uint8 dataset through the HDF5 read path: the fixture pins + // the reader's decode directly, hand-built in the classic + // layout, and must land uint8 on the values below. The writer + // path is covered by io-hdf5-narrow-write-roundtrip above, + // which writes every narrow dtype and reads it back. + path := filepath.Join(t.TempDir(), "u8.h5") + file := oracleHDF5Narrow(oracleHDF5Fixed(1, false), []uint64{5}, + []byte{0, 1, 200, 255, 42}) + if err := os.WriteFile(path, file, 0o644); err != nil { + t.Fatal(err) + } + sets, err := LoadHDF5(path) + if err != nil { + t.Fatal(err) + } + return []any{len(sets), sets[0].Path, sets[0].Values} + }}, + {"io-hdf5-int16-roundtrip", func(t *testing.T) []any { + // An int16 dataset through the same hand-built fixture path: + // the read-back must land int16 on the extremes below. + path := filepath.Join(t.TempDir(), "i16.h5") + file := oracleHDF5Narrow(oracleHDF5Fixed(2, true), []uint64{4}, + []byte{0xfd, 0xff, 0x00, 0x80, 0xff, 0x7f, 0x07, 0x00}) + if err := os.WriteFile(path, file, 0o644); err != nil { + t.Fatal(err) + } + sets, err := LoadHDF5(path) + if err != nil { + t.Fatal(err) + } + return []any{len(sets), sets[0].Path, sets[0].Values} + }}, + {"io-hdf5-bool-roundtrip", func(t *testing.T) []any { + // A boolean enumeration dataset, the convention HDF5 writers + // carry booleans in: the read-back must land core.Bool. + path := filepath.Join(t.TempDir(), "bool.h5") + file := oracleHDF5Narrow(oracleHDF5EnumBool(), []uint64{3}, []byte{1, 0, 1}) + if err := os.WriteFile(path, file, 0o644); err != nil { + t.Fatal(err) + } + sets, err := LoadHDF5(path) + if err != nil { + t.Fatal(err) + } + return []any{len(sets), sets[0].Path, sets[0].Values} + }}, + {"io-netcdf-classic-natives", func(t *testing.T) []any { + // A NetCDF classic file carrying byte, short, int, char, float + // and double values, read natively. The NetCDF writer stores + // float64, float32 and int64 only, so the fixture is written + // here as classic-format bytes and the oracle pins the landings. + path := filepath.Join(t.TempDir(), "natives.nc") + if err := os.WriteFile(path, oracleNetCDFNatives(), 0o644); err != nil { + t.Fatal(err) + } + _, vars, _, err := LoadNetCDF(path) + if err != nil { + t.Fatal(err) + } + out := []any{len(vars)} + for _, v := range vars { + out = append(out, v.Name, v.Values) + } + return out + }}, + {"optim-lm-fit", func(t *testing.T) []any { + // The Levenberg-Marquardt fit through the Fit surface: an + // exact exponential decay with a requested covariance, then a + // diagonal-Sigma run whose first observation carries a quarter + // of the variance, so the whitening moves the answer. The + // digest pins the parameters, both chi2 values, the statuses + // and the covariance. + xs := []float64{0, 1, 2, 3, 4, 5, 6, 7} + ys := make([]float64, len(xs)) + for i, x := range xs { + ys[i] = 2.5*math.Exp(-0.7*float64(x)) + 0.3 + } + residual := func(p *Array) (*Array, error) { + out := New(Float, len(xs)) + for i := range len(xs) { + out.RawFloats()[i] = ys[i] - (p.FloatAt(0)*math.Exp(-p.FloatAt(1)*float64(xs[i])) + p.FloatAt(2)) + } + return out, nil + } + p0, err := FromFloats([]float64{2, 0.5, 0.1}, 3) + if err != nil { + t.Fatal(err) + } + res, err := LevenbergMarquardtFit(residual, p0, LMOptions{RequestCovariance: true}) + if err != nil { + t.Fatal(err) + } + if res.Status != FitConverged { + t.Fatalf("status %d, want FitConverged", res.Status) + } + out := []any{res.Parameters, res.Chi2, int(res.Status), res.Covariance} + sigma, err := FromFloats([]float64{0.25, 1, 1, 1, 1, 1, 1, 1}, 8) + if err != nil { + t.Fatal(err) + } + wres, err := LevenbergMarquardtFit(residual, p0, LMOptions{Sigma: sigma}) + if err != nil { + t.Fatal(err) + } + if wres.Status != FitConverged { + t.Fatalf("weighted status %d, want FitConverged", wres.Status) + } + return append(out, wres.Parameters, wres.Chi2, int(wres.Status)) + }}, + {"core-cosm1", func(t *testing.T) []any { + // The small-argument cosine deficit across its regimes: the + // digest pins the factored series from the underflowed end to + // the crossover and the direct subtraction above it. + a, err := FromFloats([]float64{ + 1e-300, 1e-12, 1e-9, 1e-8, 1e-4, 0.01, 0.5, + math.Pi / 4, 0.79, 1, 2, -0.3, -2.5, 1e6, + }, 14) + if err != nil { + t.Fatal(err) + } + v, err := Cosm1(a) + if err != nil { + t.Fatal(err) + } + return []any{v} + }}, + {"core-substream", func(t *testing.T) []any { + // The substream family and the scalar splitmix64 chain: the + // digest pins the first draws of three members of one seed's + // family beside the chain of states the scalar surface hands + // out, so a change to the seeding moves loudly. + out := []any{} + for _, i := range []int{0, 1, 7} { + g, err := Substream(42, i) + if err != nil { + t.Fatal(err) + } + a, err := Floats(g, 8) + if err != nil { + t.Fatal(err) + } + out = append(out, a) + } + state, v := Splitmix64(0) + for range 3 { + out = append(out, fmt.Sprintf("%#016x", state), fmt.Sprintf("%#016x", v)) + state, v = Splitmix64(state) + } + return out + }}, + {"signal-chirp", func(t *testing.T) []any { + // The chirp synthesiser: the digest pins a rising sweep, a + // falling one and the constant tone that must be the plain + // sine. + rise, err := Chirp(64, 5, 60, 128) + if err != nil { + t.Fatal(err) + } + fall, err := Chirp(16, 900, 100, 8000) + if err != nil { + t.Fatal(err) + } + tone, err := Chirp(32, 11, 11, 256) + if err != nil { + t.Fatal(err) + } + return []any{rise, fall, tone} + }}, + {"core-besseljreal", func(t *testing.T) []any { + // The real-order Bessel J across its regimes: the series at a + // small argument, the climb seeded from the expansion, the + // fractional Miller walk past the argument, and the integer + // delegation, pinned beside the points themselves. + type point struct { + nu, x float64 + } + pts := []point{ + {0.5, 2}, {2.5, 0.05}, {1.5, 12.1}, {2.3, 15}, + {7.7, 40}, {20.5, 15}, {1.7, 1000}, {3, 20}, + } + out := []any{len(pts)} + for _, p := range pts { + v, err := BesselJRealOrder(p.nu, p.x) + if err != nil { + t.Fatal(err) + } + out = append(out, v) + } + return out + }}, + {"integrate-filon", func(t *testing.T) []any { + // The Filon quadrature across its regimes: polynomial + // amplitudes, where the construction is exact whatever the + // carrier does, an exponential amplitude under a carrier of + // five hundred, and the zero-frequency degeneration beside + // them. + poly := func(x float64) (float64, error) { return x, nil } + c1, s1, err := IntegrateFilon(poly, 2, 7, 500, FilonOptions{}) + if err != nil { + t.Fatal(err) + } + exp := func(x float64) (float64, error) { return math.Exp(0.5 * x), nil } + c2, s2, err := IntegrateFilon(exp, 2, 7, 500, FilonOptions{}) + if err != nil { + t.Fatal(err) + } + c3, s3, err := IntegrateFilon(exp, 0, 3, 0, FilonOptions{}) + if err != nil { + t.Fatal(err) + } + return []any{c1, s1, c2, s2, c3, s3} + }}, + {"core-narrow", func(t *testing.T) []any { + // The narrow integer dtypes end to end: every width over the + // corners of its own range, promotion across widths and across + // the class boundary, the exact int64 folds, the bool masks the + // comparisons answer and the selection and logic they feed. + i8, err := Int8sFromArray([]int8{-128, -1, 0, 1, 127, -128, 127, 0, 3, -7}, 10) + if err != nil { + t.Fatal(err) + } + u8, err := Uint8sFromArray([]uint8{0, 1, 127, 128, 255, 254, 3, 0, 9, 200}, 10) + if err != nil { + t.Fatal(err) + } + i16, err := Int16sFromArray([]int16{-32768, -1, 0, 1, 32767, -32768, 32767, 0, 500, -900}, 10) + if err != nil { + t.Fatal(err) + } + u16, err := Uint16sFromArray([]uint16{0, 1, 32767, 32768, 65535, 65534, 7, 0, 40, 60000}, 10) + if err != nil { + t.Fatal(err) + } + i32, err := Int32sFromArray([]int32{-2147483648, -1, 0, 1, 2147483647, -2147483648, 2147483647, 0, 70000, -300}, 10) + if err != nil { + t.Fatal(err) + } + u32, err := Uint32sFromArray([]uint32{0, 1, 2147483647, 2147483648, 4294967295, 4294967294, 11, 0, 5, 4000000000}, 10) + if err != nil { + t.Fatal(err) + } + cross1, err := Add(i8, u16) + if err != nil { + t.Fatal(err) + } + cross2, err := Sub(i32, u8) + if err != nil { + t.Fatal(err) + } + cross3, err := Mul(i16, u8) + if err != nil { + t.Fatal(err) + } + cross4, err := Add(u32, i16) + if err != nil { + t.Fatal(err) + } + ltMask, err := Lt(i8, u16) + if err != nil { + t.Fatal(err) + } + eqMask, err := EqI(u8, 1) + if err != nil { + t.Fatal(err) + } + leMask, err := LeI(i32, -1) + if err != nil { + t.Fatal(err) + } + geMask, err := GeF(i8, 0.5) + if err != nil { + t.Fatal(err) + } + ltfMask, err := LtF(i16, -0.5) + if err != nil { + t.Fatal(err) + } + joined, err := And(ltMask, eqMask) + if err != nil { + t.Fatal(err) + } + flipped, err := Not(eqMask) + if err != nil { + t.Fatal(err) + } + picked, err := Select(i8, joined) + if err != nil { + t.Fatal(err) + } + chosen, err := Where(eqMask, u8, i8) + if err != nil { + t.Fatal(err) + } + widened, err := Astype(u32, Float) + if err != nil { + t.Fatal(err) + } + wide, err := FromInts([]int64{2147483648, -2147483649, -1}, 3) + if err != nil { + t.Fatal(err) + } + _, narrowErr := Astype(wide, Int32) + if narrowErr == nil { + t.Fatal("Astype narrowed 2^31 into int32") + } + inRange, err := FromInts([]int64{2147483647, -2147483648, -1}, 3) + if err != nil { + t.Fatal(err) + } + narrowed, err := Astype(inRange, Int32) + if err != nil { + t.Fatal(err) + } + count, err := CountNonzero(leMask) + if err != nil { + t.Fatal(err) + } + minU8, err := Min(u8) + if err != nil { + t.Fatal(err) + } + maxI32, err := Max(i32) + if err != nil { + t.Fatal(err) + } + return []any{i8, u8, i16, u16, i32, u32, + cross1, cross2, cross3, cross4, + Sum(i8), Sum(i16), Sum(u32), minU8, maxI32, + ltMask, eqMask, leMask, geMask, ltfMask, joined, flipped, + picked, chosen, widened, narrowed, narrowErr != nil, count} + }}, +} + +// hostileNetCDF builds a CDF-1 file whose single variable spans three +// dimensions of 2^21, 2^20 and 2^20 elements: their product times the +// 8-byte element width wraps to zero in 64-bit arithmetic, so a reader +// that multiplies before it bounds would allocate what the header +// claims. The file is 112 bytes. +func hostileOracleHeader() []byte { + u32 := func(v uint32) []byte { + var b [4]byte + binary.BigEndian.PutUint32(b[:], v) + return b[:] + } + name := func(s string) []byte { + b := u32(uint32(len(s))) + b = append(b, s...) + for len(s)%4 != 0 { + b = append(b, 0) + s += " " + } + return b + } + var b []byte + b = append(b, 'C', 'D', 'F', 1) + b = append(b, u32(0)...) // numrecs + b = append(b, u32(10)...) // NC_DIMENSION + b = append(b, u32(3)...) + for i, l := range []uint32{1 << 21, 1 << 20, 1 << 20} { + b = append(b, name(string(rune('a'+i)))...) + b = append(b, u32(l)...) + } + b = append(b, u32(0)...) // absent attribute list + b = append(b, u32(0)...) + b = append(b, u32(11)...) // NC_VARIABLE + b = append(b, u32(1)...) + b = append(b, name("v")...) + b = append(b, u32(3)...) + b = append(b, u32(0)...) + b = append(b, u32(1)...) + b = append(b, u32(2)...) + b = append(b, u32(0)...) // absent variable attributes + b = append(b, u32(0)...) + b = append(b, u32(6)...) // NC_DOUBLE + b = append(b, u32(0)...) // vsize + b = append(b, u32(0)...) // begin + return b +} + +// hostileFITSTable builds a BINTABLE header with a zero-width TFORM and +// no data block: the row arithmetic must not be able to divide by the +// row size, and the truncated file must be refused. +func hostileOracleTable() []byte { + var b []byte + for _, body := range []string{ + "XTENSION= 'BINTABLE'", "BITPIX = 8", "NAXIS = 2", + "NAXIS1 = 4", "NAXIS2 = 2", "TFIELDS = 1", + "TFORM1 = 'X '", "END", + } { + b = append(b, []byte(body)...) + for len(b)%80 != 0 { + b = append(b, ' ') + } + } + return b +} + +// oracleHDF5Fixed renders a version 1 fixed-point datatype message of +// the given element width and signedness, with the bit offset and bit +// precision the HDF5 file format specification's fixed-point property +// table defines behind the header. +func oracleHDF5Fixed(size uint32, signed bool) []byte { + m := []byte{0x10, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0} // version 1, class 0 + if signed { + m[1] = 0x08 // class bit field: bit 3 marks two's complement + } + binary.LittleEndian.PutUint32(m[4:], size) + binary.LittleEndian.PutUint16(m[10:], uint16(8*size)) // bit precision + return m +} + +// oracleHDF5EnumBool renders the boolean enumeration datatype message: +// class 8 over a one-byte unsigned base, the member names padded in +// their own fields to multiples of eight bytes, the member values +// packed behind them, exactly as the HDF5 file format specification's +// enumeration class defines. +func oracleHDF5EnumBool() []byte { + m := []byte{0x18, 2, 0, 0, 1, 0, 0, 0} // class 8, two members, size 1 + m = append(m, oracleHDF5Fixed(1, false)...) // the base type + for _, name := range []string{"TRUE", "FALSE"} { + start := len(m) + m = append(m, name...) + m = append(m, 0) + for (len(m)-start)%8 != 0 { + m = append(m, 0) + } + } + return append(m, 1, 0) // TRUE = 1, FALSE = 0 +} + +// oracleHDF5Narrow builds a minimal classic-layout HDF5 file carrying +// one dataset: a version 0 superblock of eight-byte addresses and +// lengths, a version 1 object header with the given datatype message, +// contiguous storage and the payload bytes. The io tests build their +// hostile fixtures the same way; these bytes are fixed, so the digest +// of what the reader lands is stable across runs. +func oracleHDF5Narrow(dtypeMsg []byte, dims []uint64, payload []byte) []byte { + const headerAt, dataAt = 96, 448 + align8 := func(n int) int { return (n + 7) / 8 * 8 } + n := dataAt + len(payload) + f := make([]byte, n) + copy(f, []byte{0x89, 'H', 'D', 'F', '\r', '\n', 0x1a, '\n'}) + f[8] = 0 // superblock version 0, the classic layout + f[13], f[14] = 8, 8 // address and length sizes + put := func(at int, v uint64) { binary.LittleEndian.PutUint64(f[at:], v) } + put(32, math.MaxUint64) // free space undefined + put(40, uint64(n)) // end of file + put(48, math.MaxUint64) // driver information undefined + put(64, headerAt) // the root object header + space := make([]byte, 8+8*len(dims)) + space[0] = 1 // dataspace version 1 + space[1] = byte(len(dims)) + for i, d := range dims { + binary.LittleEndian.PutUint64(space[8+8*i:], d) + } + lay := make([]byte, 18) + lay[0], lay[1] = 3, 1 // layout version 3, contiguous + binary.LittleEndian.PutUint64(lay[2:], dataAt) + binary.LittleEndian.PutUint64(lay[10:], uint64(len(payload))) + msgs := []struct { + typ uint16 + body []byte + }{ + {1, space}, // dataspace + {3, dtypeMsg}, + {8, lay}, // data layout + } + off := headerAt + f[off] = 1 // object header version 1 + binary.LittleEndian.PutUint16(f[off+2:], uint16(len(msgs))) + binary.LittleEndian.PutUint32(f[off+4:], 1) // reference count + region := off + 16 + size := 0 + for _, m := range msgs { + size = align8(size + 8 + len(m.body)) + } + binary.LittleEndian.PutUint32(f[off+8:], uint32(size)) + for _, m := range msgs { + binary.LittleEndian.PutUint16(f[region:], m.typ) + binary.LittleEndian.PutUint16(f[region+2:], uint16(len(m.body))) + copy(f[region+8:], m.body) + region += align8(8 + len(m.body)) + } + copy(f[dataAt:], payload) + return f +} + +// oracleNetCDFNatives builds a CDF-1 file by hand carrying one +// variable per classic type code: NC_BYTE, NC_SHORT, NC_INT, NC_CHAR, +// NC_FLOAT and NC_DOUBLE over one shared four-element dimension, the +// integers at the extremes of every width and the floats at +// recognisable values. The NetCDF writer stores float64, float32 and +// int64 only, so the fixture bytes live here exactly as the foreign +// fixtures the io package carries. +func oracleNetCDFNatives() []byte { + u32 := func(v uint32) []byte { + var b [4]byte + binary.BigEndian.PutUint32(b[:], v) + return b[:] + } + name := func(s string) []byte { + b := u32(uint32(len(s))) + b = append(b, s...) + for len(b)%4 != 0 { + b = append(b, 0) + } + return b + } + type classicVar struct { + name string + code uint32 + data []byte + } + vars := []classicVar{ + {"b", 1, []byte{0x80, 0xff, 0x00, 0x7f}}, // NC_BYTE: -128, -1, 0, 127 + {"s", 3, []byte{0x80, 0x00, 0xff, 0xff, 0x00, 0x01, 0x7f, 0xff}}, // NC_SHORT extremes + {"i", 4, []byte{0x80, 0, 0, 0, 0xff, 0xff, 0xff, 0xff, 0, 0, 0, 1, 0x7f, 0xff, 0xff, 0xff}}, // NC_INT + {"c", 2, []byte{0x00, 0x41, 0xc8, 0xff}}, // NC_CHAR: raw bytes + {"f", 5, []byte{ // NC_FLOAT: 1.5, -2.5, 0, 3.25 + 0x3f, 0xc0, 0x00, 0x00, 0xc0, 0x20, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x40, 0x50, 0x00, 0x00}}, + {"d", 6, []byte{ // NC_DOUBLE: 1.5, -2.5, 0, 4 + 0x3f, 0xf8, 0, 0, 0, 0, 0, 0, 0xc0, 0x04, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0x40, 0x10, 0, 0, 0, 0, 0, 0}}, + } + var b []byte + b = append(b, 'C', 'D', 'F', 1) + b = append(b, u32(0)...) // numrecs + b = append(b, u32(10)...) // NC_DIMENSION + b = append(b, u32(1)...) // one dimension + b = append(b, name("n")...) + b = append(b, u32(4)...) // length 4 + b = append(b, u32(0)...) // gatt_list ABSENT + b = append(b, u32(0)...) // gatt_list ABSENT + b = append(b, u32(11)...) // NC_VARIABLE + b = append(b, u32(uint32(len(vars)))...) // six variables + // The header ends where the first payload begins: one 36-byte + // variable record per variable behind the fixed prefix. + begin := len(b) + len(vars)*36 + off := begin + for _, v := range vars { + b = append(b, name(v.name)...) + b = append(b, u32(1)...) // rank 1 + b = append(b, u32(0)...) // dimid 0 + b = append(b, u32(0)...) // variable attributes ABSENT + b = append(b, u32(0)...) // variable attributes ABSENT + b = append(b, u32(v.code)...) + b = append(b, u32(uint32(len(v.data)))...) // vsize + b = append(b, u32(uint32(off))...) + off += len(v.data) + } + for _, v := range vars { + b = append(b, v.data...) + } + return b +} + +// pinnedOracleDigests is the recorded behaviour of the harness, keyed +// by runtime.GOOS/GOARCH plus the code generation the binary was built +// with: floating-point kernels may differ between architectures (FMA +// availability, libm corners), and GOAMD64 decides whether the compiler +// may contract a multiply and an add into one fused operation, which +// changes the last bits without changing the algorithm. A deliberate +// algorithm change re-records the affected entries on the configuration +// it was made on; an accidental one fails the test there. A build whose +// combination has no block skips loudly until it is recorded, so +// determinism is measured per configuration, never claimed in general. +// +// The portable build is the product and the only build, and its +// digests are recorded at the toolchain default level (v1): the +// compiler has no auto-vectoriser, so a pinned higher level would buy +// only scalar FMA contraction, which the bit-pinned kernels already +// suppress by spelling. A build pinned to another GOAMD64 level has no +// block and skips with the recording instructions rather than failing. +// +// A re-recorded digest is sometimes the point itself: an old digest +// can have pinned a bug, and three entries exist because their +// predecessors did. signal-welch-savgol's magnitude goes through +// math.Hypot, so it no longer rounds through a squared sum. +// optim-roots-minima and optim-constrained-minima carry the L-BFGS +// line search that no longer returns its start point when a stiff +// objective rejects every trial step, the Nelder-Mead that no longer +// stops on a simplex lying on one level set, and MinimiseConstrained +// returning f at the answer instead of the augmented Lagrangian value. +var pinnedOracleDigests = map[string]map[string]string{ + "linux/amd64+v1": { + "spmd-shards": "0bec5217a9a79b182241ac0937f32e9d20c89a5b392fc797ab3787a03fc8c9ed", + "io-hdf5-narrow-write-roundtrip": "f2e87eaadcbab846e4acf5473b3023faca3015cd3b3f0919c4c3792114589ee9", + "io-hdf5-uint8-roundtrip": "54550518525945c92d970a38d2d81a5005fb8a899b597819c6060610ae77dd54", + "io-hdf5-int16-roundtrip": "b4163c6e9ed7a04ddb4a2328d394afe074237533738d19aa5ecaa68bff9e2685", + "io-hdf5-bool-roundtrip": "fa56fd9a421d2e78f17711e717f851c4b19dc1f6946d5c4e1dddad85cc42d347", + "io-netcdf-classic-natives": "c0219b60aa963176ce509f5633c22b96bc070529db1035d53a643b67d614d01b", + "core-elementwise": "0812ac01c62b084326d0e590941a46327136e9022a101fad68998781b68f7014", + "core-matmul-einsum": "a62b07d724f06603fe4c5ebf3d448b61d7638b93d1b58df605d94096d30ba4b4", + "core-sort-argsort": "24fc532b4140bc6633b5fabb8a19c708c1919932c75ae8fce048c413c94aff0f", + "core-reductions": "764eac75a3bcc1a38f7975ed4de1e5df39d70889dfb845507ef1cfe1ff456eac", + "core-shape": "c84940b2aff4e4a25e7a4292fa6d76c2a0306f328ba334e3a6085843d94fcc3b", + "core-fft": "1848e23763a0a72f282eff45c1cc54a3d2486f4ff53b12b4a8f462bcdc773061", + "core-quasirandom": "000620721b8ce6fff07287da6301b89dbcdaa7293d68a0783aae60fa4370aa6d", + "linalg-solve-inv-det": "5b73c4a57a6221f2b5abd3b6769d6aa13723db0ada0373e870d551f795c61bf3", + "linalg-factorisations": "f70c03b3759703a6fcd1ff1ef3cbab024acd000781e2065331c18f95c45b4d15", + "linalg-sparse-cholesky": "da49d6554341ffbf40500a99f241e0133528c2e6fadfc385cfa6d7d59685da28", + "linalg-sparse-lu": "3ccf6e7064b410c3c6df0a6b2697b4c66f3722aec99491ed9298db6212d65fbb", + "integrate-fem-poisson": "065cb65bb7fba9838ed6f84b703d7ffb69d877564bfbe006905e81fdea6b9175", + "linalg-sparse-complex": "d81a8dd242a23ce7026ded53c44a4fea3a58d9b20c8be4ed2c03147c740e2130", + "signal-welch-savgol": "36d6581ac2b06191b2d6eb743e042283d8cd0cbc187c54b6d40f490eef5984a4", + "signal-wavelets": "f324710c1f67c8e5bb099731c8647c944744a62e9219ef8fdbb070412df434ab", + "signal-lombscargle-conv": "4c3f3eb0f292b978d701dcea384ccbfc26be06acff348756ec0bda19c4cca20e", + "signal-filters-stencils": "00bccd2f3b5d61f07c3383280dc8111897e1eaf57a6e16a0d70081a703a8b158", + "stats-moments-quantiles": "61d196193641395bbb855685f2727cfe29521848bba111d89de9f480ef7264cc", + "stats-correlation-regression": "f398fea69ae522ffaafdf0d307edd02d657aaf7a452c6c75d174307becbf7218", + "stats-poisson-regression": "e385d9fda18acef4669ac1d1f1e8b12a5abeb682eff37a06f6ec140dd7b9f91a", + "stats-distributions": "8986f622708f3f1f28c1c00fed9225014c17a90a8e0714e2ce6e2f181ee03321", + "integrate-quadrature-ode": "66dd47bbc30de692902937078808da5322c76e6de1facaf5099fd72344f248a5", + "optim-roots-minima": "9af044a7e71416f0ab1f8dcbd4073f199163791ad28779b59c7faabdf982e6f4", + "optim-bounded-minima": "368dbb82d31c5e53c7dad9f944921e741ccccedafa7fda8d8b180f9daa58628e", + "optim-constrained-minima": "884020c56542ac84e5460f0cbd171f64904f86254ebbae6479e3bd6662cc6662", + "grad-backward": "0f3a1be94160019b796b457aa42531817685b7bc478ed041cc6c73ec5810b744", + "io-roundtrips": "de1a36e2b2e9235a40af37b4ee5823e526f6a46583c34af6245dbbd322e3f1e3", + "core-views-and-int-precision": "ca3c789ead2894ed2c53416f128b8aa665aed9dc3d4d67e38ae9d540a1ac5ccd", + "io-hdf5-fixture": "01e96137069f7d9dd7d4cd2e236de8f5fbcf4ba6121a1f121f21c041197cf117", + "io-hostile-inputs": "9dcf97a184f32623d11a73124ceb99a5709b083721e878a16d78f596718ba7b2", + "core-float16": "6af37847cc6187b64f2bb20b7fefdd3b28a039e42988708d942c41bf62576949", + "signal-kalman-arma": "874b9c7bfd3dae3111cb49576a51913972ad572963a2ab7ad1eaa3187f15196b", + "signal-windows-filtfilt-dwt": "b435d694350e7a7f06303ded85c2b63e0e562795cf69686f2ae854f2c6dc81c5", + "stats-distributions2-multipletest": "6d2f2661fa3a2bdaefbdf16f3d26b41ef69b29c93347e16e6df8542f30b14078", + "stats-regression2": "cabb52d66ccb5d92b8d0885f5a5a166c1814d364dc75a4e8cac89819820edba3", + "stats-unsupervised": "c8a17031b7133d58a6b6b6ec21a5624281d8f291717749869e41251b3e8ef68f", + "integrate-stiff-solvers": "48276e65576eeefa8d7c3da8da42ee782441462ce34cf6bb814e494c0e00668f", + "integrate-fem3d": "8323b5c89ac9187441cfac694193ef6394aff0508e71b53052d1b76fd59a535b", + "linalg-sparse-lsqr-rrqr-update": "806455d7e53765a272afca8cfef90aa233ab78e62e4ad682c085ff7e117231fc", + "optim-lp-qp-global2": "510ce65940755a164784ed4577502db2b60604aa9cbe184f22b2ab2a4c03bd57", + "io-hdf5-write": "21482614860f885f277b05666924ab3b63a2c55b8b24734cbadc4577f2c7ac90", + "optim-lm-fit": "cdbf15c3a4c5615e1a364a9fe4e5ec431d0cb65d9744cab52cbca66af51c58cf", + "core-cosm1": "f995296d5c1f285ec6f2e719ee62f8eaf7c9a8e751cea8702cbe16b3df66ddb2", + "core-substream": "ce49fae36f5e15b37aacc25d62b97c9f800d51aa08dee553cee1a5b6e1c01798", + "signal-chirp": "ff9af5783d32eed88bf2658bde5660c247a0bd30a699c403f1755da290830a4b", + "core-besseljreal": "49698a21228f091079cd65ae94cda134b16c3b0bec58a47642f6ad3fec7461d7", + "integrate-filon": "4d5242520a5c144e4bdf111df83c2f2faf4fadce0859f87738acab91e3ed93e0", + "core-narrow": "80913507636c4aabcdc57e5f90a75cd18499c86e244d84006650de3e1273da92", + "stats-contingency": "9ac61a7a29227ddae02861076c85a669fd467f466a8c1118c52965ad4ac20a1b", + "stats-mixedmodel": "6db27ea7908863a1f2d0725325d8a72a93a5ec83ad3882bca28a5d3d2b614227", + "stats-hmm": "779e2c9cb8ee1e1871b17a04a42b2d574ff3de6132430bae8c58be502d450b75", + "stats-hierarchy": "c4171ae9c61b76c676d016356ae319c5e3c15c0e32f7d0c26f0b1b38700cb96e", + }, + "linux/amd64+v1+race": { + "spmd-shards": "0bec5217a9a79b182241ac0937f32e9d20c89a5b392fc797ab3787a03fc8c9ed", + "io-hdf5-narrow-write-roundtrip": "f2e87eaadcbab846e4acf5473b3023faca3015cd3b3f0919c4c3792114589ee9", + "io-hdf5-uint8-roundtrip": "54550518525945c92d970a38d2d81a5005fb8a899b597819c6060610ae77dd54", + "io-hdf5-int16-roundtrip": "b4163c6e9ed7a04ddb4a2328d394afe074237533738d19aa5ecaa68bff9e2685", + "io-hdf5-bool-roundtrip": "fa56fd9a421d2e78f17711e717f851c4b19dc1f6946d5c4e1dddad85cc42d347", + "io-netcdf-classic-natives": "c0219b60aa963176ce509f5633c22b96bc070529db1035d53a643b67d614d01b", + "core-elementwise": "0812ac01c62b084326d0e590941a46327136e9022a101fad68998781b68f7014", + "core-matmul-einsum": "a62b07d724f06603fe4c5ebf3d448b61d7638b93d1b58df605d94096d30ba4b4", + "core-sort-argsort": "24fc532b4140bc6633b5fabb8a19c708c1919932c75ae8fce048c413c94aff0f", + "core-reductions": "764eac75a3bcc1a38f7975ed4de1e5df39d70889dfb845507ef1cfe1ff456eac", + "core-shape": "c84940b2aff4e4a25e7a4292fa6d76c2a0306f328ba334e3a6085843d94fcc3b", + "core-fft": "1848e23763a0a72f282eff45c1cc54a3d2486f4ff53b12b4a8f462bcdc773061", + "core-quasirandom": "000620721b8ce6fff07287da6301b89dbcdaa7293d68a0783aae60fa4370aa6d", + "linalg-solve-inv-det": "5b73c4a57a6221f2b5abd3b6769d6aa13723db0ada0373e870d551f795c61bf3", + "linalg-factorisations": "f70c03b3759703a6fcd1ff1ef3cbab024acd000781e2065331c18f95c45b4d15", + "linalg-sparse-cholesky": "da49d6554341ffbf40500a99f241e0133528c2e6fadfc385cfa6d7d59685da28", + "linalg-sparse-lu": "3ccf6e7064b410c3c6df0a6b2697b4c66f3722aec99491ed9298db6212d65fbb", + "integrate-fem-poisson": "065cb65bb7fba9838ed6f84b703d7ffb69d877564bfbe006905e81fdea6b9175", + "linalg-sparse-complex": "d81a8dd242a23ce7026ded53c44a4fea3a58d9b20c8be4ed2c03147c740e2130", + "signal-welch-savgol": "36d6581ac2b06191b2d6eb743e042283d8cd0cbc187c54b6d40f490eef5984a4", + "signal-wavelets": "f324710c1f67c8e5bb099731c8647c944744a62e9219ef8fdbb070412df434ab", + "signal-lombscargle-conv": "4c3f3eb0f292b978d701dcea384ccbfc26be06acff348756ec0bda19c4cca20e", + "signal-filters-stencils": "00bccd2f3b5d61f07c3383280dc8111897e1eaf57a6e16a0d70081a703a8b158", + "stats-moments-quantiles": "61d196193641395bbb855685f2727cfe29521848bba111d89de9f480ef7264cc", + "stats-correlation-regression": "f398fea69ae522ffaafdf0d307edd02d657aaf7a452c6c75d174307becbf7218", + "stats-poisson-regression": "e385d9fda18acef4669ac1d1f1e8b12a5abeb682eff37a06f6ec140dd7b9f91a", + "stats-distributions": "8986f622708f3f1f28c1c00fed9225014c17a90a8e0714e2ce6e2f181ee03321", + "integrate-quadrature-ode": "66dd47bbc30de692902937078808da5322c76e6de1facaf5099fd72344f248a5", + "optim-roots-minima": "9af044a7e71416f0ab1f8dcbd4073f199163791ad28779b59c7faabdf982e6f4", + "optim-bounded-minima": "368dbb82d31c5e53c7dad9f944921e741ccccedafa7fda8d8b180f9daa58628e", + "optim-constrained-minima": "884020c56542ac84e5460f0cbd171f64904f86254ebbae6479e3bd6662cc6662", + "grad-backward": "0f3a1be94160019b796b457aa42531817685b7bc478ed041cc6c73ec5810b744", + "io-roundtrips": "de1a36e2b2e9235a40af37b4ee5823e526f6a46583c34af6245dbbd322e3f1e3", + "core-views-and-int-precision": "ca3c789ead2894ed2c53416f128b8aa665aed9dc3d4d67e38ae9d540a1ac5ccd", + "io-hdf5-fixture": "01e96137069f7d9dd7d4cd2e236de8f5fbcf4ba6121a1f121f21c041197cf117", + "io-hostile-inputs": "9dcf97a184f32623d11a73124ceb99a5709b083721e878a16d78f596718ba7b2", + "core-float16": "6af37847cc6187b64f2bb20b7fefdd3b28a039e42988708d942c41bf62576949", + "signal-kalman-arma": "874b9c7bfd3dae3111cb49576a51913972ad572963a2ab7ad1eaa3187f15196b", + "signal-windows-filtfilt-dwt": "b435d694350e7a7f06303ded85c2b63e0e562795cf69686f2ae854f2c6dc81c5", + "stats-distributions2-multipletest": "6d2f2661fa3a2bdaefbdf16f3d26b41ef69b29c93347e16e6df8542f30b14078", + "stats-regression2": "cabb52d66ccb5d92b8d0885f5a5a166c1814d364dc75a4e8cac89819820edba3", + "stats-unsupervised": "c8a17031b7133d58a6b6b6ec21a5624281d8f291717749869e41251b3e8ef68f", + "integrate-stiff-solvers": "48276e65576eeefa8d7c3da8da42ee782441462ce34cf6bb814e494c0e00668f", + "integrate-fem3d": "8323b5c89ac9187441cfac694193ef6394aff0508e71b53052d1b76fd59a535b", + "linalg-sparse-lsqr-rrqr-update": "806455d7e53765a272afca8cfef90aa233ab78e62e4ad682c085ff7e117231fc", + "optim-lp-qp-global2": "510ce65940755a164784ed4577502db2b60604aa9cbe184f22b2ab2a4c03bd57", + "io-hdf5-write": "21482614860f885f277b05666924ab3b63a2c55b8b24734cbadc4577f2c7ac90", + "optim-lm-fit": "cdbf15c3a4c5615e1a364a9fe4e5ec431d0cb65d9744cab52cbca66af51c58cf", + "core-cosm1": "f995296d5c1f285ec6f2e719ee62f8eaf7c9a8e751cea8702cbe16b3df66ddb2", + "core-substream": "ce49fae36f5e15b37aacc25d62b97c9f800d51aa08dee553cee1a5b6e1c01798", + "signal-chirp": "ff9af5783d32eed88bf2658bde5660c247a0bd30a699c403f1755da290830a4b", + "core-besseljreal": "49698a21228f091079cd65ae94cda134b16c3b0bec58a47642f6ad3fec7461d7", + "integrate-filon": "4d5242520a5c144e4bdf111df83c2f2faf4fadce0859f87738acab91e3ed93e0", + "core-narrow": "80913507636c4aabcdc57e5f90a75cd18499c86e244d84006650de3e1273da92", + "stats-contingency": "9ac61a7a29227ddae02861076c85a669fd467f466a8c1118c52965ad4ac20a1b", + "stats-mixedmodel": "6db27ea7908863a1f2d0725325d8a72a93a5ec83ad3882bca28a5d3d2b614227", + "stats-hmm": "779e2c9cb8ee1e1871b17a04a42b2d574ff3de6132430bae8c58be502d450b75", + "stats-hierarchy": "c4171ae9c61b76c676d016356ae319c5e3c15c0e32f7d0c26f0b1b38700cb96e", + }, +} + +// oracleBuild names the code generation this binary was built with: the +// GOAMD64 level decides whether the compiler may contract a multiply and +// an add into one fused operation, and the race detector's +// instrumentation changes which loops the compiler still contracts. +// Each changes the last bits of ordinary arithmetic without changing +// the algorithm, so the digests are keyed by both alongside the +// architecture: a build whose combination has no recorded block skips +// loudly instead of failing, exactly as an unseen architecture does. +func oracleBuild() string { + level := "v1" + if bi, ok := debug.ReadBuildInfo(); ok { + for _, s := range bi.Settings { + if s.Key == "GOAMD64" && s.Value != "" { + level = s.Value + } + } + } + return "+" + level + oracleRaceSuffix +} + +// oracleKey names this architecture and code-generation block. +func oracleKey() string { + return runtime.GOOS + "/" + runtime.GOARCH + oracleBuild() +} + +func TestOracle(t *testing.T) { + record := os.Getenv("TENSOR_ORACLE_RECORD") == "1" + current := make(map[string]string, len(oracleCases)) + for _, c := range oracleCases { + d := newOracleDigest() + for _, v := range c.run(t) { + d.result(t, v) + } + current[c.name] = d.sum() + } + if record { + for _, c := range oracleCases { + t.Logf("oracle %s %s", c.name, current[c.name]) + } + t.Logf("record these into pinnedOracleDigests[%q] and commit", oracleKey()) + return + } + pinned, known := pinnedOracleDigests[oracleKey()] + if !known { + t.Skipf("no oracle digests recorded for %s; run TENSOR_ORACLE_RECORD=1 go test -run TestOracle -v . on it, "+ + "then commit the pinnedOracleDigests[%q] block", oracleKey(), oracleKey()) + } + if len(pinned) != len(oracleCases) { + t.Fatalf("pinnedOracleDigests[%q] holds %d entries, the harness runs %d; record with TENSOR_ORACLE_RECORD=1", + oracleKey(), len(pinned), len(oracleCases)) + } + for _, c := range oracleCases { + want, ok := pinned[c.name] + if !ok { + t.Errorf("oracle case %q is not pinned", c.name) + continue + } + if current[c.name] != want { + t.Errorf("oracle case %q moved:\n got %s\nwant %s\n"+ + "If the change is deliberate, re-record with TENSOR_ORACLE_RECORD=1 in the same commit.", + c.name, current[c.name], want) + } + } +} diff --git a/plot/dtypes_census_test.go b/plot/dtypes_census_test.go new file mode 100644 index 0000000..47dcb50 --- /dev/null +++ b/plot/dtypes_census_test.go @@ -0,0 +1,179 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package plot + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The dtype census for plot: the series surface, probed with Bool, the +// narrow integers and the Int anchor against a float64 baseline +// carrying exactly the widened probe values. Line reads both axes +// through the widening accessors by design, so every numeric dtype +// plots exactly what Int plots. + +var ptDtypes = []core.Dtype{core.Bool, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32, core.Int} + +type ptMaker func(vals []float64, shape ...int) *core.Array + +func ptCast(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 ptMakers(t *testing.T, dt core.Dtype) (probe, base ptMaker) { + t.Helper() + castOf := func(vals []float64) []float64 { + out := make([]float64, len(vals)) + for i, v := range vals { + out[i] = ptCast(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 +} + +// TestDtypesCensusPlot pins Line on every probe dtype against the +// float64 baseline of the same widened values: identical points, no +// panic, and the existing non-finite refusal preserved. +func TestDtypesCensusPlot(t *testing.T) { + xs := []float64{0, 1, 2, 3, 4} + ys := []float64{3, 1, 4, 1, 5} + for _, dt := range ptDtypes { + t.Run("Line/"+dt.String(), func(t *testing.T) { + probe, base := ptMakers(t, dt) + ps, perr := Line("s", probe(xs, 5), probe(ys, 5)) + bs, berr := Line("s", base(xs, 5), base(ys, 5)) + if berr != nil { + t.Fatalf("Line float baseline: %v", berr) + } + if perr != nil { + t.Fatalf("Line(%s): %v; the float baseline of the same values succeeded", dt, perr) + } + if ps.Name != bs.Name || len(ps.Points) != len(bs.Points) { + t.Fatalf("Line(%s): series %q with %d points, want %q with %d", dt, ps.Name, len(ps.Points), bs.Name, len(bs.Points)) + } + for i := range ps.Points { + if ps.Points[i] != bs.Points[i] { + t.Fatalf("Line(%s): point %d = %+v, want %+v", dt, i, ps.Points[i], bs.Points[i]) + } + } + }) + } + // The standing refusals keep their wording for the new dtypes too. + t.Run("Line refusals", func(t *testing.T) { + x8, err := core.FromInt8s([]int8{0, 1, 2, 3, 4}, 5) + if err != nil { + t.Fatal(err) + } + y8, err := core.FromInt8s([]int8{3, 1, 4, 1}, 4) + if err != nil { + t.Fatal(err) + } + if _, err := Line("s", x8, y8); err == nil { + t.Fatal("Line accepted a length mismatch") + } + empty, _ := core.FromInt8s(nil, 0) + if _, err := Line("s", empty, empty); err == nil { + t.Fatal("Line accepted empty arrays") + } + okX, _ := core.FromFloats([]float64{0, 1, 2}, 3) + nanY, _ := core.FromFloats([]float64{1, math.NaN(), 1}, 3) + if _, err := Line("s", okX, nanY); err == nil { + t.Fatal("Line accepted a non-finite point") + } + }) +} diff --git a/plot/plot.go b/plot/plot.go new file mode 100644 index 0000000..4602aee --- /dev/null +++ b/plot/plot.go @@ -0,0 +1,240 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package plot draws the deterministic SVG line charts a scientific +// paper needs: linear axes, five ticks each, one legend line per +// series, and nothing else. The output is deterministic by contract: +// the same chart always renders byte for byte the same file, so a +// figure in a paper can be regenerated and compared exactly like any +// other computed number. The package is small by intent; it draws the +// figures, it does not stage a cinema. +package plot + +import ( + "fmt" + "math" + "os" + "path/filepath" + "strings" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Point is one data point in axis units. +type Point struct { + X, Y float64 +} + +// Series is one named polyline. +type Series struct { + Name string + Points []Point +} + +// Chart is a linear-axis line chart. +type Chart struct { + Title string + XLabel string + YLabel string + Width int + Height int + Series []Series + // XRange and YRange are optional; a zero, inverted or non-finite + // span falls back to the data's own bounds. + XRange [2]float64 + YRange [2]float64 +} + +// Line returns a series joining the points (xs[i], ys[i]) of two +// arrays. Both arrays must be rank 1, of equal, non-zero length, and +// hold only finite numbers; the values are read through the promotion +// ladder, so any numeric dtype is accepted. +func Line(name string, xs, ys *core.Array) (Series, error) { + const op = "Plot" + if xs == nil || ys == nil { + return Series{}, base.Errf("%s: a nil array cannot make a series", op) + } + if xs.NDim() != 1 { + return Series{}, base.Errf("%s: the x values must be rank 1, got shape %s", op, base.ShapeText(xs.Shape())) + } + if ys.NDim() != 1 { + return Series{}, base.Errf("%s: the y values must be rank 1, got shape %s", op, base.ShapeText(ys.Shape())) + } + if xs.Len() != ys.Len() { + return Series{}, base.Errf("%s: length mismatch, %d points of x against %d points of y", + op, xs.Len(), ys.Len()) + } + if xs.Len() == 0 { + return Series{}, base.Errf("%s: an empty array cannot make a series", op) + } + pts := make([]Point, xs.Len()) + for i := range pts { + x, y := xs.FloatAt(i), ys.FloatAt(i) + if math.IsNaN(x) || math.IsInf(x, 0) || math.IsNaN(y) || math.IsInf(y, 0) { + return Series{}, base.Errf("%s: non-finite point at index %d: (%g, %g)", op, i, x, y) + } + pts[i] = Point{X: x, Y: y} + } + return Series{Name: name, Points: pts}, nil +} + +// WriteSVG renders the chart into path. The output is deterministic: +// the same chart always renders byte for byte the same file. Every +// point of every series must be finite, the contract the Line +// constructor enforces on the caller's behalf and this entry point +// enforces for a series built by hand. +func (c Chart) WriteSVG(path string) error { + w := c.Width + if w <= 0 { + w = 720 + } + h := c.Height + if h <= 0 { + h = 460 + } + const ( + left = 64.0 + right = 16.0 + top = 40.0 + bottom = 52.0 + ) + all := make([]Point, 0, 256) + for si, s := range c.Series { + for pi, p := range s.Points { + if math.IsNaN(p.X) || math.IsInf(p.X, 0) || math.IsNaN(p.Y) || math.IsInf(p.Y, 0) { + return base.Errf("Plot: series %d (%s) holds the non-finite point %d: (%g, %g)", si, s.Name, pi, p.X, p.Y) + } + } + all = append(all, s.Points...) + } + if len(all) < 2 { + return base.Errf("Plot: the chart needs at least two points, has %d", len(all)) + } + xr := c.XRange + if !(xr[0] < xr[1]) || math.IsInf(xr[0], 0) || math.IsInf(xr[1], 0) { + xr = bounds(all, true) + } + yr := c.YRange + if !(yr[0] < yr[1]) || math.IsInf(yr[0], 0) || math.IsInf(yr[1], 0) { + yr = bounds(all, false) + } + px := func(x float64) float64 { + return left + (x-xr[0])/(xr[1]-xr[0])*(float64(w)-left-right) + } + py := func(y float64) float64 { + return float64(h) - bottom - (y-yr[0])/(yr[1]-yr[0])*(float64(h)-top-bottom) + } + var b strings.Builder + b.WriteString(xmlHeader) + fmt.Fprintf(&b, "\n", w, h, w, h) + fmt.Fprintf(&b, "\n", w, h) + fmt.Fprintf(&b, "%s\n", + left, esc(c.Title)) + // Axes with five ticks each. + for k := range 5 { + t := xr[0] + (xr[1]-xr[0])*float64(k)/4 + x := px(t) + fmt.Fprintf(&b, "\n", + x, top, x, float64(h)-bottom) + fmt.Fprintf(&b, "%s\n", + x, float64(h)-bottom+16, tick(t)) + } + for k := range 5 { + t := yr[0] + (yr[1]-yr[0])*float64(k)/4 + y := py(t) + fmt.Fprintf(&b, "\n", + left, y, float64(w)-right, y) + fmt.Fprintf(&b, "%s\n", + left-6, y+4, tick(t)) + } + fmt.Fprintf(&b, "\n", + left, float64(h)-bottom, float64(w)-right, float64(h)-bottom) + fmt.Fprintf(&b, "\n", + left, top, left, float64(h)-bottom) + fmt.Fprintf(&b, "%s\n", + (left+float64(w)-right)/2, float64(h)-12, esc(c.XLabel)) + fmt.Fprintf(&b, "%s\n", + top-12, esc(c.YLabel)) + for i, s := range c.Series { + colour := colour(i) + fmt.Fprintf(&b, " 0 { + b.WriteByte(' ') + } + fmt.Fprintf(&b, "%.2f,%.2f", px(p.X), py(p.Y)) + } + b.WriteString("\"/>\n") + ly := top + 16 + float64(i)*16 + fmt.Fprintf(&b, "\n", + float64(w)-230, ly, float64(w)-214, ly, colour) + fmt.Fprintf(&b, "%s\n", + float64(w)-208, ly+4, esc(s.Name)) + } + b.WriteString("\n") + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + return base.Errf("Plot: %w", err) + } + if err := os.WriteFile(path, []byte(b.String()), 0o644); err != nil { + return base.Errf("Plot: %w", err) + } + return nil +} + +func bounds(pts []Point, xAxis bool) [2]float64 { + lo, hi := math.Inf(1), math.Inf(-1) + for _, p := range pts { + v := p.Y + if xAxis { + v = p.X + } + lo = math.Min(lo, v) + hi = math.Max(hi, v) + } + if hi == lo { + hi = lo + 1 + } + pad := 0.05 * (hi - lo) + return [2]float64{lo - pad, hi + pad} +} + +func tick(v float64) string { + if v == math.Trunc(v) && math.Abs(v) < 1e15 { + return fmt.Sprintf("%d", int64(v)) + } + return fmt.Sprintf("%g", v) +} + +func esc(s string) string { + r := strings.NewReplacer("&", "&", "<", "<", ">", ">", `"`, """) + return r.Replace(s) +} + +// seriesColours is the chart's fixed colour cycle: seven even samples +// of the Viridis perceptual-uniform map (Nathaniel J. Smith, Stéfan +// van der Walt and Eric Firing, released under CC0), read from the +// map's 256-entry table at t = 0, 7/60, ..., 0.7 by linear +// interpolation, each channel rounded to the nearest byte. The map's +// light tail is left out on purpose: the +// chart paints on white, and the pale yellows the full range ends in +// drop far below a legible contrast at stroke width, while the +// sampled range runs dark violet through blue and teal to green with +// every stroke legible. The cycle is a constant, so the same chart +// renders the same colours byte for byte, like everything else it +// draws. +var seriesColours = [7]string{ + "#440154", + "#482a79", + "#3d4d8a", + "#2f6c8e", + "#23888e", + "#20a486", + "#43bf71", +} + +func colour(i int) string { + return seriesColours[i%len(seriesColours)] +} + +const xmlHeader = "\n" diff --git a/plot/plot_test.go b/plot/plot_test.go new file mode 100644 index 0000000..94c207a --- /dev/null +++ b/plot/plot_test.go @@ -0,0 +1,184 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package plot + +import ( + "math" + "os" + "path/filepath" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +func sampleChart() Chart { + pts1 := make([]Point, 0, 50) + pts2 := make([]Point, 0, 50) + for i := range 50 { + x := float64(i) / 49 + pts1 = append(pts1, Point{X: x, Y: x * x}) + pts2 = append(pts2, Point{X: x, Y: math.Sqrt(x)}) + } + return Chart{ + Title: "Test ", + XLabel: "x", YLabel: "y", + Series: []Series{{Name: "quadratic", Points: pts1}, {Name: "sqrt", Points: pts2}}, + } +} + +func TestWriteSVG(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "out.svg") + if err := sampleChart().WriteSVG(path); err != nil { + t.Fatal(err) + } + body, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + s := string(body) + if !strings.HasPrefix(s, " (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Time-series model estimation: the autoregressive and +// ARMA staples. The pure AR fit solves the Yule-Walker equations by +// the Durbin-Levinson recursion over the biased autocovariance, the +// small solver this package's own PartialAutocorrelate runs in +// miniature (the recursion lives locally here; the domain boundary +// allows no linalg import, and the earlier recursion answers a +// different question). The mixed model is fitted by the +// Hannan-Rissanen innovations method: a high-order autoregression +// first, whose residuals stand in for the unobserved innovations, +// then one least-squares regression of the series on its own lags and +// the lagged residuals. Both report the concentrated Gaussian +// likelihood of their innovations, and with it the AIC and BIC the +// grid search in SelectARMA minimises. + +// ARMAResult holds one fitted time-series model. The coefficients +// follow the convention x_t − μ = Σ_j φ_j·(x_{t−j} − μ) + e_t + +// Σ_k θ_k·e_{t−k}, innovations white with variance +// InnovationVariance. +type ARMAResult struct { + // AR holds φ_1..φ_P, empty only for a pure MA model. + AR []float64 + // MA holds θ_1..θ_Q, empty only for a pure AR model. + MA []float64 + // InnovationVariance is the estimated variance σ² of e_t. + InnovationVariance float64 + // Mean is the sample mean μ the fit removed and reports back. + Mean float64 + // LogLikelihood is the concentrated Gaussian likelihood of the + // innovations over the whole series: −n/2·(log 2πσ² + 1) at the + // fitted innovation variance, so the criteria of different models + // on one series sit on a common footing. + LogLikelihood float64 + // AIC is −2·LogLikelihood + 2k with k = P + Q + 1, the variance + // counted as a parameter. + AIC float64 + // BIC is −2·LogLikelihood + k·log n, the same k against the + // innovation count's logarithm. + BIC float64 + // P and Q record the orders the result carries. + P int + Q int +} + +// ARMAOptions tunes the Hannan-Rissanen estimation and the grid +// search. HighAROrder is the order of the proxy autoregression whose +// residuals supply the MA stage; the zero value implies the default +// max(16, p+q+8), which the data length bounds. Criterion names the +// information criterion SelectARMA minimises: "aic" or "bic", with +// "aic" implied by the zero value. +type ARMAOptions struct { + HighAROrder int + Criterion string +} + +// armaLikelihood fills the concentrated likelihood and the criteria +// from the innovation variance, evaluated over the whole series: nEff +// is the same sample count for every model fitted to one series, +// which is what makes the criteria comparable across a grid, and k +// counts the coefficients plus the variance. The values a fit +// conditions on (an AR's first p samples, the Hannan-Rissanen +// window's proxy residuals) are estimation detail, not a data +// difference. +func armaLikelihood(res *ARMAResult, nEff int) { + k := res.P + res.Q + 1 + logL := -0.5 * float64(nEff) * (math.Log(2*math.Pi) + math.Log(res.InnovationVariance) + 1) + res.LogLikelihood = logL + res.AIC = 2*float64(k) - 2*logL + res.BIC = math.Log(float64(nEff))*float64(k) - 2*logL +} + +// armaAutocovariance returns the biased autocovariances r_0..r_maxLag +// of a real rank-1 series, the estimator the Yule-Walker theory is +// written for: the mean removed, every lag divided by the sample +// count. The series' values are read through the house autocorrelation +// and rescaled by the lag-zero power. The mean comes back too, because +// every model is fitted to the demeaned series and reported with it. +func armaAutocovariance(name string, x *core.Array, maxLag int) (r []float64, n int, mean float64, err error) { + if x.NDim() != 1 { + return nil, 0, 0, base.Errf("%s: needs a rank-1 series, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, 0, 0, base.Errf("%s: complex series are not supported", name) + } + n = x.Len() + if n < maxLag+2 { + return nil, 0, 0, base.Errf("%s: at least %d samples are needed for lag %d, got %d", + name, maxLag+2, maxLag, n) + } + vals := widenFloats(x) + if err := kfFinite(name, "the series", vals); err != nil { + return nil, 0, 0, err + } + for _, v := range vals { + mean += v + } + mean /= float64(n) + power := 0.0 + for _, v := range vals { + d := v - mean + power += d * d + } + variance := power / float64(n) + acf, err := Autocorrelate(x, maxLag) + if err != nil { + return nil, 0, 0, err + } + r = make([]float64, maxLag+1) + for k := range r { + r[k] = acf.FloatAt(k) * variance + } + return r, n, mean, nil +} + +// yuleWalker solves the Yule-Walker equations R·φ = r for the +// order-p autoregression's coefficients, R the Toeplitz matrix of the +// autocovariances r[0..p], through the Durbin-Levinson recursion: p² +// work, no matrix factored. It returns φ_1..φ_p and the innovation +// variance r₀·(1 − Σφ_j·ρ_j). A non-positive recursion denominator +// means the autocovariances do not belong to a stationary process; +// the same recursion, run for its reflection coefficients, is what +// PartialAutocorrelate reports. +func yuleWalker(name string, r []float64, p int) (phi []float64, sigma2 float64, err error) { + phi = make([]float64, p+1) // 1-based: phi[k] holds φ_k + prev := make([]float64, p+1) + for k := 1; k <= p; k++ { + num, den := r[k], r[0] + for j := 1; j < k; j++ { + num -= prev[j] * r[k-j] + den -= prev[j] * r[j] + } + if !(den > 0) || math.IsNaN(den) { + return nil, 0, base.Errf("%s: the Durbin-Levinson recursion broke down at order %d: the autocovariances do not describe a stationary process", name, k) + } + phi[k] = num / den + for j := 1; j < k; j++ { + phi[j] = prev[j] - phi[k]*prev[k-j] + } + copy(prev, phi) + } + out := make([]float64, p) + copy(out, phi[1:]) + sigma2 = r[0] + for j := 1; j <= p; j++ { + sigma2 -= phi[j] * r[j] + } + if !(sigma2 > 0) || math.IsNaN(sigma2) { + return nil, 0, base.Errf("%s: the innovation variance came out %g: the autocovariances do not describe a stationary process", name, sigma2) + } + return out, sigma2, nil +} + +// armaLstsq solves the small least-squares problem min ‖A·c − b‖² +// through the normal equations and Gaussian elimination with partial +// pivoting. The width stays at a model order's size, where the +// squandered conditioning costs nothing that matters; a pivot too +// small against the normal matrix's largest entry reports a +// degenerate regression rather than returning coefficients for it. +func armaLstsq(name string, design [][]float64, target []float64) ([]float64, error) { + width := len(design[0]) + g := make([]float64, width*width) + rhs := make([]float64, width) + for i, row := range design { + for a := range width { + rhs[a] += row[a] * target[i] + for b := range width { + g[a*width+b] += row[a] * row[b] + } + } + } + big := 0.0 + for _, v := range g { + big = max(big, math.Abs(v)) + } + for a := range width { + piv := a + for c := a + 1; c < width; c++ { + if math.Abs(g[c*width+a]) > math.Abs(g[piv*width+a]) { + piv = c + } + } + if !(math.Abs(g[piv*width+a]) > 1e-12*big) { + return nil, base.Errf("%s: the regression design is degenerate: its normal equations have no pivot at column %d", name, a+1) + } + if piv != a { + for b := range width { + g[a*width+b], g[piv*width+b] = g[piv*width+b], g[a*width+b] + } + rhs[a], rhs[piv] = rhs[piv], rhs[a] + } + for c := a + 1; c < width; c++ { + f := g[c*width+a] / g[a*width+a] + for b := a; b < width; b++ { + g[c*width+b] -= f * g[a*width+b] + } + rhs[c] -= f * rhs[a] + } + } + coef := make([]float64, width) + for a := width - 1; a >= 0; a-- { + total := rhs[a] + for b := a + 1; b < width; b++ { + total -= g[a*width+b] * coef[b] + } + coef[a] = total / g[a*width+a] + } + return coef, nil +} + +// EstimateAR fits the order-p autoregression x_t − μ = Σ_j +// φ_j·(x_{t−j} − μ) + e_t to the real rank-1 series x: the +// Yule-Walker equations over the biased autocovariance, solved by the +// Durbin-Levinson recursion. The result carries the coefficients, the +// innovation variance, the concentrated Gaussian likelihood of the +// whole series, and the AIC and BIC that compare models fitted to +// one series. An order below 1, a series too short for the lag, and +// non-finite or complex data are errors; autocovariances that do not +// describe a stationary process break the recursion and are reported. +func EstimateAR(x *core.Array, order int) (*ARMAResult, error) { + const name = "EstimateAR" + if order < 1 { + return nil, base.Errf("%s: the order must be at least 1, got %d", name, order) + } + r, n, mean, err := armaAutocovariance(name, x, order) + if err != nil { + return nil, err + } + phi, sigma2, err := yuleWalker(name, r, order) + if err != nil { + return nil, err + } + res := &ARMAResult{AR: phi, InnovationVariance: sigma2, Mean: mean, P: order} + armaLikelihood(res, n) + return res, nil +} + +// EstimateARMA fits the ARMA(p, q) model x_t − μ = Σ_j φ_j·(x_{t−j} − +// μ) + e_t + Σ_k θ_k·e_{t−k} by the Hannan-Rissanen innovations +// method: an order-m autoregression first (m from ARMAOptions' +// HighAROrder, defaulting to max(16, p+q+8)), whose residuals stand +// in for the unobserved innovations, then one least-squares +// regression of the demeaned series on its own lags and the lagged +// residuals. The result carries the coefficients in the documented +// convention, the regression residuals' variance as the innovation +// variance and the likelihood and criteria computed from it over the +// whole series, so the criteria line up with the pure AR fit's and +// with the other cells of a SelectARMA grid. +// +// The honest caveat: the method conditions on proxy residuals, so its +// estimates carry more sampling noise than a maximum-likelihood fit +// would, the MA side most of all; tolerances set against it should be +// correspondingly looser. A pure AR model (q = 0) is routed to the +// Yule-Walker fit, which solves that problem exactly; p + q must be +// at least 1, the proxy order at least p+q+1, and the series long +// enough to leave a regression window worth fitting. +func EstimateARMA(x *core.Array, p, q int, opts ARMAOptions) (*ARMAResult, error) { + const name = "EstimateARMA" + if p < 0 || q < 0 { + return nil, base.Errf("%s: the orders must not be negative, got %d and %d", name, p, q) + } + if p+q < 1 { + return nil, base.Errf("%s: at least one coefficient is needed, got orders %d and %d", name, p, q) + } + if q == 0 { + return EstimateAR(x, p) + } + m := opts.HighAROrder + if m == 0 { + m = max(16, p+q+8) + } + if m < p+q+1 { + return nil, base.Errf("%s: the proxy AR order must be at least p+q+1 = %d, got %d", name, p+q+1, m) + } + r, n, mean, err := armaAutocovariance(name, x, m) + if err != nil { + return nil, err + } + rows := n - m - q + if rows < p+q+8 { + return nil, base.Errf("%s: %d samples leave only %d regression rows for %d coefficients; shorten the proxy or lengthen the series", + name, n, rows, p+q) + } + phiAR, _, err := yuleWalker(name, r[:m+1], m) + if err != nil { + return nil, err + } + demeaned := make([]float64, n) + for i := range n { + demeaned[i] = x.FloatAt(i) - mean + } + // The proxy AR's residuals, the innovations' stand-ins: e_t for + // t ≥ m. + resid := make([]float64, n-m) + for t := m; t < n; t++ { + v := demeaned[t] + for j := 1; j <= m; j++ { + v -= phiAR[j-1] * demeaned[t-j] + } + resid[t-m] = v + } + // The regression: x_t − μ on x_{t−1..t−p} − μ and e_{t−1..t−q}, + // over the window t = m+q..n−1 the proxies cover. + T := n - m - q + design := make([][]float64, T) + target := make([]float64, T) + for i := range T { + t := m + q + i + row := make([]float64, 0, p+q) + for j := 1; j <= p; j++ { + row = append(row, demeaned[t-j]) + } + for k := 1; k <= q; k++ { + row = append(row, resid[t-k-m]) + } + design[i] = row + target[i] = demeaned[t] + } + coef, err := armaLstsq(name, design, target) + if err != nil { + return nil, err + } + rss := 0.0 + for i := range T { + fitted := 0.0 + for a, v := range design[i] { + fitted += v * coef[a] + } + d := target[i] - fitted + rss += d * d + } + sigma2 := rss / float64(T) + if !(sigma2 > 0) || math.IsNaN(sigma2) { + return nil, base.Errf("%s: the innovation variance came out %g: the fit carries no information", name, sigma2) + } + res := &ARMAResult{AR: coef[:p], MA: coef[p:], InnovationVariance: sigma2, Mean: mean, P: p, Q: q} + armaLikelihood(res, n) + return res, nil +} + +// SelectARMA searches the order grid 0..maxAR × 0..maxMA for the ARMA +// model the chosen information criterion ranks best: every cell is +// fitted by EstimateARMA, the criterion (ARMAOptions' Criterion, +// "aic" by default) is computed from the innovation variance, and the +// smallest score wins. The (0, 0) cell carries no dynamics and is +// skipped; a cell whose fit fails is skipped too, and a grid where +// nothing can be fitted reports the last failure. The result carries +// both criteria of the winning model, so the runner-up is one more +// sweep away. +// +// The two criteria answer different questions: the AIC's linear +// parameter penalty makes it efficient but willing to overfit by an +// order with probability that does not vanish with the sample count, +// while the BIC's logarithmic penalty makes it consistent, choosing +// the true orders with probability tending to one. Where they +// disagree on long series, the BIC is the one to trust. +func SelectARMA(x *core.Array, maxAR, maxMA int, opts ARMAOptions) (*ARMAResult, error) { + const name = "SelectARMA" + if maxAR < 0 || maxMA < 0 { + return nil, base.Errf("%s: the grid bounds must not be negative, got %d and %d", name, maxAR, maxMA) + } + if maxAR+maxMA < 1 { + return nil, base.Errf("%s: the grid must hold at least one candidate, got %d×%d", name, maxAR+1, maxMA+1) + } + criterion := opts.Criterion + if criterion == "" { + criterion = "aic" + } + if criterion != "aic" && criterion != "bic" { + return nil, base.Errf("%s: the criterion must be %q or %q, got %q", name, "aic", "bic", opts.Criterion) + } + var ( + best *ARMAResult + bestScore float64 + lastErr error + ) + for p := range maxAR + 1 { + for q := range maxMA + 1 { + if p == 0 && q == 0 { + continue + } + res, err := EstimateARMA(x, p, q, opts) + if err != nil { + lastErr = err + continue + } + score := res.AIC + if criterion == "bic" { + score = res.BIC + } + if best == nil || score < bestScore { + best, bestScore = res, score + } + } + } + if best == nil { + return nil, base.Errf("%s: no model on the grid could be fitted: %w", name, lastErr) + } + return best, nil +} + +// ARMASpectrum evaluates the fitted model's theoretical power +// spectrum at the nFreq frequencies evenly spaced from 0 to the +// Nyquist frequency 0.5 inclusive, on the one-sided convention the +// house periodogram prints: the bins strictly between DC and Nyquist +// carry the folded double weight, DC and Nyquist the single one. The +// value is σ²·|Θ(e^{−2πif})|² / |Φ(e^{−2πif})|² with that weighting, +// so model and data line up bin for bin against WelchPSD at fs = 1: +// a white-noise sequence of variance σ² estimates σ² at the edges and +// 2σ² between them, and the fitted model's spectrum describes the +// data's own periodogram without a scale fudge. Fewer than two +// frequencies, a nil model, and a non-positive innovation variance or +// non-finite coefficient are errors; a pole of Φ on the unit circle +// names its frequency instead of answering infinities. +func ARMASpectrum(res *ARMAResult, nFreq int) (freqs, psd *core.Array, err error) { + const name = "ARMASpectrum" + if res == nil { + return nil, nil, base.Errf("%s: the model is nil", name) + } + if nFreq < 2 { + return nil, nil, base.Errf("%s: at least two frequencies are needed, got %d", name, nFreq) + } + if !(res.InnovationVariance > 0) || math.IsNaN(res.InnovationVariance) || math.IsInf(res.InnovationVariance, 0) { + return nil, nil, base.Errf("%s: the innovation variance must be positive and finite, got %g", + name, res.InnovationVariance) + } + vals := append(append([]float64{}, res.AR...), res.MA...) + if err := kfFinite(name, "the coefficients", vals); err != nil { + return nil, nil, err + } + freqs = core.New(core.Float, nFreq) + psd = core.New(core.Float, nFreq) + fRow := freqs.RawFloats() + pRow := psd.RawFloats() + for k := range nFreq { + f := 0.5 * float64(k) / float64(nFreq-1) + omega := 2 * math.Pi * f + // |Φ|² and |Θ|² by real arithmetic: the phase parts multiply + // out to squares of sine sums. Both loops read a sine and a + // cosine of the same argument; math.Sincos returns exactly the + // pair math.Sin and math.Cos produce (verified bit-for-bit), so + // the sums are unchanged while the trig work halves. + rePhi, imPhi := 1.0, 0.0 + for j, phi := range res.AR { + s, c := math.Sincos(float64(j+1) * omega) + rePhi -= phi * c + imPhi += phi * s + } + reTheta, imTheta := 1.0, 0.0 + for k2, theta := range res.MA { + s, c := math.Sincos(float64(k2+1) * omega) + reTheta += theta * c + imTheta -= theta * s + } + den := rePhi*rePhi + imPhi*imPhi + if den == 0 { + return nil, nil, base.Errf("%s: the model has a pole on the unit circle at frequency %g", name, f) + } + fRow[k] = f + pRow[k] = res.InnovationVariance * (reTheta*reTheta + imTheta*imTheta) / den + if k > 0 && k < nFreq-1 { + pRow[k] *= 2 // the one-sided fold between DC and Nyquist + } + } + return freqs, psd, nil +} diff --git a/signal/arma_test.go b/signal/arma_test.go new file mode 100644 index 0000000..3981cf6 --- /dev/null +++ b/signal/arma_test.go @@ -0,0 +1,427 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// arSeries generates a zero-mean AR(p) series: x_t = Σ φ_j·x_{t−j} + +// e_t with standard normal innovations, the first burn values of the +// burn-in discarded. +func arSeries(g *core.Generator, phi []float64, n, burn int) []float64 { + p := len(phi) + x := make([]float64, n+burn) + for i := p; i < len(x); i++ { + v := g.NormalUnit() + for j := range p { + v += phi[j] * x[i-1-j] + } + x[i] = v + } + return x[burn:] +} + +// arma11Series generates a zero-mean ARMA(1,1) series: x_t = +// φ·x_{t−1} + e_t + θ·e_{t−1}, the burn-in discarded. +func arma11Series(g *core.Generator, phi, theta float64, n, burn int) []float64 { + x := make([]float64, n) + var xPrev, ePrev float64 + for i := -burn; i < n; i++ { + e := g.NormalUnit() + v := phi*xPrev + e + theta*ePrev + xPrev, ePrev = v, e + if i >= 0 { + x[i] = v + } + } + return x +} + +// TestEstimateAR2YuleWalker recovers a known AR(2) from a long +// generated series, and then demands the recovered coefficients solve +// the Yule-Walker equations on the empirical autocovariance the fit +// itself read: the Durbin-Levinson solve is exact, so the residuals +// sit at rounding level. The partial autocorrelation, produced by the +// same recursion answering a different question, must agree: lag one +// carries φ1, lag two the signature φ2/(1 − φ1²)... for an AR(2) the +// second partial is φ2 and the rest are noise. +func TestEstimateAR2YuleWalker(t *testing.T) { + const ( + phi1 = 0.8 + phi2 = -0.4 + n = 60000 + ) + g := core.NewGenerator(7) + x := arSeries(g, []float64{phi1, phi2}, n, 500) + res, err := EstimateAR(mustFloats(t, x), 2) + if err != nil { + t.Fatalf("EstimateAR: %v", err) + } + if math.Abs(res.AR[0]-phi1) > 0.02 || math.Abs(res.AR[1]-phi2) > 0.02 { + t.Fatalf("AR(2) coefficients (%.4f, %.4f), want (%.2f, %.2f)", res.AR[0], res.AR[1], phi1, phi2) + } + if math.Abs(res.InnovationVariance-1) > 0.1 { + t.Fatalf("innovation variance %.4f, want 1", res.InnovationVariance) + } + if math.Abs(res.Mean) > 0.05 { + t.Fatalf("mean %.4f, want 0", res.Mean) + } + // The Yule-Walker equations on the empirical autocovariance. + r, _, _, err := armaAutocovariance("test", mustFloats(t, x), 2) + if err != nil { + t.Fatalf("armaAutocovariance: %v", err) + } + for k := range 2 { + got := res.AR[k]*r[0] + res.AR[1-k]*r[1] - r[k+1] + if math.Abs(got) > 1e-8*r[0] { + t.Fatalf("Yule-Walker residual at lag %d = %g, want rounding only", k+1, got) + } + } + // The PACF, the same recursion read for its reflection + // coefficients, must agree with the fit: lag two's partial is the + // order-two coefficient itself, lag one's is the lag-one + // autocorrelation φ1/(1 − φ2), both exact relations of the same + // autocovariances the Yule-Walker solve consumed. + pacf, err := PartialAutocorrelate(mustFloats(t, x), 3) + if err != nil { + t.Fatalf("PartialAutocorrelate: %v", err) + } + if math.Abs(pacf.FloatAt(1)-res.AR[1]) > 1e-9 { + t.Fatalf("PACF lag 2 = %.10f against the fitted φ2 %.10f", pacf.FloatAt(1), res.AR[1]) + } + rho1 := res.AR[0] / (1 - res.AR[1]) + if math.Abs(pacf.FloatAt(0)-rho1) > 1e-9 { + t.Fatalf("PACF lag 1 = %.10f against φ1/(1 − φ2) = %.10f", pacf.FloatAt(0), rho1) + } + if math.Abs(pacf.FloatAt(2)) > 0.02 { + t.Fatalf("PACF lag 3 = %.4f, an AR(2) cuts off after lag 2", pacf.FloatAt(2)) + } + // The criteria carry the concentrated likelihood of the whole + // series with one parameter per coefficient plus the variance. + wantAIC := float64(n)*(math.Log(2*math.Pi)+math.Log(res.InnovationVariance)+1) + 2*3 + if math.Abs(res.AIC-wantAIC) > 1e-6 { + t.Fatalf("AIC %.6f, want %.6f", res.AIC, wantAIC) + } + wantBIC := float64(n)*(math.Log(2*math.Pi)+math.Log(res.InnovationVariance)+1) + math.Log(float64(n))*3 + if math.Abs(res.BIC-wantBIC) > 1e-6 { + t.Fatalf("BIC %.6f, want %.6f", res.BIC, wantBIC) + } +} + +// TestEstimateARErrors pins the input gates and the stationary-data +// requirement. +func TestEstimateARErrors(t *testing.T) { + x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 6) + if _, err := EstimateAR(x, 0); err == nil { + t.Error("order zero accepted") + } + if _, err := EstimateAR(x, 5); err == nil { + t.Error("order at the sample count accepted") + } + rank2, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := EstimateAR(rank2, 1); err == nil { + t.Error("rank-2 series accepted") + } + bad, _ := core.FromFloats([]float64{1, math.NaN(), 3, 4}, 4) + if _, err := EstimateAR(bad, 1); err == nil { + t.Error("non-finite series accepted") + } + constant, _ := core.FromFloats([]float64{2, 2, 2, 2, 2, 2}, 6) + if _, err := EstimateAR(constant, 1); err == nil { + t.Error("constant series accepted") + } +} + +// TestEstimateARMA11HannanRissanen recovers a known ARMA(1,1) from a +// long generated series. The honest caveat, documented with the +// function: Hannan-Rissen conditions on proxy residuals, so the +// estimates carry more sampling noise than a maximum-likelihood fit +// would and the tolerances here are an order looser than the AR(2) +// ones above, the MA side most of all. +func TestEstimateARMA11HannanRissanen(t *testing.T) { + const ( + phi = 0.6 + theta = 0.4 + n = 80000 + ) + g := core.NewGenerator(11) + x := arma11Series(g, phi, theta, n, 500) + res, err := EstimateARMA(mustFloats(t, x), 1, 1, ARMAOptions{}) + if err != nil { + t.Fatalf("EstimateARMA: %v", err) + } + if res.P != 1 || res.Q != 1 { + t.Fatalf("orders (%d, %d), want (1, 1)", res.P, res.Q) + } + if math.Abs(res.AR[0]-phi) > 0.05 { + t.Fatalf("AR coefficient %.4f, want %.2f within 0.05", res.AR[0], phi) + } + if math.Abs(res.MA[0]-theta) > 0.05 { + t.Fatalf("MA coefficient %.4f, want %.2f within 0.05", res.MA[0], theta) + } + if math.Abs(res.InnovationVariance-1) > 0.15 { + t.Fatalf("innovation variance %.4f, want 1 within 0.15", res.InnovationVariance) + } + if math.Abs(res.Mean) > 0.05 { + t.Fatalf("mean %.4f, want 0", res.Mean) + } + // The criteria carry the concentrated whole-series likelihood + // here too, not one windowed by the high-order AR proxy: a + // regression to the windowed count shifts AIC by dozens of points + // and breaks cross-model comparability, so the closed form pins + // the value on the ARMA path exactly as it does on the AR one. + wantLL := -float64(n) / 2 * (math.Log(2*math.Pi) + math.Log(res.InnovationVariance) + 1) + if math.Abs(res.LogLikelihood-wantLL) > 1e-6 { + t.Fatalf("log likelihood %.6f, want %.6f", res.LogLikelihood, wantLL) + } + wantAIC := -2*res.LogLikelihood + 2*3 + if math.Abs(res.AIC-wantAIC) > 1e-6 { + t.Fatalf("AIC %.6f, want %.6f", res.AIC, wantAIC) + } +} + +// TestEstimateARMAErrors pins the order and data gates, including the +// routing of a pure AR model to the Yule-Walker fit. +func TestEstimateARMAErrors(t *testing.T) { + x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}, 10) + if _, err := EstimateARMA(x, 0, 0, ARMAOptions{}); err == nil { + t.Error("orders (0, 0) accepted") + } + if _, err := EstimateARMA(x, -1, 1, ARMAOptions{}); err == nil { + t.Error("negative order accepted") + } + if _, err := EstimateARMA(x, 1, 1, ARMAOptions{HighAROrder: 1}); err == nil { + t.Error("proxy order below p+q+1 accepted") + } + if _, err := EstimateARMA(x, 1, 1, ARMAOptions{}); err == nil { + t.Error("series too short for the regression window accepted") + } + // A pure AR request is routed to the Yule-Walker fit: orders below + // its own gate are refused there. + if _, err := EstimateARMA(x, 0, 1, ARMAOptions{}); err == nil { + t.Error("degenerate ARMA with p = 0 accepted on a short series") + } + constant, _ := core.FromFloats([]float64{2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2}, 20) + if _, err := EstimateARMA(constant, 1, 1, ARMAOptions{}); err == nil { + t.Error("constant series accepted") + } +} + +// TestSelectARMAInformationCriteria generates an AR(2) and demands the +// grid search pick exactly (2, 0) by both criteria, then an ARMA(1,1) +// and demands both criteria pick exactly (1, 1). Both pins are +// deterministic on the seeded generator: the streams are bit-stable, +// so the choice is a regression pin, not a restatement of the +// criteria's large-sample guarantees. The seeds are chosen for wide +// decision margins (the closest competitor sits 1.999 AIC points +// behind on the AR(2) grid, 1.991 AIC and 11.3 BIC points on the +// ARMA(1,1) grid), because the AIC's documented willingness to +// overfit by one order makes exact selection a coin weighted by the +// sample: with margin 2 the decision cannot move on arithmetic-level +// perturbations. +func TestSelectARMAInformationCriteria(t *testing.T) { + g := core.NewGenerator(29) + x := arSeries(g, []float64{0.8, -0.4}, 60000, 500) + for _, criterion := range []string{"aic", "bic"} { + res, err := SelectARMA(mustFloats(t, x), 3, 1, ARMAOptions{Criterion: criterion}) + if err != nil { + t.Fatalf("SelectARMA(%s): %v", criterion, err) + } + if res.P != 2 || res.Q != 0 { + t.Fatalf("%s picked ARMA(%d, %d), want ARMA(2, 0)", criterion, res.P, res.Q) + } + } + // Both criteria of the winning model are reported. + res, err := SelectARMA(mustFloats(t, x), 3, 1, ARMAOptions{Criterion: "bic"}) + if err != nil { + t.Fatalf("SelectARMA: %v", err) + } + if !(res.BIC > res.AIC) { + t.Fatalf("BIC %.4f does not exceed AIC %.4f at n = 60000", res.BIC, res.AIC) + } + // The ARMA(1,1) series on a wider grid: both criteria pin (1, 1). + g2 := core.NewGenerator(4) + y := arma11Series(g2, 0.6, 0.4, 80000, 500) + for _, criterion := range []string{"aic", "bic"} { + best, err := SelectARMA(mustFloats(t, y), 2, 2, ARMAOptions{Criterion: criterion}) + if err != nil { + t.Fatalf("SelectARMA(%s): %v", criterion, err) + } + if best.P != 1 || best.Q != 1 { + t.Fatalf("%s picked ARMA(%d, %d), want ARMA(1, 1)", criterion, best.P, best.Q) + } + } +} + +// TestSelectARMAErrors pins the grid and criterion gates, including +// the reporting when no cell can be fitted. +func TestSelectARMAErrors(t *testing.T) { + x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 6) + if _, err := SelectARMA(x, 0, 0, ARMAOptions{}); err == nil { + t.Error("an empty grid accepted") + } + if _, err := SelectARMA(x, -1, 2, ARMAOptions{}); err == nil { + t.Error("a negative grid bound accepted") + } + if _, err := SelectARMA(x, 1, 1, ARMAOptions{Criterion: "sic"}); err == nil { + t.Error("an unknown criterion accepted") + } + constant, _ := core.FromFloats([]float64{2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2}, 18) + if _, err := SelectARMA(constant, 2, 2, ARMAOptions{}); err == nil { + t.Error("a grid where nothing can be fitted reported no error") + } +} + +// TestARMASpectrumMatchesPeriodogram evaluates the fitted AR(2) model's +// theoretical spectrum against Welch's estimate of a long series +// generated from that fitted model: the frequency grids line up bin +// for bin, and on the broad band between 0.03 and 0.47 +// cycles-per-sample the two log-spectra must track each other well +// inside the averaging fluctuation of the estimate. +func TestARMASpectrumMatchesPeriodogram(t *testing.T) { + const n = 40000 + g := core.NewGenerator(31) + x := arSeries(g, []float64{0.8, -0.4}, n, 500) + fitted, err := EstimateAR(mustFloats(t, x), 2) + if err != nil { + t.Fatalf("EstimateAR: %v", err) + } + // A long series from the fitted coefficients: the fitted model's + // own output, whose periodogram the theory must describe. + const n2 = 1 << 18 + sigma := math.Sqrt(fitted.InnovationVariance) + y := make([]float64, n2+500) + for i := 2; i < len(y); i++ { + y[i] = fitted.AR[0]*y[i-1] + fitted.AR[1]*y[i-2] + sigma*g.NormalUnit() + } + const segment = 1024 + freqs, psd, err := WelchPSD(mustFloats(t, y[500:]), 1, segment, segment/2, "hann") + if err != nil { + t.Fatalf("WelchPSD: %v", err) + } + bins := segment/2 + 1 + tFreqs, tPsd, err := ARMASpectrum(fitted, bins) + if err != nil { + t.Fatalf("ARMASpectrum: %v", err) + } + for k := range bins { + if math.Abs(freqs.FloatAt(k)-tFreqs.FloatAt(k)) > 1e-12 { + t.Fatalf("frequency grids disagree at bin %d", k) + } + } + worst := 0.0 + for k := range bins { + f := freqs.FloatAt(k) + if f < 0.03 || f > 0.47 { + continue + } + d := math.Abs(math.Log(psd.FloatAt(k)) - math.Log(tPsd.FloatAt(k))) + if d > worst { + worst = d + } + } + if worst > 0.2 { + t.Fatalf("log-spectral gap %.4f on the broad band, want under 0.2", worst) + } + // The resonance: cos ω₀ = φ1(φ2 − 1)/(4φ2) locates the AR(2) + // peak; both the Welch estimate and the theoretical spectrum must + // peak inside ±0.03 cycles-per-sample of it, and the theoretical + // argmax must sit within two bins of the formula's own optimum + // (the peak is broad, so bins are the honest resolution). + phi1, phi2 := fitted.AR[0], fitted.AR[1] + cosw0 := phi1 * (phi2 - 1) / (4 * phi2) + f0 := math.Acos(cosw0) / (2 * math.Pi) + if f0 < 0.02 || f0 > 0.48 { + t.Fatalf("the resonance formula left the band: %g", f0) + } + binOf := func(a *core.Array) int { + best, at := math.Inf(-1), -1 + for k := range bins { + f := freqs.FloatAt(k) + if f < 0.02 || f > 0.48 { + continue + } + if a.FloatAt(k) > best { + best, at = a.FloatAt(k), k + } + } + return at + } + if d := math.Abs(float64(binOf(psd))/float64(segment) - f0); d > 0.03 { + t.Fatalf("the empirical peak sits %.4f from the resonance %g", d, f0) + } + if d := math.Abs(float64(binOf(tPsd))/float64(segment) - f0); d > 2.0/float64(segment) { + t.Fatalf("the theoretical peak sits %.4f from the resonance %g", d, f0) + } + // The theoretical values, checked against the formula evaluated + // independently on a fine grid: the ARMA(1,0) polynomial arithmetic + // is the whole function, so this pins it to rounding. + const fine = 4097 + fFreqs, fPsd, err := ARMASpectrum(fitted, fine) + if err != nil { + t.Fatalf("ARMASpectrum: %v", err) + } + for k := range fine { + f := 0.5 * float64(k) / float64(fine-1) + w := 2 * math.Pi * f + re := 1 - phi1*math.Cos(w) - phi2*math.Cos(2*w) + im := phi1*math.Sin(w) + phi2*math.Sin(2*w) + want := fitted.InnovationVariance / (re*re + im*im) + if k > 0 && k < fine-1 { + want *= 2 + } + if math.Abs(fPsd.FloatAt(k)-want) > 1e-10*want { + t.Fatalf("spectrum %.12g at bin %d, formula %.12g", fPsd.FloatAt(k), k, want) + } + if math.Abs(fFreqs.FloatAt(k)-f) > 1e-12 { + t.Fatalf("frequency %.12g at bin %d, want %.12g", fFreqs.FloatAt(k), k, f) + } + } + // The flat model answers the innovation variance at the edges and + // its double between them: the one-sided fold the Welch estimate + // of the same white noise prints. + flat := &ARMAResult{InnovationVariance: 2.5} + _, flatPsd, err := ARMASpectrum(flat, 16) + if err != nil { + t.Fatalf("ARMASpectrum: %v", err) + } + for k := range 16 { + want := 5.0 + if k == 0 || k == 15 { + want = 2.5 + } + if flatPsd.FloatAt(k) != want { + t.Fatalf("the flat model's spectrum is %g at bin %d, want %g", flatPsd.FloatAt(k), k, want) + } + } +} + +// TestARMASpectrumErrors pins the gates, including the pole on the +// unit circle being named instead of answered with infinities. +func TestARMASpectrumErrors(t *testing.T) { + if _, _, err := ARMASpectrum(nil, 8); err == nil { + t.Error("a nil model accepted") + } + res := &ARMAResult{AR: []float64{0.5}, InnovationVariance: 1} + if _, _, err := ARMASpectrum(res, 1); err == nil { + t.Error("a single frequency accepted") + } + res.InnovationVariance = -1 + if _, _, err := ARMASpectrum(res, 8); err == nil { + t.Error("a negative innovation variance accepted") + } + res.InnovationVariance = 1 + res.AR = []float64{math.NaN()} + if _, _, err := ARMASpectrum(res, 8); err == nil { + t.Error("a non-finite coefficient accepted") + } + res.AR = []float64{1} + if _, _, err := ARMASpectrum(res, 8); err == nil { + t.Error("a pole on the unit circle accepted") + } +} diff --git a/signal/bench_perf_test.go b/signal/bench_perf_test.go new file mode 100644 index 0000000..bb7d14d --- /dev/null +++ b/signal/bench_perf_test.go @@ -0,0 +1,355 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the kernels whose scheduling and per-element work was +// reworked: the spectral Poisson solves, the Bluestein route for a +// non-power-of-two length, the DCT/DST table path, the CWT wavelet +// build, the nonuniform transform and the rank filters. Every input is a +// fixed literal pattern, so the numbers are comparable between runs and +// machines. + +// benchPerfFloats builds a deterministic float array of the requested +// rank from a fixed literal pattern. +func benchPerfFloats(b *testing.B, n int, shape ...int) *core.Array { + b.Helper() + v := make([]float64, n) + for i := range v { + v[i] = math.Sin(float64(i)*0.017) + 0.25*float64(i%29) - 3 + } + a, err := core.FromFloats(v, shape...) + if err != nil { + b.Fatal(err) + } + return a +} + +// poissonSource builds the right-hand side of a Poisson solve on a +// rows×cols grid: a smooth interior bump that vanishes on the boundary, +// the shape the Dirichlet solve is built for. +func poissonSource(b *testing.B, rows, cols int) *core.Array { + b.Helper() + v := make([]float64, rows*cols) + for r := range rows { + for c := range cols { + v[r*cols+c] = math.Sin(math.Pi*float64(r)/float64(rows-1)) * + math.Sin(2*math.Pi*float64(c)/float64(cols-1)) + } + } + a, err := core.FromFloats(v, rows, cols) + if err != nil { + b.Fatal(err) + } + return a +} + +// BenchmarkSolvePoissonDirichlet measures the 512×512 sine-basis solve: +// the per-mode division by the stencil eigenvalues and the two +// separable DST-I passes. The grid is large enough for the eigenvalue +// loop to be a visible share of the total. +func BenchmarkSolvePoissonDirichlet(b *testing.B) { + f := poissonSource(b, 512, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := SolvePoissonDirichlet(f, 1, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkSolvePoissonNeumann is the cosine-basis twin on a 256×256 +// grid. +func BenchmarkSolvePoissonNeumann(b *testing.B) { + f := poissonSource(b, 256, 256) + b.ReportAllocs() + for b.Loop() { + if _, err := SolvePoissonNeumann(f, 1, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkFFTBluestein1000 measures a 1-D transform at a length that +// is not a power of two, the only route to Bluestein's chirp-z +// transform. The chirp and the kernel spectrum depend only on the +// length and the direction, which is what the plan cache holds. +func BenchmarkFFTBluestein1000(b *testing.B) { + x := benchPerfFloats(b, 1000, 1000) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT(x); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkFFTBluestein10000 is the same route at a length whose padded +// size is large enough for the kernel transform to dominate. +func BenchmarkFFTBluestein10000(b *testing.B) { + x := benchPerfFloats(b, 10000, 10000) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT(x); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkDCT2_1024 measures the type-2 cosine transform of a +// 1024-sample vector: the per-sample phase table, one padded inverse +// FFT and the per-bin rotation table. +func BenchmarkDCT2_1024(b *testing.B) { + x := benchPerfFloats(b, 1024, 1024) + b.ReportAllocs() + for b.Loop() { + if _, err := DCT(x, 2); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkDST1_1023 measures the type-1 sine transform, whose phase +// table is a constant fill and whose bin rotations are the widest. +func BenchmarkDST1_1023(b *testing.B) { + x := benchPerfFloats(b, 1023, 1023) + b.ReportAllocs() + for b.Loop() { + if _, err := DST(x, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkCWTMorletScales16 measures sixteen Morlet scales over a +// 512-sample signal: the per-scale wavelet build (a cosine, a sine and +// an exponential per sample, with the time axis wrapped once per +// sample) and the FFT pair around it. +func BenchmarkCWTMorletScales16(b *testing.B) { + x := benchPerfFloats(b, 512, 512) + scales := make([]float64, 16) + for i := range scales { + scales[i] = 1 + float64(i) + } + b.ReportAllocs() + for b.Loop() { + if _, err := CWT(x, Morlet, scales, 0.25); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkNUFFTType1_4096 measures the type-1 nonuniform transform: +// 4096 samples spread onto an eight-times oversampled grid of 2048 +// points, one inverse FFT and the per-bin deconvolution. +func BenchmarkNUFFTType1_4096(b *testing.B) { + count := 4096 + xs := make([]float64, count) + cs := make([]complex128, count) + for j := range count { + xs[j] = float64(j)/float64(count) - 0.5 + 1e-4 + cs[j] = complex(math.Sin(float64(j)*0.05), math.Cos(float64(j)*0.11)) + } + x, err := core.FromFloats(xs, count) + if err != nil { + b.Fatal(err) + } + c, err := core.FromComplexes(cs, count) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := NUFFTType1(x, c, 512); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMedianFilterWide31 measures the running median over a +// 32768-sample signal with a 31-sample window: every output is the +// middle order statistic of its own window. +func BenchmarkMedianFilterWide31(b *testing.B) { + x := benchPerfFloats(b, 1<<15, 1<<15) + b.ReportAllocs() + for b.Loop() { + if _, err := MedianFilter(x, 31); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkRankFilterWide31_7 is the general-rank twin: the same window +// read at rank 7 rather than at the middle. +func BenchmarkRankFilterWide31_7(b *testing.B) { + x := benchPerfFloats(b, 1<<15, 1<<15) + b.ReportAllocs() + for b.Loop() { + if _, err := RankFilter(x, 31, 7); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMedianFilter2D measures the 3×3 running median over a +// 256×256 image, whose windows truncate at every edge. +func BenchmarkMedianFilter2D(b *testing.B) { + img := benchPerfFloats(b, 256*256, 256, 256) + b.ReportAllocs() + for b.Loop() { + if _, err := MedianFilter2D(img, 3); err != nil { + b.Fatal(err) + } + } +} + +// rankFilterPerWindow is the per-window sort RankFilter used before the +// sliding mirror, kept here as the reference the current kernel is +// checked against bit for bit. A window is collected and sorted whole +// for every output point, and the rank reads the sorted slice. +func rankFilterPerWindow(x *core.Array, window, k int) *core.Array { + n := x.Len() + src := widenFloats(x) + out := core.New(core.Float, n) + dst := out.RawFloats() + win := make([]float64, window) + half := window / 2 + for i := range n { + lo := max(i-half, 0) + hi := min(i+half+1, n) + count := hi - lo + nan := false + for j := lo; j < hi; j++ { + v := src[j] + if math.IsNaN(v) { + nan = true + break + } + win[j-lo] = v + } + if nan { + dst[i] = math.NaN() + continue + } + slices.Sort(win[:count]) + dst[i] = win[k*count/window] + } + return out +} + +// TestRankFilterMirrorMatchesPerWindow pins the sliding mirror to the +// per-window sort bit for bit. The windows where the sort's own tie +// order decides the answer, a zero next to a −0 and any NaN, are the +// reason the mirror is dropped there; the patterns below hold ties, +// signed zeros, duplicates, NaNs, constant runs and ramps, over every +// signal length and window the gate admits. +func TestRankFilterMirrorMatchesPerWindow(t *testing.T) { + seed := uint64(0x2545F4914F6CDD1D) + next := func() uint64 { + seed ^= seed << 13 + seed ^= seed >> 7 + seed ^= seed << 17 + return seed + } + fill := func(pattern string, n int) []float64 { + vals := make([]float64, n) + for i := range vals { + u := next() + switch pattern { + case "random": + vals[i] = float64(int64(u%2001)-1000) / 8 + case "duplicates": + vals[i] = float64(u % 4) + case "signed zeros": + switch u % 5 { + case 0: + vals[i] = math.Copysign(0, -1) + case 1: + vals[i] = 0 + default: + vals[i] = float64(u%3) - 1 + } + case "NaN": + vals[i] = float64(u%7) - 3 + if u%11 == 0 { + vals[i] = math.NaN() + } + case "constant": + vals[i] = 2.5 + default: + vals[i] = float64(i) + } + } + return vals + } + values := func(t *testing.T, vals []float64) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, len(vals)) + if err != nil { + t.Fatal(err) + } + return a + } + for _, pattern := range []string{"random", "duplicates", "signed zeros", "NaN", "constant", "ramp"} { + for n := 1; n <= 48; n++ { + x := values(t, fill(pattern, n)) + for window := 3; window <= n; window += 2 { + for k := range window { + got, err := RankFilter(x, window, k) + if err != nil { + t.Fatalf("%s n=%d window=%d k=%d: %v", pattern, n, window, k, err) + } + want := rankFilterPerWindow(x, window, k) + g, w := got.RawFloats(), want.RawFloats() + for i := range g { + if math.Float64bits(g[i]) != math.Float64bits(w[i]) { + t.Fatalf("%s n=%d window=%d k=%d sample %d: got %v want %v", + pattern, n, window, k, i, g[i], w[i]) + } + } + } + } + } + } + // Long runs, so the mirror has slid far before the check. + for _, pattern := range []string{"random", "duplicates", "signed zeros", "NaN"} { + x := values(t, fill(pattern, 5000)) + for _, window := range []int{3, 5, 31, 997} { + for _, k := range []int{0, 1, window / 2, window - 2, window - 1} { + got, err := RankFilter(x, window, k) + if err != nil { + t.Fatal(err) + } + want := rankFilterPerWindow(x, window, k) + g, w := got.RawFloats(), want.RawFloats() + for i := range g { + if math.Float64bits(g[i]) != math.Float64bits(w[i]) { + t.Fatalf("%s window=%d k=%d sample %d: got %v want %v", + pattern, window, k, i, g[i], w[i]) + } + } + } + } + } +} + +// BenchmarkFFTRealInput4096 measures the 1-D forward transform of a +// real (float64) array: the payload is materialised rather than shared, +// so the entry point adds no defensive copy of its own. +func BenchmarkFFTRealInput4096(b *testing.B) { + x := benchPerfFloats(b, 4096, 4096) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT(x); err != nil { + b.Fatal(err) + } + } +} diff --git a/signal/bench_test.go b/signal/bench_test.go new file mode 100644 index 0000000..e2d95ff --- /dev/null +++ b/signal/bench_test.go @@ -0,0 +1,138 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +func BenchmarkConv2DParallel(b *testing.B) { benchConv2DN(b) } + +func BenchmarkConv2DSerial(b *testing.B) { + core.SetNumCPU(1) + defer core.SetNumCPU(0) + benchConv2DN(b) +} + +func benchConv2DN(b *testing.B) { + in, _ := core.FromFloats(make([]float64, 8*64*28*28), 8, 64, 28, 28) + k, _ := core.FromFloats(make([]float64, 64*64*3*3), 64, 64, 3, 3) + bi, _ := core.FromFloats(make([]float64, 64), 64) + b.ResetTimer() + for b.Loop() { + if _, err := Conv2D(in, k, bi, 1, 1); err != nil { + b.Fatal(err) + } + } +} + +// FFT benchmarks: the radix-2 kernel at a power-of-two length, and the +// 2-D separable path whose per-line scratch and payload accessors +// dominate at practical sizes. Run with -bench before and after any +// change to fft.go. + +func BenchmarkFFT4096(b *testing.B) { + vals := make([]complex128, 4096) + for i := range vals { + vals[i] = complex(float64(i%17)-8, float64(i%5)-2) + } + a := complexFromArrayMust(vals, []int{len(vals)}) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkFFT2_128x128(b *testing.B) { + vals := make([]complex128, 128*128) + for i := range vals { + vals[i] = complex(float64(i%17)-8, float64(i%5)-2) + } + a := complexFromArrayMust(vals, []int{128, 128}) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT2(a); err != nil { + b.Fatal(err) + } + } +} + +// Larger sizes: the 1-D entry where per-stage parallelism and the +// cached bit-reversal table pay off, and the 2-D entry where the +// unit-stride row pass dominates. + +func BenchmarkFFT65536(b *testing.B) { + vals := make([]complex128, 65536) + for i := range vals { + vals[i] = complex(float64(i%17)-8, float64(i%5)-2) + } + a := complexFromArrayMust(vals, []int{len(vals)}) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkFFT2_512x512(b *testing.B) { + vals := make([]complex128, 512*512) + for i := range vals { + vals[i] = complex(float64(i%17)-8, float64(i%5)-2) + } + a := complexFromArrayMust(vals, []int{512, 512}) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT2(a); err != nil { + b.Fatal(err) + } + } +} + +// Serial twins: one worker, so the parallel contribution of each size +// reads directly off the pair. + +func BenchmarkFFT65536Serial(b *testing.B) { + core.SetNumCPU(1) + defer core.SetNumCPU(0) + benchFFTLine(b, 65536) +} + +func BenchmarkFFT2_512x512Serial(b *testing.B) { + core.SetNumCPU(1) + defer core.SetNumCPU(0) + benchFFT2Square(b, 512) +} + +func benchFFTLine(b *testing.B, n int) { + vals := make([]complex128, n) + for i := range vals { + vals[i] = complex(float64(i%17)-8, float64(i%5)-2) + } + a := complexFromArrayMust(vals, []int{n}) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT(a); err != nil { + b.Fatal(err) + } + } +} + +func benchFFT2Square(b *testing.B, side int) { + vals := make([]complex128, side*side) + for i := range vals { + vals[i] = complex(float64(i%17)-8, float64(i%5)-2) + } + a := complexFromArrayMust(vals, []int{side, side}) + b.ReportAllocs() + for b.Loop() { + if _, err := FFT2(a); err != nil { + b.Fatal(err) + } + } +} diff --git a/signal/chirp.go b/signal/chirp.go new file mode 100644 index 0000000..793e78c --- /dev/null +++ b/signal/chirp.go @@ -0,0 +1,60 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// Chirp returns n samples of a linear frequency sweep from f0 through +// f1, sampled at rate samples per unit of time. The sweep runs over +// the whole sample span and reaches f1 exactly at the last sample; the +// first sample sits at phase zero. Every phase comes from the closed +// form 2π·(f0·t + (f1−f0)·t²/(2T)) at that sample's own time, never +// from a recursive oscillator, so the samples carry no accumulated +// drift: sample k is as accurate as the formula at t_k and nothing +// that happened at earlier samples can bend it. That property is what +// aliasing studies of a swept tone need, which is where the transform +// side of this package keeps meeting chirps in the wild. +// +// The sweep edges live strictly below the Nyquist frequency rate/2: an +// edge on or beyond it folds onto lower frequencies, and the error +// names both. The edges may be negative, which reverses the direction +// of rotation without weakening the guard, which watches their +// magnitudes. +// +// Errors: n below 1, a non-positive or non-finite rate, a non-finite +// sweep edge, an edge at or beyond Nyquist. +func Chirp(n int, f0, f1, rate float64) ([]float64, error) { + if n < 1 { + return nil, base.Errf("Chirp: n must be at least 1, got %d", n) + } + if math.IsNaN(rate) || math.IsInf(rate, 0) || rate <= 0 { + return nil, base.Errf("Chirp: the sample rate must be finite and positive, got %g", rate) + } + if math.IsNaN(f0) || math.IsInf(f0, 0) || math.IsNaN(f1) || math.IsInf(f1, 0) { + return nil, base.Errf("Chirp: the sweep edges must be finite, got %g and %g", f0, f1) + } + nyquist := rate / 2 + edge := math.Max(math.Abs(f0), math.Abs(f1)) + if edge >= nyquist { + return nil, base.Errf("Chirp: the sweep edge %g reaches the Nyquist frequency %g at rate %g", edge, nyquist, rate) + } + out := make([]float64, n) + // The span up to the last sample; a one-sample chirp has no span + // and no slope, and its single sample is the phase-zero zero. + span := float64(n-1) / rate + slope := 0.0 + if n > 1 { + slope = (f1 - f0) / span + } + const twoPi = 2 * math.Pi + for i := range out { + t := float64(i) / rate + out[i] = math.Sin(twoPi * (f0*t + 0.5*slope*t*t)) + } + return out, nil +} diff --git a/signal/chirp_test.go b/signal/chirp_test.go new file mode 100644 index 0000000..6d0d4ed --- /dev/null +++ b/signal/chirp_test.go @@ -0,0 +1,158 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "math/big" + "testing" +) + +// TestChirpConstantToneIsAPureSine pins the degenerate sweep: with f0 +// equal to f1 the slope is exactly zero and every sample must be bit +// for bit the plain sine at that frequency. +func TestChirpConstantToneIsAPureSine(t *testing.T) { + const f, rate = 17.3, 800.0 + got, err := Chirp(400, f, f, rate) + if err != nil { + t.Fatalf("Chirp: %v", err) + } + for i := range got { + tt := float64(i) / rate + if want := math.Sin(2 * math.Pi * (f*tt + 0.5*0*tt*tt)); got[i] != want { + t.Fatalf("sample %d = %.17g, want the plain sine %.17g", i, got[i], want) + } + } + if got[0] != 0 { + t.Fatalf("the first sample = %.17g, want the phase-zero zero", got[0]) + } +} + +// chirpPhaseBig evaluates the sweep phase 2π(f0·t + (f1−f0)t²/(2T)) in +// 256-bit arithmetic from the exact float64 edges, the referent the +// samples are held against. +func chirpPhaseBig(i int, f0, f1, span float64, rate int64) *big.Float { + const prec = 256 + t := new(big.Float).SetPrec(prec).SetInt64(int64(i)) + t.Quo(t, new(big.Float).SetPrec(prec).SetInt64(rate)) + f0b := new(big.Float).SetPrec(prec).SetFloat64(f0) + slope := new(big.Float).SetPrec(prec).SetFloat64(f1 - f0) + slope.Quo(slope, new(big.Float).SetPrec(prec).SetFloat64(span)) + phase := new(big.Float).SetPrec(prec).Mul(slope, t) + phase.Mul(phase, t) + phase.Quo(phase, new(big.Float).SetPrec(prec).SetInt64(2)) + phase.Add(phase, new(big.Float).SetPrec(prec).Mul(f0b, t)) + return phase.Mul(phase, new(big.Float).SetPrec(prec).SetFloat64(2*math.Pi)) +} + +// TestChirpPhaseAgainstBigFloat holds selected samples against the +// phase evaluated in 256-bit arithmetic: the float64 route may lose +// only the rounding its own phase arithmetic costs, a few parts in +// 1e16 of a phase that stays near 2π·40. +func TestChirpPhaseAgainstBigFloat(t *testing.T) { + const n, rate = 200, 1000 + const f0, f1 = 3, 40 + span := float64(n-1) / rate + got, err := Chirp(n, f0, f1, float64(rate)) + if err != nil { + t.Fatalf("Chirp: %v", err) + } + for _, i := range []int{0, 1, 37, 100, 150, 198, 199} { + phase, _ := chirpPhaseBig(i, f0, f1, span, rate).Float64() + want := math.Sin(phase) + if math.Abs(got[i]-want) > 5e-13 { + t.Fatalf("sample %d = %.17g, want sin(256-bit phase) %.17g", i, got[i], want) + } + } +} + +// TestChirpSweepDirection counts zero crossings per eighth of a rising +// sweep: the count must climb from roughly f0's rate to roughly f1's, +// which pins both the direction of the sweep and its approximate +// linearity. +func TestChirpSweepDirection(t *testing.T) { + const n, rate = 8000, 8000.0 + const f0, f1 = 100, 900 + got, err := Chirp(n, f0, f1, rate) + if err != nil { + t.Fatalf("Chirp: %v", err) + } + crossings := func(from, to int) int { + c := 0 + for i := from; i < to; i++ { + if got[i] >= 0 && got[i+1] < 0 { + c++ + } + } + return c + } + first := crossings(0, n/8) + last := crossings(7*n/8, n-1) + // The eighth spans 1000 samples; the instantaneous frequency runs + // from about 200 to about 887 Hz, so the crossing counts must + // climb and sit near those rates. + if first < 15 || first > 40 || last < 85 || last > 110 { + t.Fatalf("crossings per eighth: first %d, last %d, want a climb from near 200 Hz to near 890 Hz", first, last) + } + if last <= first { + t.Fatalf("the sweep does not rise: %d crossings at the end against %d at the start", last, first) + } +} + +// TestChirpErrors pins the guard rails: the length, the rate, the +// finite edges and the Nyquist boundary the sweep must stay under, +// counting the magnitude for negative edges too. +func TestChirpErrors(t *testing.T) { + if _, err := Chirp(0, 1, 2, 100); err == nil { + t.Fatal("n = 0: want an error") + } + if _, err := Chirp(1, 1, 2, 0); err == nil { + t.Fatal("a zero rate: want an error") + } + if _, err := Chirp(1, 1, 2, math.Inf(1)); err == nil { + t.Fatal("an infinite rate: want an error") + } + if _, err := Chirp(1, math.NaN(), 2, 100); err == nil { + t.Fatal("a NaN edge: want an error") + } + if _, err := Chirp(10, 1, math.Inf(-1), 100); err == nil { + t.Fatal("an infinite edge: want an error") + } + if _, err := Chirp(10, 1, 50, 100); err == nil { + t.Fatal("an edge on the Nyquist frequency: want an error") + } + if _, err := Chirp(10, 60, 1, 100); err == nil { + t.Fatal("a negative edge beyond Nyquist in magnitude: want an error") + } + // The one-sample chirp is legal and carries the phase-zero zero. + one, err := Chirp(1, 7, 9, 100) + if err != nil { + t.Fatalf("Chirp(1, ...): %v", err) + } + if len(one) != 1 || one[0] != 0 { + t.Fatalf("the one-sample chirp = %v, want [0]", one) + } + // A downward sweep inside the guard is legal. + if _, err := Chirp(16, 900, 100, 8000); err != nil { + t.Fatalf("Chirp downward: %v", err) + } +} + +// TestChirpDeterministic redraws one sweep and requires the same bits, +// the contract every signal builder here carries. +func TestChirpDeterministic(t *testing.T) { + a, err := Chirp(256, 4, 90, 512) + if err != nil { + t.Fatal(err) + } + b, err := Chirp(256, 4, 90, 512) + if err != nil { + t.Fatal(err) + } + for i := range a { + if math.Float64bits(a[i]) != math.Float64bits(b[i]) { + t.Fatalf("sample %d moved between identical calls", i) + } + } +} diff --git a/signal/complex_refusal_pins_test.go b/signal/complex_refusal_pins_test.go new file mode 100644 index 0000000..3db15c9 --- /dev/null +++ b/signal/complex_refusal_pins_test.go @@ -0,0 +1,119 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Complex-refusal pins: complex inputs that reached FloatAt's nil-int +// branch, degenerate kernels and empty windows that answered NaN, and +// contract gaps the sibling entry points had already closed. + +// TestResampleComplexRefusal: Resample and Decimate widened a complex +// series through widenFloats, which panicked instead of refusing. +func TestResampleComplexRefusal(t *testing.T) { + c := mustComplexes(t, []complex128{1 + 1i, 2, 3 - 1i, 4}, 4) + if _, err := Resample(c, 2, 1, 0); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("Resample on complex: err = %v", err) + } + if _, err := Decimate(c, 2, 0); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("Decimate on complex: err = %v", err) + } +} + +// TestNUFFTComplexCoordinates: complex coordinates reached FloatAt's +// nil-int branch and panicked. +func TestNUFFTComplexCoordinates(t *testing.T) { + x := mustComplexes(t, []complex128{0.1 + 0.2i, -0.2}, 2) + c := mustComplexes(t, []complex128{1, 2i}, 2) + if _, err := NUFFTType1(x, c, 8); err == nil || !strings.Contains(err.Error(), "real") { + t.Fatalf("NUFFTType1 with complex coordinates: err = %v", err) + } +} + +// TestAnalyticSignalComplexRefusal: the analytic signal is defined for +// a real series; a complex one was silently transformed. +func TestAnalyticSignalComplexRefusal(t *testing.T) { + c := mustComplexes(t, []complex128{1, 2, 3, 4}, 4) + if _, err := AnalyticSignal(c); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("AnalyticSignal on complex: err = %v", err) + } +} + +// TestConvComplexBiasRefusal: the convolutions gate the input and the +// kernel on dtype but widened the bias through FloatAt, whose complex +// branch reads a nil payload and panics. A complex bias is refused +// like a complex kernel. +func TestConvComplexBiasRefusal(t *testing.T) { + bias := mustComplexes(t, []complex128{1 + 2i}, 1) + in2 := mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 1, 1, 2, 4) + ker2 := mustFloats(t, []float64{1, 0, 0, 1}, 1, 1, 2, 2) + in3 := mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 1, 1, 2, 2, 2) + ker3 := mustFloats(t, []float64{1, 1, 1, 1, 1, 1, 1, 1}, 1, 1, 2, 2, 2) + in1 := mustFloats(t, []float64{1, 2, 3, 4}, 1, 1, 4) + ker1 := mustFloats(t, []float64{1, 1}, 1, 1, 2) + for name, fn := range map[string]func() (*core.Array, error){ + "Conv1D": func() (*core.Array, error) { return Conv1D(in1, ker1, bias, 1, 0, 1) }, + "Conv2D": func() (*core.Array, error) { return Conv2D(in2, ker2, bias, 1, 0) }, + "Conv3D": func() (*core.Array, error) { return Conv3D(in3, ker3, bias, 1, [3]int{0, 0, 0}, [3]int{1, 1, 1}) }, + "ConvTranspose2D": func() (*core.Array, error) { return ConvTranspose2D(in2, ker2, bias, 1, 0) }, + } { + if _, err := fn(); err == nil || !strings.Contains(err.Error(), "complex bias") { + t.Fatalf("%s with a complex bias: err = %v", name, err) + } + } +} + +// TestDecimateOneTap: taps = 1 divided by taps−1 = 0 in the Kaiser +// window and every output sample came back NaN; a one-tap kernel is +// the identity, so decimation keeps every factor-th sample. +func TestDecimateOneTap(t *testing.T) { + x := mustFromFloats(t, []float64{10, 11, 12, 13, 14, 15, 16, 17}, 8) + out, err := Decimate(x, 2, 1) + if err != nil { + t.Fatalf("Decimate taps=1: %v", err) + } + if out.Len() != 4 { + t.Fatalf("output length %d, want 4", out.Len()) + } + for i := range 4 { + if got := out.FloatAt(i); got != float64(10+2*i) { + t.Fatalf("out[%d] = %v, want %d", i, got, 10+2*i) + } + } +} + +// TestAdaptivePoolEmptySpatial: a zero-length spatial dimension left +// every window empty and max answered -Inf while avg divided by zero; +// both are errors now, in every rank. +func TestAdaptivePoolEmptySpatial(t *testing.T) { + empty1D, err := core.Zeros(core.Float, 1, 1, 0) + if err != nil { + t.Fatal(err) + } + empty2D, err := core.Zeros(core.Float, 1, 1, 0, 4) + if err != nil { + t.Fatal(err) + } + empty3D, err := core.Zeros(core.Float, 1, 1, 2, 0, 4) + if err != nil { + t.Fatal(err) + } + for name, fn := range map[string]func() (*core.Array, error){ + "max1D": func() (*core.Array, error) { return AdaptiveMaxPool1D(empty1D, 2) }, + "max2D": func() (*core.Array, error) { return AdaptiveMaxPool2D(empty2D, 2, 2) }, + "max3D": func() (*core.Array, error) { return AdaptiveMaxPool3D(empty3D, 2, 2, 2) }, + "avg1D": func() (*core.Array, error) { return AdaptiveAvgPool1D(empty1D, 2) }, + "avg2D": func() (*core.Array, error) { return AdaptiveAvgPool2D(empty2D, 2, 2) }, + "avg3D": func() (*core.Array, error) { return AdaptiveAvgPool3D(empty3D, 2, 2, 2) }, + } { + if _, err := fn(); err == nil || !strings.Contains(err.Error(), "spatial") { + t.Fatalf("%s on an empty spatial dimension: err = %v", name, err) + } + } +} diff --git a/signal/conv.go b/signal/conv.go new file mode 100644 index 0000000..4624196 --- /dev/null +++ b/signal/conv.go @@ -0,0 +1,1011 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "sourcedock.dev/petrbalvin/tensor/internal/engine" + +// 2-D convolution. The forward pass walks every output +// position and sums the contribution of every (cIn, kH, kW) input +// patch times the kernel weights. NCHW layout: input is (N, C_in, +// H, W), kernel is (C_out, C_in, kH, kW), output is (N, C_out, H_out, +// W_out). padding applies to both spatial dims (a single int for now); +// stride is also a single int. bias is added per output channel. +// +// Performance note: this is a direct O(N * C_out * C_in * H_out * +// W_out * kH * kW) implementation, not an im2col + GEMM. The kernels +// hold two constraints fixed: +// +// - Tap order. Each output element accumulates its taps in exactly +// the (cIn, kH, kW) nesting the kernels have always used; float64 +// addition is not associative, so any other order changes bits. +// Loop interchange that keeps this per-element order (moving the +// output position innermost) is safe; reordering taps is not. +// - Parallel split. Work is distributed over output rows, which own +// disjoint output elements, so no split can influence a sum. +// +// Non-float inputs widen once into a scratch payload before the loops: +// FloatAt widens exactly, so the multiplied values are bit-identical +// to the per-element accessor reads. + +// Conv2D performs a 2-D convolution forward pass. +func Conv2D(input, kernel, bias *core.Array, stride, padding int) (*core.Array, error) { + return conv2DImpl(input, kernel, bias, stride, padding, 1) +} + +// Conv2DGroups performs a 2-D convolution with `groups` for grouped +// (groups=2) and depthwise (groups=C_in) convolutions. When groups=1 +// the call is identical to Conv2D. +func Conv2DGroups(input, kernel, bias *core.Array, stride, padding, groups int) (*core.Array, error) { + return conv2DImpl(input, kernel, bias, stride, padding, groups) +} + +// conv2DImpl is the shared implementation behind Conv2D and +// Conv2DGroups. groups=1 is the standard dense convolution; groups>1 +// is a depthwise convolution when groups == c_in, or a grouped +// convolution otherwise (the kernel's in-channels is c_in/groups per +// group). +func conv2DImpl(input, kernel, bias *core.Array, stride, padding, groups int) (*core.Array, error) { + if input.Dtype() == core.Complex { + return nil, base.Errf("Conv2D: complex arrays are not supported") + } + if kernel.Dtype() == core.Complex { + return nil, base.Errf("Conv2D: complex kernel is not supported") + } + if input.NDim() != 4 { + return nil, base.Errf("Conv2D: input must be 4-D (N, C_in, H, W), got shape %s", base.ShapeText(input.Shape())) + } + if kernel.NDim() != 4 { + return nil, base.Errf("Conv2D: kernel must be 4-D (C_out, C_in, kH, kW), got shape %s", base.ShapeText(kernel.Shape())) + } + if stride < 1 { + return nil, base.Errf("Conv2D: stride must be at least 1, got %d", stride) + } + if padding < 0 { + return nil, base.Errf("Conv2D: padding must be non-negative, got %d", padding) + } + n, cIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3] + cOut, cInK, kH, kW := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2], kernel.Shape()[3] + if groups < 1 { + return nil, base.Errf("Conv2D: groups must be at least 1, got %d", groups) + } + if cIn%groups != 0 { + return nil, base.Errf("Conv2D: input channels %d must be divisible by groups %d", cIn, groups) + } + if cOut%groups != 0 { + return nil, base.Errf("Conv2D: output channels %d must be divisible by groups %d", cOut, groups) + } + cInPerGroup := cIn / groups + cOutPerGroup := cOut / groups + if cInK != cInPerGroup { + return nil, base.Errf("Conv2D: kernel in-channels %d does not match cIn/groups %d", cInK, cInPerGroup) + } + // Guard the numerators before the division: a negative numerator + // truncates toward zero and masquerades as a 1-wide output. + hNum := hIn + 2*padding - kH + wNum := wIn + 2*padding - kW + if hNum < 0 || wNum < 0 { + return nil, base.Errf("Conv2D: kernel %dx%d with padding %d does not fit the %dx%d input", kH, kW, padding, hIn, wIn) + } + hOut := hNum/stride + 1 + wOut := wNum/stride + 1 + if hOut < 1 || wOut < 1 { + return nil, base.Errf("Conv2D: output is empty (H_out=%d, W_out=%d)", hOut, wOut) + } + if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) { + return nil, base.Errf("Conv2D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape())) + } + // The bias widens through FloatAt like the payloads, which reads + // only real dtypes: a complex bias is refused instead of widened. + if bias != nil && bias.Dtype() == core.Complex { + return nil, base.Errf("Conv2D: complex bias is not supported") + } + out, oerr := core.Zeros(core.Float, []int{n, cOut, hOut, wOut}...) + if oerr != nil { + return nil, oerr + } + // The fast path streams raw float64 payloads; other dtypes widen + // once into scratch (exact, see widenFloats) so the loop below + // never branches on dtype. Dispatching once here keeps the + // per-element work free of accessor calls and bounds on the shape + // of the read. + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + kerF := kernel.RawFloats() + if kerF == nil || kernel.Strided() { + kerF = widenFloats(kernel) + } + var biasF []float64 + if bias != nil { + biasF = widenFloats(bias) + } + outF := out.RawFloats() + // Blocked over output channels: one work item is one (batch, group, + // output row, channel block) and its accumulator covers every output + // channel of the block at once. A tap's input row then feeds every + // channel of the block from the lowest cache instead of being read + // again for each of them, which is what the per-channel walk spent + // its bandwidth on. The tap walk stays (cIn, kH, kW) ascending and + // every output element is updated by exactly one tap of the walk, so + // each sum sees the same addends in the same order as before, and the + // bias is added last, as before. When the row is too wide for the + // block to stay in cache the block shrinks to one channel, which is + // the per-channel walk again. + ocBlock := max(1, min(cOutPerGroup, convAccCap/wOut)) + blocksPerRow := (cOutPerGroup + ocBlock - 1) / ocBlock + // Per-channel kernel offsets within the block: the weight base of the + // block's first channel is added per item, the rest follow from the + // channel index, so a tap reads its weight without multiplying the + // channel out. + kOff := make([]int, ocBlock) + for i := range kOff { + kOff[i] = i * cInPerGroup * kH * kW + } + items := n * hOut * groups * blocksPerRow + floor := workFloorFor(wOut * ocBlock * cInPerGroup * kH * kW) + engine.ParallelMin(items, floor, func(bs, be int) { + acc := engine.GetFloat64Buf(ocBlock * wOut) + defer engine.PutFloat64Buf(acc) + for item := bs; item < be; item++ { + ob := item % blocksPerRow + rest := item / blocksPerRow + g := rest % groups + rest /= groups + oh := rest % hOut + batch := rest / hOut + oc0 := g*cOutPerGroup + ob*ocBlock + oc1 := min(oc0+ocBlock, (g+1)*cOutPerGroup) + clear(acc) + chanBase := batch*cIn*hIn*wIn + g*cInPerGroup*hIn*wIn + ocBase := oc0 * cInPerGroup * kH * kW + for icInGroup := range cInPerGroup { + inChan := chanBase + icInGroup*hIn*wIn + tapBase := icInGroup * kH * kW + for kh := range kH { + ih := oh*stride + kh - padding + if ih < 0 || ih >= hIn { + continue + } + inRow := inChan + ih*wIn + for kw := range kW { + // The outputs this tap reaches: ow*stride + // must land the input column inside the row. + // The reach is a property of the tap alone, so + // it is computed once for the whole block. + lo := 0 + if d := padding - kw; d > 0 { + lo = (d + stride - 1) / stride + } + hi := wOut + if m := wIn - 1 + padding - kw; m < 0 { + continue + } else if t := m/stride + 1; t < hi { + hi = t + } + tap := kw - padding + if hi <= lo { + // The tap reaches no output here (its + // whole window lies in the padding): the + // unit-stride sub-slices below would get + // a high bound under their low one. + continue + } + tapOff := tapBase + kh*kW + kw + if stride == 1 { + // Unit stride lets both sides range + // over paired sub-slices: no bounds + // checks and no index multiply in + // the hot loop. + src := inF[inRow+lo+tap : inRow+hi+tap] + for o := range oc1 - oc0 { + k := kerF[ocBase+kOff[o]+tapOff] + dst := acc[o*wOut+lo : o*wOut+hi] + for j, v := range src { + dst[j] += v * k + } + } + continue + } + for o := range oc1 - oc0 { + k := kerF[ocBase+kOff[o]+tapOff] + base := o * wOut + for ow := lo; ow < hi; ow++ { + acc[base+ow] += inF[inRow+ow*stride+tap] * k + } + } + } + } + } + for o := range oc1 - oc0 { + oc := oc0 + o + b := 0.0 + if biasF != nil { + b = biasF[oc] + } + off := (batch*cOut + oc) * hOut * wOut + base := o * wOut + for ow := range wOut { + outF[off+oh*wOut+ow] = acc[base+ow] + b + } + } + } + }) + return out, nil +} + +// conv1DRowBlock is the output-length grain the 1-D kernel splits rows +// into when a row alone would starve the worker pool (a single long +// signal with a handful of output channels is exactly that case). +const conv1DRowBlock = 512 + +// convAccCap bounds one work item's channel-block accumulator, in +// float64 entries: a block of channels·row values is what the taps of +// a whole row feed, and it stays cheap only while it fits the lowest +// cache. Above the cap the block covers fewer channels, down to one, +// where the kernel degenerates to the per-channel walk it replaced. +// Every channel-blocked kernel shares it. +const convAccCap = 1 << 12 + +// Conv1D performs a 1-D convolution forward pass. input is (N, C_in, +// L), kernel is (C_out, C_in, kL), output is (N, C_out, L_out). NCL +// layout. +func Conv1D(input, kernel, bias *core.Array, stride, padding, dilation int) (*core.Array, error) { + if input.Dtype() == core.Complex { + return nil, base.Errf("Conv1D: complex arrays are not supported") + } + if kernel.Dtype() == core.Complex { + return nil, base.Errf("Conv1D: complex kernel is not supported") + } + if input.NDim() != 3 { + return nil, base.Errf("Conv1D: input must be 3-D (N, C_in, L), got shape %s", base.ShapeText(input.Shape())) + } + if kernel.NDim() != 3 { + return nil, base.Errf("Conv1D: kernel must be 3-D (C_out, C_in, kL), got shape %s", base.ShapeText(kernel.Shape())) + } + if stride < 1 || dilation < 1 || padding < 0 { + return nil, base.Errf("Conv1D: stride/dilation must be ≥1, padding ≥0") + } + if dilation != 1 { + return nil, base.Errf("Conv1D: dilation != 1 is not implemented yet") + } + n, cIn, lIn := input.Shape()[0], input.Shape()[1], input.Shape()[2] + cOut, cInK, kL := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2] + if cInK != cIn { + return nil, base.Errf("Conv1D: kernel in-channels %d does not match input %d", cInK, cIn) + } + lNum := lIn + 2*padding - kL + if lNum < 0 { + return nil, base.Errf("Conv1D: kernel length %d with padding %d does not fit the length-%d input", kL, padding, lIn) + } + lOut := lNum/stride + 1 + if lOut < 1 { + return nil, base.Errf("Conv1D: output is empty (L_out=%d)", lOut) + } + if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) { + return nil, base.Errf("Conv1D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape())) + } + if bias != nil && bias.Dtype() == core.Complex { + return nil, base.Errf("Conv1D: complex bias is not supported") + } + out, oerr := core.Zeros(core.Float, []int{n, cOut, lOut}...) + if oerr != nil { + return nil, oerr + } + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + kerF := kernel.RawFloats() + if kerF == nil || kernel.Strided() { + kerF = widenFloats(kernel) + } + var biasF []float64 + if bias != nil { + biasF = widenFloats(bias) + } + outF := out.RawFloats() + // Work items are output row segments (batch, oc, a block of ol). + // Segmenting keeps a single long signal with few output channels + // from starving the pool; segments are disjoint, and each output + // accumulates its taps in the (cIn, kL) order. The per-tap spans + // are shape-only and are clipped to the row block per item. + spans := convRowSpans(kL, padding, stride, lIn, lOut) + blocks := (lOut + conv1DRowBlock - 1) / conv1DRowBlock + items := n * cOut * blocks + engine.ParallelMin(items, workFloorFor(min(conv1DRowBlock, lOut)*cIn*kL), func(bs, be int) { + acc := engine.GetFloat64Buf(min(conv1DRowBlock, lOut)) + defer engine.PutFloat64Buf(acc) + for item := bs; item < be; item++ { + block := item % blocks + row := item / blocks + oc := row % cOut + batch := row / cOut + ols := block * conv1DRowBlock + ole := min(ols+conv1DRowBlock, lOut) + clear(acc) + chanBase := batch * cIn * lIn + kerBase := oc * cIn * kL + // The interior of the block reached by every tap of the + // kernel: [fa, fb) gets one fused sweep per input channel + // with the running sum held in a register, instead of one + // accumulator pass per tap. Each element still receives + // its taps in (cIn, kL) order, so the accumulated bits + // are unchanged. Outside the register-sized kernel rows + // the per-tap walk below covers the whole block. + fa := max(spans[0].lo, ols) + fb := min(spans[kL-1].hi, ole) + fused := stride == 1 && kL >= 2 && kL <= 7 && fb > fa + for ic := range cIn { + inRow := chanBase + ic*lIn + kerRow := kerBase + ic*kL + for kl := range kL { + sp := spans[kl] + if !sp.valid { + continue + } + k := kerF[kerRow+kl] + // Clip the tap's precomputed span to this row + // block and to the edges of the fused region; + // the clip only narrows. + lo := max(sp.lo, ols) + hi := min(sp.hi, ole) + if fused { + hi = min(hi, fa) + } + tap := sp.tap + if hi > lo { + if stride == 1 { + // Unit stride: paired sub-slices + // keep the hot loop free of bounds + // checks and index multiplies. + src := inF[inRow+lo+tap : inRow+hi+tap] + dst := acc[lo-ols : hi-ols] + for j, v := range src { + dst[j] += v * k + } + } else { + for ol := lo; ol < hi; ol++ { + acc[ol-ols] += inF[inRow+ol*stride+tap] * k + } + } + } + if !fused { + continue + } + lo = max(sp.lo, fb) + hi = min(sp.hi, ole) + if hi <= lo { + continue + } + src := inF[inRow+lo+tap : inRow+hi+tap] + dst := acc[lo-ols : hi-ols] + for j, v := range src { + dst[j] += v * k + } + } + if !fused { + continue + } + // The fused interior sweep: one shifted window of the + // input row, its sub-slices in lockstep, taps in + // registers. + src := inF[inRow+fa-padding : inRow+fb-padding+kL-1] + dst := acc[fa-ols : fb-ols] + m := len(dst) + switch kL { + case 2: + k0, k1 := kerF[kerRow], kerF[kerRow+1] + s1 := src[1 : 1+m] + for i, v0 := range src[:m] { + d := dst[i] + v0*k0 + dst[i] = d + s1[i]*k1 + } + case 3: + k0, k1, k2 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2] + s1 := src[1 : 1+m] + s2 := src[2 : 2+m] + for i, v0 := range src[:m] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + dst[i] = d + s2[i]*k2 + } + case 4: + k0, k1, k2, k3 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3] + s1 := src[1 : 1+m] + s2 := src[2 : 2+m] + s3 := src[3 : 3+m] + for i, v0 := range src[:m] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + dst[i] = d + s3[i]*k3 + } + case 5: + k0, k1, k2, k3, k4 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4] + s1 := src[1 : 1+m] + s2 := src[2 : 2+m] + s3 := src[3 : 3+m] + s4 := src[4 : 4+m] + for i, v0 := range src[:m] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + d += s3[i] * k3 + dst[i] = d + s4[i]*k4 + } + case 6: + k0, k1, k2, k3, k4, k5 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5] + s1 := src[1 : 1+m] + s2 := src[2 : 2+m] + s3 := src[3 : 3+m] + s4 := src[4 : 4+m] + s5 := src[5 : 5+m] + for i, v0 := range src[:m] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + d += s3[i] * k3 + d += s4[i] * k4 + dst[i] = d + s5[i]*k5 + } + case 7: + k0, k1, k2, k3, k4, k5, k6 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5], kerF[kerRow+6] + s1 := src[1 : 1+m] + s2 := src[2 : 2+m] + s3 := src[3 : 3+m] + s4 := src[4 : 4+m] + s5 := src[5 : 5+m] + s6 := src[6 : 6+m] + for i, v0 := range src[:m] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + d += s3[i] * k3 + d += s4[i] * k4 + d += s5[i] * k5 + dst[i] = d + s6[i]*k6 + } + } + } + b := 0.0 + if biasF != nil { + b = biasF[oc] + } + off := batch*cOut*lOut + oc*lOut + for ol := ols; ol < ole; ol++ { + outF[off+ol] = acc[ol-ols] + b + } + } + }) + return out, nil +} + +// convGatherSpan is the reach of one transposed-convolution kernel +// column: the outputs it feeds sit stride apart on the phase lattice, +// at lattice indices [j0, j0+count), and read the input row from +// base on. first is the lowest output column reached and phase is its +// residue modulo the stride, which decides which accumulator partition +// the tap lands in. A tap with an empty reach is marked invalid. +type convGatherSpan struct { + first, j0, count, base, phase int + valid bool +} + +// convGatherSpans precomputes the per-tap gather spans for one +// transposed kernel row. The values depend on the kernel column and +// the shape alone, so the hot loop reads them from this table instead +// of dividing per output position. +func convGatherSpans(kW, padding, stride, wIn, wOut int) []convGatherSpan { + spans := make([]convGatherSpan, kW) + for kw := range kW { + lo := 0 + if d := kw - padding; d > 0 { + lo = d + } + hi := wOut + if m := wIn*stride + kw - padding; m < hi { + hi = m + } + if hi <= lo { + continue + } + r := (kw - padding - lo) % stride + if r < 0 { + r += stride + } + first := lo + r + count := (hi - first + stride - 1) / stride + if count <= 0 { + continue + } + spans[kw].first = first + spans[kw].count = count + spans[kw].j0 = first / stride + spans[kw].base = (first + padding - kw) / stride + spans[kw].phase = first % stride + spans[kw].valid = true + } + return spans +} + +// convTapSpan is the reach of one kernel tap along an output row: the +// tap lands on outputs [lo, hi) and reads the input row offset by tap +// from them. A tap whose whole window falls into the padding is marked +// invalid, so the hot loop skips it without recomputing why. +type convTapSpan struct { + lo, hi, tap int + valid bool +} + +// convRowSpans precomputes the per-tap spans of one kernel row for +// stride and padding. The bounds depend on the kernel column and the +// shape alone, never on the output position, so the innermost loops +// load them from this table instead of dividing once per tap. +func convRowSpans(kW, padW, stride, wIn, wOut int) []convTapSpan { + spans := make([]convTapSpan, kW) + for kw := range kW { + lo := 0 + if d := padW - kw; d > 0 { + lo = (d + stride - 1) / stride + } + // lo and tap are stored whatever the reach: the fused region + // bounds read spans[0].lo and spans[kW-1].hi, so an extreme tap + // that reaches nothing must leave an empty window (hi 0) rather + // than a zero lo beside a stale hi. + spans[kw].lo = lo + spans[kw].tap = kw - padW + if m := wIn - 1 + padW - kw; m < 0 { + continue + } else if t := m/stride + 1; t < wOut { + spans[kw].hi = t + } else { + spans[kw].hi = wOut + } + if spans[kw].hi <= lo { + // The tap reaches no output here (its whole window + // lies in the padding): the unit-stride sub-slices + // would get a high bound under their low one. + spans[kw].hi = 0 + continue + } + spans[kw].valid = true + } + return spans +} + +// Conv3D performs a 3-D convolution forward pass. input is (N, C_in, +// D, H, W), kernel is (C_out, C_in, kD, kH, kW), output is (N, C_out, +// D_out, H_out, W_out). NCDHW layout. +func Conv3D(input, kernel, bias *core.Array, stride int, padding, dilation [3]int) (*core.Array, error) { + if input.Dtype() == core.Complex { + return nil, base.Errf("Conv3D: complex arrays are not supported") + } + if kernel.Dtype() == core.Complex { + return nil, base.Errf("Conv3D: complex kernel is not supported") + } + if input.NDim() != 5 { + return nil, base.Errf("Conv3D: input must be 5-D, got shape %s", base.ShapeText(input.Shape())) + } + if kernel.NDim() != 5 { + return nil, base.Errf("Conv3D: kernel must be 5-D (C_out, C_in, kD, kH, kW), got shape %s", base.ShapeText(kernel.Shape())) + } + if stride < 1 { + return nil, base.Errf("Conv3D: stride must be at least 1, got %d", stride) + } + for d, dil := range dilation { + if dil != 1 { + return nil, base.Errf("Conv3D: dilation != 1 not implemented yet (dim %d)", d) + } + } + n, cIn, dIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3], input.Shape()[4] + cOut, cInK, kD, kH, kW := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2], kernel.Shape()[3], kernel.Shape()[4] + if cInK != cIn { + return nil, base.Errf("Conv3D: kernel in-channels %d does not match input %d", cInK, cIn) + } + padD, padH, padW := padding[0], padding[1], padding[2] + // Guard the numerators before the division: a negative numerator + // truncates toward zero and masquerades as a 1-wide output. + dNum := dIn + 2*padD - kD + hNum := hIn + 2*padH - kH + wNum := wIn + 2*padW - kW + if dNum < 0 || hNum < 0 || wNum < 0 { + return nil, base.Errf("Conv3D: kernel %dx%dx%d with padding %v does not fit the input", kD, kH, kW, padding) + } + dOut := dNum/stride + 1 + hOut := hNum/stride + 1 + wOut := wNum/stride + 1 + if dOut < 1 || hOut < 1 || wOut < 1 { + return nil, base.Errf("Conv3D: output is empty") + } + if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) { + return nil, base.Errf("Conv3D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape())) + } + if bias != nil && bias.Dtype() == core.Complex { + return nil, base.Errf("Conv3D: complex bias is not supported") + } + out, oerr := core.Zeros(core.Float, []int{n, cOut, dOut, hOut, wOut}...) + if oerr != nil { + return nil, oerr + } + strideHW := wIn + strideDHW := hIn * wIn + strideCDHW := dIn * hIn * wIn + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + kerF := kernel.RawFloats() + if kerF == nil || kernel.Strided() { + kerF = widenFloats(kernel) + } + var biasF []float64 + if bias != nil { + biasF = widenFloats(bias) + } + outF := out.RawFloats() + // One work item per output depth-height row (batch, oc, od, oh). + // Rows are disjoint; the taps of one output accumulate in the + // (cIn, kD, kH, kW) order. The channel-blocked shape was tried + // here and measured: it paid 12 to 14 percent at 32 and 128 + // channels and lost 12 percent at 4, where the per-tap channel + // loop costs more than the re-reads it saves, so the per-channel + // walk stays. The per-tap column spans are shape-only, so they + // are read from a table built once per call. + spans := convRowSpans(kW, padW, stride, wIn, wOut) + items := n * cOut * dOut * hOut + engine.ParallelMin(items, workFloorFor(wOut*cIn*kD*kH*kW), func(bs, be int) { + acc := engine.GetFloat64Buf(wOut) + defer engine.PutFloat64Buf(acc) + for item := bs; item < be; item++ { + oh := item % hOut + t := item / hOut + od := t % dOut + t /= dOut + oc := t % cOut + batch := t / cOut + clear(acc) + for ic := range cIn { + inChan := batch*cIn*strideCDHW + ic*strideCDHW + kerChan := oc*cInK*kD*kH*kW + ic*kD*kH*kW + for kd := range kD { + id := od*stride + kd - padD + if id < 0 || id >= dIn { + continue + } + inDepth := inChan + id*strideDHW + kerDepth := kerChan + kd*kH*kW + for kh := range kH { + ih := oh*stride + kh - padH + if ih < 0 || ih >= hIn { + continue + } + inRow := inDepth + ih*strideHW + kerRow := kerDepth + kh*kW + if stride == 1 && kW >= 2 && kW <= 7 { + // Unit stride with a + // register-sized kernel row. The + // interior [a, b) is reached by + // every tap, so one sweep adds all + // taps there with the running sum + // held in a register instead of + // re-reading the accumulator row + // once per tap. Each element still + // receives its taps in kw order, so + // the accumulated bits are unchanged. + a := spans[0].lo + b := spans[kW-1].hi + if b <= a { + // No interior: an empty window + // hands every tap's span to the + // second part below in one + // piece, exactly once. + a, b = 0, 0 + } + for kw := range kW { + sp := spans[kw] + if !sp.valid { + continue + } + k := kerF[kerRow+kw] + lo := sp.lo + hi := min(sp.hi, a) + if hi > lo { + src := inF[inRow+lo+sp.tap : inRow+hi+sp.tap] + dst := acc[lo:hi] + for j, v := range src { + dst[j] += v * k + } + } + lo = max(sp.lo, b) + hi = sp.hi + if hi > lo { + src := inF[inRow+lo+sp.tap : inRow+hi+sp.tap] + dst := acc[lo:hi] + for j, v := range src { + dst[j] += v * k + } + } + } + if b <= a { + continue + } + // The interior sweep. The taps + // live in registers; the sub-slices + // of one shifted window keep the + // loop free of bounds checks and + // index multiplies. + src := inF[inRow+a-padW : inRow+b-padW+kW-1] + dst := acc[a:b] + n := len(dst) + switch kW { + case 2: + k0, k1 := kerF[kerRow], kerF[kerRow+1] + s1 := src[1 : 1+n] + for i, v0 := range src[:n] { + d := dst[i] + v0*k0 + dst[i] = d + s1[i]*k1 + } + case 3: + k0, k1, k2 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2] + s1 := src[1 : 1+n] + s2 := src[2 : 2+n] + for i, v0 := range src[:n] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + dst[i] = d + s2[i]*k2 + } + case 4: + k0, k1, k2, k3 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3] + s1 := src[1 : 1+n] + s2 := src[2 : 2+n] + s3 := src[3 : 3+n] + for i, v0 := range src[:n] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + dst[i] = d + s3[i]*k3 + } + case 5: + k0, k1, k2, k3, k4 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4] + s1 := src[1 : 1+n] + s2 := src[2 : 2+n] + s3 := src[3 : 3+n] + s4 := src[4 : 4+n] + for i, v0 := range src[:n] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + d += s3[i] * k3 + dst[i] = d + s4[i]*k4 + } + case 6: + k0, k1, k2, k3, k4, k5 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5] + s1 := src[1 : 1+n] + s2 := src[2 : 2+n] + s3 := src[3 : 3+n] + s4 := src[4 : 4+n] + s5 := src[5 : 5+n] + for i, v0 := range src[:n] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + d += s3[i] * k3 + d += s4[i] * k4 + dst[i] = d + s5[i]*k5 + } + case 7: + k0, k1, k2, k3, k4, k5, k6 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5], kerF[kerRow+6] + s1 := src[1 : 1+n] + s2 := src[2 : 2+n] + s3 := src[3 : 3+n] + s4 := src[4 : 4+n] + s5 := src[5 : 5+n] + s6 := src[6 : 6+n] + for i, v0 := range src[:n] { + d := dst[i] + v0*k0 + d += s1[i] * k1 + d += s2[i] * k2 + d += s3[i] * k3 + d += s4[i] * k4 + d += s5[i] * k5 + dst[i] = d + s6[i]*k6 + } + } + continue + } + for kw := range kW { + sp := spans[kw] + if !sp.valid { + continue + } + k := kerF[kerRow+kw] + for ow := sp.lo; ow < sp.hi; ow++ { + acc[ow] += inF[inRow+ow*stride+sp.tap] * k + } + } + } + } + } + b := 0.0 + if biasF != nil { + b = biasF[oc] + } + off := ((batch*cOut+oc)*dOut+od)*hOut*wOut + oh*wOut + for ow := range wOut { + outF[off+ow] = acc[ow] + b + } + } + }) + return out, nil +} + +// ConvTranspose2D performs the transposed convolution (sometimes +// called "deconvolution"). Input (N, C_in, H, W), kernel (C_in, +// C_out, kH, kW), output (N, C_out, H_out, W_out) where H_out = +// (H-1)*stride - 2*padding + kH (and analogously for W). bias is +// optional, applied per output channel. +func ConvTranspose2D(input, kernel, bias *core.Array, stride, padding int) (*core.Array, error) { + if input.Dtype() == core.Complex || kernel.Dtype() == core.Complex { + return nil, base.Errf("ConvTranspose2D: complex arrays are not supported") + } + if input.NDim() != 4 { + return nil, base.Errf("ConvTranspose2D: input must be 4-D (N, C_in, H, W), got shape %s", base.ShapeText(input.Shape())) + } + if kernel.NDim() != 4 { + return nil, base.Errf("ConvTranspose2D: kernel must be 4-D (C_in, C_out, kH, kW), got shape %s", base.ShapeText(kernel.Shape())) + } + if stride < 1 { + return nil, base.Errf("ConvTranspose2D: stride must be at least 1, got %d", stride) + } + if padding < 0 { + return nil, base.Errf("ConvTranspose2D: padding must be non-negative, got %d", padding) + } + n, cIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3] + cInK, cOut, kH, kW := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2], kernel.Shape()[3] + if cInK != cIn { + return nil, base.Errf("ConvTranspose2D: kernel in-channels %d does not match input %d", cInK, cIn) + } + hOut := (hIn-1)*stride - 2*padding + kH + wOut := (wIn-1)*stride - 2*padding + kW + if hOut < 1 || wOut < 1 { + return nil, base.Errf("ConvTranspose2D: output is empty") + } + if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) { + return nil, base.Errf("ConvTranspose2D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape())) + } + if bias != nil && bias.Dtype() == core.Complex { + return nil, base.Errf("ConvTranspose2D: complex bias is not supported") + } + out, oerr := core.Zeros(core.Float, []int{n, cOut, hOut, wOut}...) + if oerr != nil { + return nil, oerr + } + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + kerF := kernel.RawFloats() + if kerF == nil || kernel.Strided() { + kerF = widenFloats(kernel) + } + var biasF []float64 + if bias != nil { + biasF = widenFloats(bias) + } + outF := out.RawFloats() + // The scatter form (walking input positions and adding into every + // output they touch) becomes a gather here: each output row sums + // the input positions that map onto it. Blocked over output + // channels, as Conv2D is: one work item is one (batch, output row, + // channel block) and its accumulator holds every output channel of + // the block at once, so a gathered input value is spent on every + // channel of the block instead of being gathered again for each of + // them. For a single output the contributions still arrive in the + // (cIn, kH, kW) order the scatter visited them, and the bias is + // added last, so the accumulated bits are unchanged. Rows are + // disjoint, so the parallel split cannot touch a sum. + // + // The accumulator is partitioned by phase: a tap with reversed + // column kw - padding only ever touches outputs congruent to it + // modulo the stride, so each phase owns a contiguous sub-lattice + // of every output row. A tap's gather then reads a contiguous + // input window into a contiguous accumulator window, with no + // divisions and no strided stores in the hot loop. + ocBlock := max(1, min(cOut, convAccCap/wOut)) + blocksPerRow := (cOut + ocBlock - 1) / ocBlock + rowCap := (wOut + stride - 1) / stride + spans := convGatherSpans(kW, padding, stride, wIn, wOut) + items := n * hOut * blocksPerRow + engine.ParallelMin(items, workFloorFor(wOut*ocBlock*cIn*kH*kW), func(bs, be int) { + acc := engine.GetFloat64Buf(ocBlock * stride * rowCap) + defer engine.PutFloat64Buf(acc) + var ihMap [16]int + mapped := min(kH, len(ihMap)) + for item := bs; item < be; item++ { + ob := item % blocksPerRow + rest := item / blocksPerRow + oh := rest % hOut + batch := rest / hOut + oc0 := ob * ocBlock + oc1 := min(oc0+ocBlock, cOut) + clear(acc) + // Map (oh) back to the input rows once per item: the + // scatter only paired positions with oh + padding - kh + // divisible by the stride and inside the input, and the + // mapping does not depend on the channel. + for kh := range mapped { + ih := -1 + if ihNum := oh + padding - kh; ihNum%stride == 0 { + if t := ihNum / stride; t >= 0 && t < hIn { + ih = t + } + } + ihMap[kh] = ih + } + for ic := range cIn { + inChan := batch*cIn*hIn*wIn + ic*hIn*wIn + icBase := ic * cOut * kH * kW + for kh := range kH { + var ih int + if kh < mapped { + ih = ihMap[kh] + } else { + ihNum := oh + padding - kh + if ihNum%stride != 0 { + continue + } + ih = ihNum / stride + if ih < 0 || ih >= hIn { + continue + } + } + if ih < 0 { + continue + } + inRow := inChan + ih*wIn + khOff := icBase + kh*kW + for kw := range kW { + sp := spans[kw] + if !sp.valid { + continue + } + tapOff := khOff + oc0*kH*kW + kw + src := inF[inRow+sp.base : inRow+sp.base+sp.count] + for o := range oc1 - oc0 { + k := kerF[tapOff+o*kH*kW] + off := (o*stride + sp.phase) * rowCap + dst := acc[off+sp.j0 : off+sp.j0+sp.count] + for j, v := range src { + dst[j] += v * k + } + } + } + } + } + for o := range oc1 - oc0 { + oc := oc0 + o + b := 0.0 + if biasF != nil { + b = biasF[oc] + } + off := (batch*cOut + oc) * hOut * wOut + for p := range stride { + phaseBase := (o*stride + p) * rowCap + j := 0 + for ow := p; ow < wOut; ow += stride { + outF[off+oh*wOut+ow] = acc[phaseBase+j] + b + j++ + } + } + } + } + }) + return out, nil +} diff --git a/signal/conv_bench_extra_test.go b/signal/conv_bench_extra_test.go new file mode 100644 index 0000000..926a4dc --- /dev/null +++ b/signal/conv_bench_extra_test.go @@ -0,0 +1,74 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "testing" + + core "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Convolution benchmarks guard the direct kernels; the 2-D pair +// already lives in bench_test.go, these add the 1-D path and pooling +// at inference-like sizes. + +func benchSig(b *testing.B, seed, n int, shape ...int) *core.Array { + b.Helper() + v := make([]float64, n) + for i := range v { + v[i] = float64(i%23)*float64(seed%7)*0.5 + float64(i%11) - 5 + } + a, err := core.FromFloats(v, shape...) + if err != nil { + b.Fatal(err) + } + return a +} + +// BenchmarkConv1D measures a 1×16×4096 input with 8 kernels of width +// 7, stride 1. +func BenchmarkConv1D(b *testing.B) { + input := benchSig(b, 1, 16*4096, 1, 16, 4096) + kernel := benchSig(b, 2, 8*16*7, 8, 16, 7) + b.ReportAllocs() + for b.Loop() { + if _, err := Conv1D(input, kernel, nil, 1, 0, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkMaxPool2D measures a 8×64×28×28 input through 2×2 pooling. +func BenchmarkMaxPool2D(b *testing.B) { + input := benchSig(b, 3, 8*64*28*28, 8, 64, 28, 28) + b.ReportAllocs() + for b.Loop() { + if _, err := MaxPool2D(input, 2, 2, 0); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkAvgPool2D is the averaging twin. +func BenchmarkAvgPool2D(b *testing.B) { + input := benchSig(b, 4, 8*64*28*28, 8, 64, 28, 28) + b.ReportAllocs() + for b.Loop() { + if _, err := AvgPool2D(input, 2, 2, 0, false); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkAutocorrelate measures the FFT-based 1-D correlation at a +// speech-like length. +func BenchmarkAutocorrelate(b *testing.B) { + x := benchSig(b, 5, 8192, 8192) + b.ReportAllocs() + for b.Loop() { + if _, err := Autocorrelate(x, 512); err != nil { + b.Fatal(err) + } + } +} diff --git a/signal/conv_bench_test.go b/signal/conv_bench_test.go new file mode 100644 index 0000000..b15cb1e --- /dev/null +++ b/signal/conv_bench_test.go @@ -0,0 +1,48 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "fmt" + "testing" +) + +// Channel-blocked kernel benchmarks at inference-like sizes. The +// channel count is the axis the kernels divide work along, so each +// size climbs a cache level: 4 channels stay in L1, 32 cross L2 and +// 128 make the whole-volume walk pay for memory bandwidth. + +// BenchmarkConv3D measures a (1, C, 64, 64, 64) input through C×C +// kernels of 3×3×3, stride 1, padding 1. +func BenchmarkConv3D(b *testing.B) { + for _, ch := range []int{4, 32, 128} { + b.Run(fmt.Sprintf("%d channels", ch), func(b *testing.B) { + input := benchSig(b, 1, ch*64*64*64, 1, ch, 64, 64, 64) + kernel := benchSig(b, 2, ch*ch*27, ch, ch, 3, 3, 3) + b.ReportAllocs() + for b.Loop() { + if _, err := Conv3D(input, kernel, nil, 1, [3]int{1, 1, 1}, [3]int{1, 1, 1}); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// BenchmarkConvTranspose2D measures a (1, C, 128, 128) input through +// C×C kernels of 3×3 at the upsampling stride of 2, padding 1. +func BenchmarkConvTranspose2D(b *testing.B) { + for _, ch := range []int{4, 32, 128} { + b.Run(fmt.Sprintf("%d channels", ch), func(b *testing.B) { + input := benchSig(b, 3, ch*128*128, 1, ch, 128, 128) + kernel := benchSig(b, 4, ch*ch*9, ch, ch, 3, 3) + b.ReportAllocs() + for b.Loop() { + if _, err := ConvTranspose2D(input, kernel, nil, 2, 1); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/signal/conv_fused_pin_test.go b/signal/conv_fused_pin_test.go new file mode 100644 index 0000000..eed2d4e --- /dev/null +++ b/signal/conv_fused_pin_test.go @@ -0,0 +1,302 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The fused interiors, the phase-lattice gather and the span tables +// are reach-and-order refactors of the per-tap walks: every output +// element must keep the exact accumulated bits. The fixtures are +// small integers, so every product and sum is exact in float64 and the +// comparison against the references below is exact regardless of +// summation order; the references walk the documented shapes and the +// (cIn, k...) tap order. + +func pinInts(n int, seed int) []float64 { + v := make([]float64, n) + x := seed + for i := range n { + x = (x*1103515245 + 12345) & 0x7fffffff + v[i] = float64(x%7 - 3) + } + return v +} + +// refConv1D is the documented forward pass: out[n][oc][ol] is the bias +// plus the (ic, kL)-ordered tap sum of in[n][ic][ol*stride + kL - +// padding] against ker[oc][ic][kL]. +func refConv1D(in []float64, n, cIn, lIn int, ker []float64, cOut, kL int, bias []float64, stride, padding int) []float64 { + lOut := (lIn+2*padding-kL)/stride + 1 + out := make([]float64, n*cOut*lOut) + for b := range n { + for oc := range cOut { + for ol := range lOut { + sum := 0.0 + for ic := range cIn { + for kl := range kL { + il := ol*stride + kl - padding + if il < 0 || il >= lIn { + continue + } + sum += in[b*cIn*lIn+ic*lIn+il] * ker[oc*cIn*kL+ic*kL+kl] + } + } + if bias != nil { + sum += bias[oc] + } + out[b*cOut*lOut+oc*lOut+ol] = sum + } + } + } + return out +} + +// refConv3D is the documented NCDHW forward pass in (ic, kD, kH, kW) +// tap order. +func refConv3D(in []float64, n, cIn, dIn, hIn, wIn int, ker []float64, cOut, kD, kH, kW int, bias []float64, stride, padD, padH, padW int) []float64 { + dOut := (dIn+2*padD-kD)/stride + 1 + hOut := (hIn+2*padH-kH)/stride + 1 + wOut := (wIn+2*padW-kW)/stride + 1 + out := make([]float64, n*cOut*dOut*hOut*wOut) + plane := dIn * hIn * wIn + for b := range n { + for oc := range cOut { + for od := range dOut { + for oh := range hOut { + for ow := range wOut { + sum := 0.0 + for ic := range cIn { + for kd := range kD { + id := od*stride + kd - padD + if id < 0 || id >= dIn { + continue + } + for kh := range kH { + ih := oh*stride + kh - padH + if ih < 0 || ih >= hIn { + continue + } + for kw := range kW { + iw := ow*stride + kw - padW + if iw < 0 || iw >= wIn { + continue + } + sum += in[b*cIn*plane+ic*plane+id*hIn*wIn+ih*wIn+iw] * + ker[oc*cIn*kD*kH*kW+ic*kD*kH*kW+kd*kH*kW+kh*kW+kw] + } + } + } + } + if bias != nil { + sum += bias[oc] + } + out[b*cOut*dOut*hOut*wOut+oc*dOut*hOut*wOut+od*hOut*wOut+oh*wOut+ow] = sum + } + } + } + } + } + return out +} + +// refConvTranspose2D is the documented gather: out[n][oc][oh][ow] is +// the bias plus the (ic, kH, kW) tap sum over the pairs whose reversed +// column lands in the input, reading in[n][ic][ih][iw] against +// ker[ic][oc][kH][kW]. +func refConvTranspose2D(in []float64, n, cIn, hIn, wIn int, ker []float64, cOut, kH, kW int, bias []float64, stride, padding int) []float64 { + hOut := (hIn-1)*stride - 2*padding + kH + wOut := (wIn-1)*stride - 2*padding + kW + out := make([]float64, n*cOut*hOut*wOut) + plane := hIn * wIn + for b := range n { + for oc := range cOut { + for oh := range hOut { + for ow := range wOut { + sum := 0.0 + for ic := range cIn { + for kh := range kH { + hn := oh + padding - kh + if hn%stride != 0 { + continue + } + ih := hn / stride + if ih < 0 || ih >= hIn { + continue + } + for kw := range kW { + wn := ow + padding - kw + if wn%stride != 0 { + continue + } + iw := wn / stride + if iw < 0 || iw >= wIn { + continue + } + sum += in[b*cIn*plane+ic*plane+ih*wIn+iw] * + ker[ic*cOut*kH*kW+oc*kH*kW+kh*kW+kw] + } + } + } + if bias != nil { + sum += bias[oc] + } + out[b*cOut*hOut*wOut+oc*hOut*wOut+oh*wOut+ow] = sum + } + } + } + } + return out +} + +func pinConvEqual(t *testing.T, name string, got, want []float64) { + t.Helper() + if len(got) != len(want) { + t.Fatalf("%s: length %d, want %d", name, len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("%s[%d] = %v, want %v", name, i, got[i], want[i]) + } + } +} + +func TestConv1DFusedInteriorMatchesReference(t *testing.T) { + cases := []struct { + name string + n, cIn, cOut, lIn, kL int + stride, padding int + }{ + {"fused3-long", 2, 3, 2, 3000, 3, 1, 1}, + {"fused7", 1, 2, 3, 200, 7, 1, 3}, + {"fused2-verylong", 1, 1, 1, 5000, 2, 1, 0}, + {"stride2", 2, 2, 2, 100, 5, 2, 2}, + {"span-extreme", 1, 1, 1, 1, 8, 1, 5}, + } + for _, tc := range cases { + in := pinInts(tc.n*tc.cIn*tc.lIn, 1) + ker := pinInts(tc.cOut*tc.cIn*tc.kL, 2) + bias := pinInts(tc.cOut, 3) + ain, err := core.FromFloats(in, tc.n, tc.cIn, tc.lIn) + if err != nil { + t.Fatalf("%s input: %v", tc.name, err) + } + aker, err := core.FromFloats(ker, tc.cOut, tc.cIn, tc.kL) + if err != nil { + t.Fatalf("%s kernel: %v", tc.name, err) + } + abias, err := core.FromFloats(bias, tc.cOut) + if err != nil { + t.Fatalf("%s bias: %v", tc.name, err) + } + out, oerr := Conv1D(ain, aker, abias, tc.stride, tc.padding, 1) + if oerr != nil { + t.Fatalf("%s: %v", tc.name, oerr) + } + want := refConv1D(in, tc.n, tc.cIn, tc.lIn, ker, tc.cOut, tc.kL, bias, tc.stride, tc.padding) + pinConvEqual(t, tc.name, out.RawFloats(), want) + } +} + +func TestConv3DFusedInteriorMatchesReference(t *testing.T) { + cases := []struct { + name string + n, cIn, cOut int + dIn, hIn, wIn int + kD, kH, kW int + stride int + padD, padH, padW int + }{ + {"fused3", 2, 2, 3, 6, 7, 8, 3, 3, 3, 1, 1, 1, 1}, + {"mixed-k", 1, 2, 2, 4, 5, 6, 2, 3, 4, 1, 0, 1, 1}, + {"fused7", 1, 1, 1, 8, 8, 8, 7, 7, 7, 1, 3, 3, 3}, + } + for _, tc := range cases { + in := pinInts(tc.n*tc.cIn*tc.dIn*tc.hIn*tc.wIn, 11) + ker := pinInts(tc.cOut*tc.cIn*tc.kD*tc.kH*tc.kW, 12) + ain, err := core.FromFloats(in, tc.n, tc.cIn, tc.dIn, tc.hIn, tc.wIn) + if err != nil { + t.Fatalf("%s input: %v", tc.name, err) + } + aker, err := core.FromFloats(ker, tc.cOut, tc.cIn, tc.kD, tc.kH, tc.kW) + if err != nil { + t.Fatalf("%s kernel: %v", tc.name, err) + } + out, oerr := Conv3D(ain, aker, nil, tc.stride, [3]int{tc.padD, tc.padH, tc.padW}, [3]int{1, 1, 1}) + if oerr != nil { + t.Fatalf("%s: %v", tc.name, oerr) + } + want := refConv3D(in, tc.n, tc.cIn, tc.dIn, tc.hIn, tc.wIn, ker, tc.cOut, tc.kD, tc.kH, tc.kW, nil, tc.stride, tc.padD, tc.padH, tc.padW) + pinConvEqual(t, tc.name, out.RawFloats(), want) + } +} + +func TestConvTranspose2DPhaseLatticeMatchesReference(t *testing.T) { + cases := []struct { + name string + n, cIn, cOut, hIn, wIn int + kH, kW int + stride, padding int + }{ + {"lattice2", 2, 2, 3, 5, 6, 3, 3, 2, 1}, + {"lattice3", 1, 2, 2, 4, 5, 2, 2, 3, 2}, + {"unit-stride", 2, 3, 2, 6, 6, 3, 3, 1, 1}, + {"tall-k", 1, 1, 2, 3, 3, 5, 4, 2, 0}, + } + for _, tc := range cases { + in := pinInts(tc.n*tc.cIn*tc.hIn*tc.wIn, 21) + ker := pinInts(tc.cIn*tc.cOut*tc.kH*tc.kW, 22) + bias := pinInts(tc.cOut, 23) + ain, err := core.FromFloats(in, tc.n, tc.cIn, tc.hIn, tc.wIn) + if err != nil { + t.Fatalf("%s input: %v", tc.name, err) + } + aker, err := core.FromFloats(ker, tc.cIn, tc.cOut, tc.kH, tc.kW) + if err != nil { + t.Fatalf("%s kernel: %v", tc.name, err) + } + abias, err := core.FromFloats(bias, tc.cOut) + if err != nil { + t.Fatalf("%s bias: %v", tc.name, err) + } + out, oerr := ConvTranspose2D(ain, aker, abias, tc.stride, tc.padding) + if oerr != nil { + t.Fatalf("%s: %v", tc.name, oerr) + } + want := refConvTranspose2D(in, tc.n, tc.cIn, tc.hIn, tc.wIn, ker, tc.cOut, tc.kH, tc.kW, bias, tc.stride, tc.padding) + pinConvEqual(t, tc.name, out.RawFloats(), want) + } +} + +func TestConvRowSpansFusedWindowAvoidsInvalidTaps(t *testing.T) { + // The fused interior takes its window bounds from spans[0].lo and + // spans[kW-1].hi, so an invalid tap must leave an empty window and + // the window must never cover a tap that reaches nothing: those + // are the two states the per-tap walk skips by construction. + for kW := 2; kW <= 8; kW++ { + for padW := 0; padW <= kW+1; padW++ { + for _, stride := range []int{1, 2, 3} { + for wIn := 1; wIn <= 9; wIn++ { + for wOut := 1; wOut <= 12; wOut++ { + spans := convRowSpans(kW, padW, stride, wIn, wOut) + fa, fb := spans[0].lo, spans[kW-1].hi + for kw := range kW { + if !spans[kw].valid && spans[kw].hi > spans[kw].lo { + t.Fatalf("kW=%d padW=%d stride=%d wIn=%d wOut=%d: tap %d invalid with window [%d, %d)", + kW, padW, stride, wIn, wOut, kw, spans[kw].lo, spans[kw].hi) + } + if fb > fa && !spans[kw].valid { + t.Fatalf("kW=%d padW=%d stride=%d wIn=%d wOut=%d: fused window [%d, %d) covers invalid tap %d", + kW, padW, stride, wIn, wOut, fa, fb, kw) + } + } + } + } + } + } + } +} diff --git a/signal/conv_pooling_test.go b/signal/conv_pooling_test.go new file mode 100644 index 0000000..eb4a957 --- /dev/null +++ b/signal/conv_pooling_test.go @@ -0,0 +1,481 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "math/rand/v2" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +func TestConv2DBadBias(t *testing.T) { + // Conv2D must reject a bias whose length does not match the output + // channels, like the other convolution entry points. + in := mustFromFloats(t, []float64{1, 2, 3, 4}, 1, 1, 2, 2) + ker := mustFromFloats(t, []float64{1, 0, 0, 1}, 1, 1, 2, 2) + bad := mustFromFloats(t, []float64{1, 2, 3}, 3) // one output channel, 3 biases + if _, err := Conv2D(in, ker, bad, 1, 0); err == nil { + t.Error("Conv2D: expected error for mismatched bias") + } + good := mustFromFloats(t, []float64{1}, 1) + if _, err := Conv2D(in, ker, good, 1, 0); err != nil { + t.Fatalf("Conv2D with valid bias: %v", err) + } +} + +func TestConv1D(t *testing.T) { + // 1×1×5 input, 1×1×3 kernel. + in := mustFromFloats(t, []float64{1, 2, 3, 4, 5}, 1, 1, 5) + ker := mustFromFloats(t, []float64{1, 0, -1}, 1, 1, 3) + got, err := Conv1D(in, ker, nil, 1, 0, 1) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 3 { + t.Errorf("Conv1D output L: %d, want 3", got.Shape()[2]) + } + // (1,2,3) convolved with (1,0,-1) = [1*1 + 2*0 + 3*(-1), 2*1 + 3*0 + 4*(-1), 3*1 + 4*0 + 5*(-1)] + // = [-2, -2, -2] + for i, w := range []float64{-2, -2, -2} { + v, _ := core.FloatAt(got, 0, 0, i) + if v != w { + t.Errorf("Conv1D [%d]: got %v, want %v", i, v, w) + } + } +} + +func TestConv3D(t *testing.T) { + // 1×1×1×2×2 input, 1×1×1×2×2 kernel (basic case). + in := mustFromFloats(t, []float64{1, 2, 3, 4}, 1, 1, 1, 2, 2) + ker := mustFromFloats(t, []float64{1, 1, 1, 1}, 1, 1, 1, 2, 2) + got, err := Conv3D(in, ker, nil, 1, [3]int{0, 0, 0}, [3]int{1, 1, 1}) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 1 || got.Shape()[1] != 1 || got.Shape()[2] != 1 || got.Shape()[3] != 1 || got.Shape()[4] != 1 { + t.Errorf("Conv3D shape: %v", got.Shape()) + } + // Sum of 1+2+3+4 = 10. + v, _ := core.FloatAt(got, 0, 0, 0, 0, 0) + if v != 10 { + t.Errorf("Conv3D: got %v, want 10", v) + } +} + +func TestConv2DGroups(t *testing.T) { + // 2 channels, groups=2 gives depthwise convolution. + in := mustFromFloats(t, []float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + }, 1, 2, 2, 2) + ker := mustFromFloats(t, []float64{ + 1, 0, // kernel for channel 0 + 0, 1, // kernel for channel 1 + }, 2, 1, 1, 2) + got, err := Conv2DGroups(in, ker, nil, 1, 0, 2) + if err != nil { + t.Fatal(err) + } + if got.Shape()[1] != 2 { + t.Errorf("Conv2DGroups C_out: %d, want 2", got.Shape()[1]) + } + // Channel 0: kernel [[1,0],[0,1]] applied to [[1,2],[3,4]] = [[1+4, 2+0],[3+0, 4+0]] = [[5,2],[3,4]] + // Wait: Conv2D with kH=1,kW=2 means 1-row × 2-col kernel. Output H=2, W=1. + // Channel 0 input [[1,2],[3,4]] conv with [[1,0]] (kH=1, kW=2) gives [[1·1+2·0, ...], [3·1+4·0, ...]] = [[1, ?], [3, ?]]. + v00, _ := core.FloatAt(got, 0, 0, 0, 0) + v10, _ := core.FloatAt(got, 0, 0, 1, 0) + // Channel 1: kernel [[0,1]] applied to [[5,6],[7,8]] = [[5·0+6·1, ...], [7·0+8·1, ...]] = [[6, ?], [8, ?]]. + v01, _ := core.FloatAt(got, 0, 1, 0, 0) + v11, _ := core.FloatAt(got, 0, 1, 1, 0) + if v00 != 1 || v10 != 3 || v01 != 6 || v11 != 8 { + t.Errorf("Conv2DGroups: got ch0 [[%v,?],[%v,?]], ch1 [[%v,?],[%v,?]], want [[1,_],[3,_]], [[6,_],[8,_]]", + v00, v10, v01, v11) + } +} + +func TestConvTranspose2D(t *testing.T) { + // 1×1×1×1 input, 1×1×2×2 kernel: output should be 2×2. + in := mustFromFloats(t, []float64{1}, 1, 1, 1, 1) + ker := mustFromFloats(t, []float64{ + 1, 2, + 3, 4, + }, 1, 1, 2, 2) + got, err := ConvTranspose2D(in, ker, nil, 1, 0) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 2 || got.Shape()[3] != 2 { + t.Errorf("ConvTranspose2D shape: %v", got.Shape()) + } + // With single input and identity-position kernel, output == kernel. + for i, w := range []float64{1, 2, 3, 4} { + v, _ := core.FloatAt(got, 0, 0, i/2, i%2) + if math.Abs(v-w) > 1e-9 { + t.Errorf("ConvTranspose2D [%d]: got %v, want %v", i, v, w) + } + } +} + +func TestMaxPool1D(t *testing.T) { + in := mustFromFloats(t, []float64{1, 3, 2, 4, 1}, 1, 1, 5) + got, err := MaxPool1D(in, 2, 2, 0) + if err != nil { + t.Fatal(err) + } + // 2 windows: [1,3] gives 3, [2,4] gives 4. + v0, _ := core.FloatAt(got, 0, 0, 0) + v1, _ := core.FloatAt(got, 0, 0, 1) + if v0 != 3 || v1 != 4 { + t.Errorf("MaxPool1D: got %v %v, want 3 4", v0, v1) + } +} + +func TestAvgPool3D(t *testing.T) { + in := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 1, 1, 2, 2, 2) + got, err := AvgPool3D(in, [3]int{2, 2, 2}, [3]int{2, 2, 2}, [3]int{0, 0, 0}, true) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 1 || got.Shape()[3] != 1 || got.Shape()[4] != 1 { + t.Errorf("AvgPool3D shape: %v", got.Shape()) + } + v, _ := core.FloatAt(got, 0, 0, 0, 0, 0) + // Average of all 8 = 36 / 8 = 4.5. + if math.Abs(v-4.5) > 1e-9 { + t.Errorf("AvgPool3D: got %v, want 4.5", v) + } +} + +func TestAdaptiveMaxPool2D(t *testing.T) { + in := mustFromFloats(t, []float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + 13, 14, 15, 16, + }, 1, 1, 4, 4) + got, err := AdaptiveMaxPool2D(in, 2, 2) + if err != nil { + t.Fatal(err) + } + v00, _ := core.FloatAt(got, 0, 0, 0, 0) + v01, _ := core.FloatAt(got, 0, 0, 0, 1) + v11, _ := core.FloatAt(got, 0, 0, 1, 1) + if v00 != 6 || v01 != 8 || v11 != 16 { + t.Errorf("AdaptiveMaxPool2D: got %v %v %v, want 6 8 16", v00, v01, v11) + } +} + +func TestGlobalAvgPool2D(t *testing.T) { + in := mustFromFloats(t, []float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + }, 1, 1, 2, 4) + got, err := GlobalAvgPool2D(in) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0, 0, 0, 0) + // Average of 1..8 = 36/8 = 4.5. + if math.Abs(v-4.5) > 1e-9 { + t.Errorf("GlobalAvgPool2D: got %v, want 4.5", v) + } +} + +func TestGlobalMaxPool2D(t *testing.T) { + in := mustFromFloats(t, []float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + }, 1, 1, 2, 4) + got, err := GlobalMaxPool2D(in) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0, 0, 0, 0) + if v != 8 { + t.Errorf("GlobalMaxPool2D: got %v, want 8", v) + } +} + +func TestGlobalAvgPool1D(t *testing.T) { + in := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 1, 1, 6) + got, err := GlobalAvgPool1D(in) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 1 { + t.Errorf("GlobalAvgPool1D L_out: %d", got.Shape()[2]) + } + v, _ := core.FloatAt(got, 0, 0, 0) + if math.Abs(v-3.5) > 1e-9 { + t.Errorf("GlobalAvgPool1D: got %v, want 3.5", v) + } +} + +func TestGlobalMaxPool1D(t *testing.T) { + in := mustFromFloats(t, []float64{1, 5, 3, 4, 2, 6}, 1, 1, 6) + got, err := GlobalMaxPool1D(in) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0, 0, 0) + if v != 6 { + t.Errorf("GlobalMaxPool1D: got %v, want 6", v) + } +} + +func TestGlobalAvgPool3D(t *testing.T) { + in := mustFromFloats(t, []float64{ + 1, 2, 3, 4, 5, 6, 7, 8, + }, 1, 1, 2, 2, 2) + got, err := GlobalAvgPool3D(in) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0, 0, 0, 0, 0) + // Average of 1..8 = 4.5. + if math.Abs(v-4.5) > 1e-9 { + t.Errorf("GlobalAvgPool3D: got %v, want 4.5", v) + } +} + +func TestGlobalMaxPool3D(t *testing.T) { + in := mustFromFloats(t, []float64{ + 1, 2, 3, 4, 5, 6, 7, 8, + }, 1, 1, 2, 2, 2) + got, err := GlobalMaxPool3D(in) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0, 0, 0, 0, 0) + if v != 8 { + t.Errorf("GlobalMaxPool3D: got %v, want 8", v) + } +} + +func TestAdaptiveMaxPool1D(t *testing.T) { + in := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 1, 1, 6) + got, err := AdaptiveMaxPool1D(in, 2) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 2 { + t.Errorf("AdaptiveMaxPool1D L_out: %d", got.Shape()[2]) + } + // Window [0..3) gives max(1,2,3) = 3; [3..6) gives max(4,5,6) = 6. + v0, _ := core.FloatAt(got, 0, 0, 0) + v1, _ := core.FloatAt(got, 0, 0, 1) + if v0 != 3 || v1 != 6 { + t.Errorf("AdaptiveMaxPool1D: got %v %v, want 3 6", v0, v1) + } +} + +func TestAdaptiveAvgPool1D(t *testing.T) { + in := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 1, 1, 6) + got, err := AdaptiveAvgPool1D(in, 2) + if err != nil { + t.Fatal(err) + } + // [1..4) gives 6/3 = 2; [4..7) gives 15/3 = 5. + v0, _ := core.FloatAt(got, 0, 0, 0) + v1, _ := core.FloatAt(got, 0, 0, 1) + if math.Abs(v0-2) > 1e-9 || math.Abs(v1-5) > 1e-9 { + t.Errorf("AdaptiveAvgPool1D: got %v %v, want 2 5", v0, v1) + } +} + +// TestAdaptivePoolCeilWindows pins the floor-start, ceil-end window +// convention on sizes that do not divide: window o covers +// floor(o·in/out) to ceil((o+1)·in/out), so every input sample lands +// in some window (the old floor ends silently dropped input samples). +func TestAdaptivePoolCeilWindows(t *testing.T) { + in := mustFromFloats(t, []float64{1, 2, 3, 4, 5}, 1, 1, 5) + maxGot, err := AdaptiveMaxPool1D(in, 2) + if err != nil { + t.Fatal(err) + } + // Windows [0,3) and [2,5): max 3 and 5 (the floor convention gave 2). + v0, _ := core.FloatAt(maxGot, 0, 0, 0) + v1, _ := core.FloatAt(maxGot, 0, 0, 1) + if v0 != 3 || v1 != 5 { + t.Errorf("AdaptiveMaxPool1D 5 to 2: got %v %v, want 3 5", v0, v1) + } + avgGot, err := AdaptiveAvgPool1D(in, 2) + if err != nil { + t.Fatal(err) + } + // (1+2+3)/3 and (3+4+5)/3. + a0, _ := core.FloatAt(avgGot, 0, 0, 0) + a1, _ := core.FloatAt(avgGot, 0, 0, 1) + if math.Abs(a0-2) > 1e-9 || math.Abs(a1-4) > 1e-9 { + t.Errorf("AdaptiveAvgPool1D 5 to 2: got %v %v, want 2 4", a0, a1) + } + // Upsampling repeats inputs: 2 to 4 gives [0,1),[0,1),[1,2),[1,2). + two := mustFromFloats(t, []float64{7, 9}, 1, 1, 2) + up, err := AdaptiveMaxPool1D(two, 4) + if err != nil { + t.Fatal(err) + } + wantUp := []float64{7, 7, 9, 9} + for i, w := range wantUp { + if v, _ := core.FloatAt(up, 0, 0, i); v != w { + t.Errorf("AdaptiveMaxPool1D 2 to 4 [%d]: got %v, want %v", i, v, w) + } + } + // The 2-D twins follow the same ceil ends. + sq := mustFromFloats(t, []float64{ + 1, 2, 3, 4, 5, + 6, 7, 8, 9, 10, + 11, 12, 13, 14, 15, + 16, 17, 18, 19, 20, + 21, 22, 23, 24, 25, + }, 1, 1, 5, 5) + max2D, err := AdaptiveMaxPool2D(sq, 2, 2) + if err != nil { + t.Fatal(err) + } + // Windows are rows/cols [0,3) and [2,5): the quadrant maxima are + // 13, 15, 23, 25. + want2D := []float64{13, 15, 23, 25} + for i, w := range want2D { + oy, ox := i/2, i%2 + if v, _ := core.FloatAt(max2D, 0, 0, oy, ox); v != w { + t.Errorf("AdaptiveMaxPool2D 5 to 2 [%d,%d]: got %v, want %v", oy, ox, v, w) + } + } + avg2D, err := AdaptiveAvgPool2D(sq, 2, 2) + if err != nil { + t.Fatal(err) + } + // Top-left window average (1+2+3+6+7+8+11+12+13)/9 = 7. + if v, _ := core.FloatAt(avg2D, 0, 0, 0, 0); math.Abs(v-7) > 1e-9 { + t.Errorf("AdaptiveAvgPool2D 5 to 2 [0,0]: got %v, want 7", v) + } +} + +// TestPoolAndConvNumeratorGuard pins the output-size contract: a +// kernel that exceeds the padded input must be an error, not a +// 1-wide output from a negative numerator truncating toward zero. +func TestPoolAndConvNumeratorGuard(t *testing.T) { + in4 := mustFromFloats(t, make([]float64, 16), 1, 1, 4, 4) + if _, err := MaxPool2D(in4, 5, 2, 0); err == nil { + t.Error("MaxPool2D: expected an error when the kernel exceeds the input") + } + if _, err := AvgPool2D(in4, 5, 2, 0, true); err == nil { + t.Error("AvgPool2D: expected an error when the kernel exceeds the input") + } + lane := mustFromFloats(t, make([]float64, 4), 1, 1, 4) + if _, err := MaxPool1D(lane, 5, 2, 0); err == nil { + t.Error("MaxPool1D: expected an error when the kernel exceeds the input") + } + cube := mustFromFloats(t, make([]float64, 8), 1, 1, 2, 2, 2) + if _, err := MaxPool3D(cube, [3]int{5, 1, 1}, [3]int{2, 1, 1}, [3]int{0, 0, 0}); err == nil { + t.Error("MaxPool3D: expected an error when the kernel exceeds the input") + } + kernel, _ := core.FromFloats(make([]float64, 25), 1, 1, 5, 5) + if _, err := Conv2D(in4, kernel, nil, 2, 0); err == nil { + t.Error("Conv2D: expected an error when the kernel exceeds the input") + } + k1, _ := core.FromFloats(make([]float64, 5), 1, 1, 5) + if _, err := Conv1D(lane, k1, nil, 2, 0, 1); err == nil { + t.Error("Conv1D: expected an error when the kernel exceeds the input") + } + k3, _ := core.FromFloats(make([]float64, 125), 1, 1, 5, 5, 5) + if _, err := Conv3D(cube, k3, nil, 2, [3]int{0, 0, 0}, [3]int{1, 1, 1}); err == nil { + t.Error("Conv3D: expected an error when the kernel exceeds the input") + } +} + +// TestConv2DGroupsPartialChannelBlock pins the 2-D kernel's +// output-channel blocking when the last block of a group is partial: +// six output channels in two groups of three, on a row wide enough that +// one block covers two channels, so the group's third channel is a +// block of its own. Reading the block count as a floor instead of a +// ceiling drops that trailing block and leaves the channel at the zero +// the output array was initialised with, which the scalar walk below +// sees at once. +func TestConv2DGroupsPartialChannelBlock(t *testing.T) { + const ( + n, cIn, hIn, wIn = 1, 2, 3, 2000 + groups = 2 + cOut, kH, kW = 6, 2, 3 + stride, padding = 1, 1 + ) + cInPerGroup := cIn / groups + cOutPerGroup := cOut / groups + if block := convAccCap / wIn; block < 1 || block >= cOutPerGroup { + t.Fatalf("the case needs a partial block: a block of %d against %d channels per group", block, cOutPerGroup) + } + rng := rand.New(rand.NewPCG(11, 13)) + in := make([]float64, n*cIn*hIn*wIn) + for i := range in { + in[i] = rng.NormFloat64() + } + ker := make([]float64, cOut*cInPerGroup*kH*kW) + for i := range ker { + ker[i] = rng.NormFloat64() + } + bias := make([]float64, cOut) + for i := range bias { + bias[i] = rng.NormFloat64() + } + x := mustFromFloats(t, in, n, cIn, hIn, wIn) + k := mustFromFloats(t, ker, cOut, cInPerGroup, kH, kW) + b := mustFromFloats(t, bias, cOut) + got, err := Conv2DGroups(x, k, b, stride, padding, groups) + if err != nil { + t.Fatalf("Conv2DGroups: %v", err) + } + hOut := (hIn+2*padding-kH)/stride + 1 + wOut := (wIn+2*padding-kW)/stride + 1 + if got.Shape()[1] != cOut || got.Shape()[2] != hOut || got.Shape()[3] != wOut { + t.Fatalf("shape %v, want [%d %d %d %d]", got.Shape(), n, cOut, hOut, wOut) + } + // A scalar walk in the same (input channel, tap row, tap column) + // order the kernel accumulates its taps in. + worst, scale := 0.0, 0.0 + for oc := range cOut { + g := oc / cOutPerGroup + live := false + for oh := range hOut { + for ow := range wOut { + want := 0.0 + for ic := range cInPerGroup { + for kh := range kH { + ih := oh*stride + kh - padding + if ih < 0 || ih >= hIn { + continue + } + for kw := range kW { + iw := ow*stride + kw - padding + if iw < 0 || iw >= wIn { + continue + } + want += in[((g*cInPerGroup+ic)*hIn+ih)*wIn+iw] * + ker[((oc*cInPerGroup+ic)*kH+kh)*kW+kw] + } + } + } + want += bias[oc] + gv := got.FloatAt(oc*hOut*wOut + oh*wOut + ow) + if gv != 0 { + live = true + } + if d := math.Abs(gv - want); d > worst { + worst = d + } + if a := math.Abs(want); a > scale { + scale = a + } + } + } + if !live { + t.Errorf("output channel %d (group %d, the group's last channel) is identically zero", oc, g) + } + } + if worst > 1e-12*scale { + t.Fatalf("worst deviation from the scalar walk %.6g (scale %.6g), want the blocked sweep to agree with it", worst, scale) + } +} diff --git a/signal/conv_transpose_test.go b/signal/conv_transpose_test.go new file mode 100644 index 0000000..0a296c5 --- /dev/null +++ b/signal/conv_transpose_test.go @@ -0,0 +1,120 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" +) + +// convTranspose2DReference computes the transpose convolution straight +// from its scatter definition, each output element summing the input +// positions that map onto it in the (input channel, kernel row, +// kernel column) order, the bias added last. It is the test's own +// arithmetic, not a second kernel: the loops mirror the definition +// the documentation states. +func convTranspose2DReference(in []float64, k []float64, bias []float64, + n, cIn, hIn, wIn, cOut, kH, kW, stride, padding int, +) []float64 { + hOut := (hIn-1)*stride - 2*padding + kH + wOut := (wIn-1)*stride - 2*padding + kW + out := make([]float64, n*cOut*hOut*wOut) + for b := range n { + for oc := range cOut { + for oh := range hOut { + for ow := range wOut { + sum := 0.0 + for ic := range cIn { + for kh := range kH { + ih := oh + padding - kh + if ih < 0 || ih%stride != 0 { + continue + } + ih /= stride + if ih >= hIn { + continue + } + for kw := range kW { + iw := ow + padding - kw + if iw < 0 || iw%stride != 0 { + continue + } + iw /= stride + if iw >= wIn { + continue + } + sum += in[((b*cIn+ic)*hIn+ih)*wIn+iw] * k[((ic*cOut+oc)*kH+kh)*kW+kw] + } + } + } + if bias != nil { + sum += bias[oc] + } + out[((b*cOut+oc)*hOut+oh)*wOut+ow] = sum + } + } + } + } + return out +} + +// TestConvTranspose2DChannelBlocking pins the channel-blocked kernel +// against the scatter definition on shapes the single-channel test +// cannot reach: several input and output channels (so the kernel's +// channel stride matters), a stride above one, a bias, and an output +// wider than one channel block, which leaves a partial block at the +// end. +func TestConvTranspose2DChannelBlocking(t *testing.T) { + seed := 1 + rand := func() float64 { + seed = (7*seed + 3) % 97 + return float64(seed)/97 - 0.5 + } + fill := func(n int) []float64 { + v := make([]float64, n) + for i := range v { + v[i] = rand() + } + return v + } + + t.Run("several channels, stride two, bias", func(t *testing.T) { + const n, cIn, cOut, k = 2, 2, 3, 2 + const hIn, wIn, stride, padding = 3, 3, 2, 1 + in := fill(n * cIn * hIn * wIn) + kk := fill(cIn * cOut * k * k) + bias := fill(cOut) + got, err := ConvTranspose2D(mustFromFloats(t, in, n, cIn, hIn, wIn), mustFromFloats(t, kk, cIn, cOut, k, k), mustFromFloats(t, bias, cOut), stride, padding) + if err != nil { + t.Fatalf("ConvTranspose2D: %v", err) + } + want := convTranspose2DReference(in, kk, bias, n, cIn, hIn, wIn, cOut, k, k, stride, padding) + outF := got.RawFloats() + for i := range want { + if math.Float64bits(outF[i]) != math.Float64bits(want[i]) { + t.Fatalf("element %d = %v, want %v", i, outF[i], want[i]) + } + } + }) + t.Run("partial channel block", func(t *testing.T) { + // An output row of 1025 caps one block at three channels, so + // four output channels run as a full block and a partial one. + const n, cIn, cOut, k = 1, 1, 4, 3 + const hIn, wIn, stride, padding = 1025, 1025, 1, 1 + in := fill(n * cIn * hIn * wIn) + kk := fill(cIn * cOut * k * k) + bias := fill(cOut) + got, err := ConvTranspose2D(mustFromFloats(t, in, n, cIn, hIn, wIn), mustFromFloats(t, kk, cIn, cOut, k, k), mustFromFloats(t, bias, cOut), stride, padding) + if err != nil { + t.Fatalf("ConvTranspose2D: %v", err) + } + want := convTranspose2DReference(in, kk, bias, n, cIn, hIn, wIn, cOut, k, k, stride, padding) + outF := got.RawFloats() + for i := range want { + if math.Float64bits(outF[i]) != math.Float64bits(want[i]) { + t.Fatalf("element %d = %v, want %v", i, outF[i], want[i]) + } + } + }) +} diff --git a/signal/conv_window_pins_test.go b/signal/conv_window_pins_test.go new file mode 100644 index 0000000..e72caf5 --- /dev/null +++ b/signal/conv_window_pins_test.go @@ -0,0 +1,801 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Pins for the convolution fast-path window arithmetic, the Decimate +// output-length rule, the float32 and int dtype paths of +// Resample/Decimate/IDWT and the SolvePoissonNeumann solve. + +package signal + +import ( + "fmt" + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestConv1DStride1TapWindowPastBlock drives a kernel tap whose whole +// window lies outside an output block at unit stride: the stride == 1 +// fast path must skip the tap instead of slicing a range whose high +// bound is under its low one (the finding's "slice bounds out of +// range" panic). The result is pinned against the direct convolution +// definition, so the skipped taps are also checked to contribute +// nothing. +func TestConv1DStride1TapWindowPastBlock(t *testing.T) { + const ( + n = 202 + factor = 3 + padding = 200 + ) + // 202 samples, kernel [1, 2, 3], padding 200: lOut = 600 spans two + // 512-wide blocks, and block 1 sits past the last tap's reach. + vals := make([]float64, n) + for i := range vals { + vals[i] = float64(i%7) - 3 + } + x := mustFloats(t, vals, 1, 1, n) + k := mustFloats(t, []float64{1, 2, 3}, 1, 1, factor) + out, err := Conv1D(x, k, nil, 1, padding, 1) + if err != nil { + t.Fatalf("Conv1D: %v", err) + } + if got := out.Shape(); got[0] != 1 || got[1] != 1 || got[2] != 600 { + t.Fatalf("Conv1D: shape %v, want [1 1 600]", got) + } + for ol := range 600 { + want := 0.0 + for kl := range factor { + idx := ol + kl - padding + if idx < 0 || idx >= n { + continue + } + want += vals[idx] * float64(kl+1) + } + if got := out.FloatAt(ol); got != want { + t.Fatalf("Conv1D: out[%d] = %g, want %g (the direct definition)", ol, got, want) + } + } +} + +// TestConv2DStride1PaddingOnlyTapWindow drives a 2-D kernel tap whose +// window lies entirely in the padding: the stride == 1 fast path must +// skip it rather than slice rank 1 into the tap row. +func TestConv2DStride1PaddingOnlyTapWindow(t *testing.T) { + x := mustFloats(t, []float64{1}, 1, 1, 1, 1) + k := mustFloats(t, []float64{1, 2, 3, 4, 5}, 1, 1, 1, 5) + out, err := Conv2D(x, k, nil, 1, 2) + if err != nil { + t.Fatalf("Conv2D: %v", err) + } + if got := out.Shape(); got[0] != 1 || got[1] != 1 || got[2] != 5 || got[3] != 1 { + t.Fatalf("Conv2D: shape %v, want [1 1 5 1]", got) + } + // Only the kernel's centre tap overlaps the single input sample; + // the taps at kw 0, 1, 3 and 4 reach no output at all. + for oh := range 5 { + want := 0.0 + if oh == 2 { + want = 3 // x[0] * k[2] + } + if got := out.FloatAt(oh); got != want { + t.Fatalf("Conv2D: out[0,0,%d,0] = %g, want %g", oh, got, want) + } + } +} + +// TestConv3DStride1PaddingOnlyTapWindow drives the 3-D twin of the +// same padding-only tap window. +func TestConv3DStride1PaddingOnlyTapWindow(t *testing.T) { + x := mustFloats(t, []float64{1}, 1, 1, 1, 1, 1) + k := mustFloats(t, []float64{1, 2, 3, 4, 5}, 1, 1, 1, 1, 5) + out, err := Conv3D(x, k, nil, 1, [3]int{0, 0, 2}, [3]int{1, 1, 1}) + if err != nil { + t.Fatalf("Conv3D: %v", err) + } + if got := out.Shape(); got[0] != 1 || got[1] != 1 || got[2] != 1 || got[3] != 1 || got[4] != 1 { + t.Fatalf("Conv3D: shape %v, want [1 1 1 1 1]", got) + } + if got := out.FloatAt(0); got != 3 { + t.Fatalf("Conv3D: out[0] = %g, want 3 (x[0] * k[2])", got) + } +} + +// TestDecimateOutLenOutOfRange calls Decimate with a custom tap count +// whose filter delay plus first kept index runs to or past the end of +// the filtered signal: the truncated division used to round a negative +// numerator toward zero, producing one output sample read out of +// range. +func TestDecimateOutLenOutOfRange(t *testing.T) { + cases := []struct { + name string + n, factor, taps int + }{ + {"ReportedSmallFactor", 26, 8, 25}, + {"FactorLargerThanN", 10, 100, 3}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + vals := make([]float64, tc.n) + for i := range vals { + vals[i] = math.Sin(0.3 * float64(i)) + } + x := mustFloats(t, vals, tc.n) + out, err := Decimate(x, tc.factor, tc.taps) + if err != nil { + t.Fatalf("Decimate(n=%d, factor=%d, taps=%d): %v", tc.n, tc.factor, tc.taps, err) + } + want := decimateOutLen(tc.n, tc.factor, tc.taps) + if out.Len() != want { + t.Fatalf("Decimate(n=%d, factor=%d, taps=%d) = %d samples, want %d", + tc.n, tc.factor, tc.taps, out.Len(), want) + } + t.Logf("Decimate(n=%d, factor=%d, taps=%d) -> %d samples", tc.n, tc.factor, tc.taps, out.Len()) + }) + } +} + +// TestResampleDecimateFloat32Input feeds float32 and int series into +// Resample and Decimate: both read the payload without a dtype guard, +// so a non-float input panicked on the nil float64 payload instead of +// widening through FloatAt like DWT and the rest of the package. The +// results are pinned to the float64 run on the same values. +func TestResampleDecimateFloat32Input(t *testing.T) { + const n = 256 + f64 := make([]float64, n) + f32 := make([]float32, n) + for i := range n { + f32[i] = float32(math.Cos(2 * math.Pi * 3 * float64(i) / float64(n))) + // The float64 twin holds the widened float32 values, so the + // two runs see bit-identical inputs. + f64[i] = float64(f32[i]) + } + a32 := mustFloat32s(t, f32, n) + a64 := mustFloats(t, f64, n) + + t.Run("ResampleFloat32", func(t *testing.T) { + got, err := Resample(a32, 3, 1, 0) + if err != nil { + t.Fatalf("Resample(float32): %v", err) + } + want, err := Resample(a64, 3, 1, 0) + if err != nil { + t.Fatalf("Resample(float64): %v", err) + } + if got.Len() != want.Len() { + t.Fatalf("Resample(float32) = %d samples, want %d", got.Len(), want.Len()) + } + for i := range want.Len() { + // FloatAt widens exactly, so both runs must agree bit + // for bit on the same input values. + if g, w := got.FloatAt(i), want.FloatAt(i); g != w { + t.Fatalf("Resample sample %d = %g, want %g", i, g, w) + } + } + }) + + t.Run("DecimateFloat32", func(t *testing.T) { + got, err := Decimate(a32, 2, 0) + if err != nil { + t.Fatalf("Decimate(float32): %v", err) + } + want, err := Decimate(a64, 2, 0) + if err != nil { + t.Fatalf("Decimate(float64): %v", err) + } + if got.Len() != want.Len() { + t.Fatalf("Decimate(float32) = %d samples, want %d", got.Len(), want.Len()) + } + // FilterApply keeps the float32 dtype, so its output samples + // are the float64 ones rounded to float32 on store; the + // decimation picks the same ones. The comparison is exact. + for i := range want.Len() { + w := float64(float32(want.FloatAt(i))) + if g := got.FloatAt(i); g != w { + t.Fatalf("Decimate sample %d = %g, want %g", i, g, w) + } + } + }) + + t.Run("Int", func(t *testing.T) { + ivals := make([]int64, n) + fvals := make([]float64, n) + for i := range n { + ivals[i] = int64(i%7) - 3 + fvals[i] = float64(ivals[i]) + } + ai, err := core.FromInts(ivals, n) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + af := mustFloats(t, fvals, n) + gi, err := Resample(ai, 2, 1, 0) + if err != nil { + t.Fatalf("Resample(int): %v", err) + } + gf, err := Resample(af, 2, 1, 0) + if err != nil { + t.Fatalf("Resample(float64): %v", err) + } + for i := range gf.Len() { + if g, w := gi.FloatAt(i), gf.FloatAt(i); g != w { + t.Fatalf("Resample int sample %d = %g, want %g", i, g, w) + } + } + di, err := Decimate(ai, 4, 0) + if err != nil { + t.Fatalf("Decimate(int): %v", err) + } + df, err := Decimate(af, 4, 0) + if err != nil { + t.Fatalf("Decimate(float64): %v", err) + } + for i := range df.Len() { + if g, w := di.FloatAt(i), df.FloatAt(i); g != w { + t.Fatalf("Decimate int sample %d = %g, want %g", i, g, w) + } + } + }) +} + +// TestIDWTFloat32Coefficients inverts a float32 coefficient array: +// DWT accepts float32 through FloatAt, and IDWT reading the nil +// float64 payload instead returned all zeros with no error. The +// float32 and float64 runs must agree sample for sample. +func TestIDWTFloat32Coefficients(t *testing.T) { + // A genuine packed coefficient array, so the assertion covers a + // realistic layout: the DWT of a ramp at one level. + ramp := make([]float64, 8) + for i := range ramp { + ramp[i] = float64(i + 1) + } + coef, err := DWT(mustFloats(t, ramp, 8), 1) + if err != nil { + t.Fatalf("DWT: %v", err) + } + f32 := make([]float32, 8) + f64 := make([]float64, 8) + for i := range 8 { + v := coef.FloatAt(i) + f32[i] = float32(v) + f64[i] = float64(f32[i]) + } + got, err := IDWT(mustFloat32s(t, f32, 8), 1) + if err != nil { + t.Fatalf("IDWT(float32): %v", err) + } + want, err := IDWT(mustFloats(t, f64, 8), 1) + if err != nil { + t.Fatalf("IDWT(float64): %v", err) + } + nonzero := 0 + for i := range 8 { + g, w := got.FloatAt(i), want.FloatAt(i) + if g != w { + t.Fatalf("IDWT sample %d = %g, want %g", i, g, w) + } + if g != 0 { + nonzero++ + } + } + if nonzero == 0 { + t.Fatal("IDWT returned all zeros for a float32 coefficient array") + } +} + +// mustFloat32s builds a float32 array, failing the test on a bad shape. +func mustFloat32s(t *testing.T, vals []float32, shape ...int) *core.Array { + t.Helper() + a, err := core.FromFloat32s(vals, shape...) + if err != nil { + t.Fatalf("FromFloat32s(%v, %v): %v", vals, shape, err) + } + return a +} + +// decimateOutLen mirrors the documented output rule for the tap count +// the caller passes: the filter delay plus the first kept sample must +// leave whole factors behind it. It is what the regression test +// compares the library against, written from the rule rather than from +// the implementation. +func decimateOutLen(n, factor, taps int) int { + if taps%2 == 0 { + taps++ + } + if taps >= n { + return 0 + } + delay := (taps - 1) / 2 + // A kept sample needs m + delay <= n - 1 for some multiple m of + // factor with m >= delay, i.e. floor((n-1-delay)/factor) - skip + 1 + // samples exist below the end of the signal. + skip := (delay + factor - 1) / factor + lastMultiple := (n - 1 - delay) / factor + if lastMultiple < skip { + return 0 + } + return lastMultiple - skip + 1 +} + +// assertFinite fails the test on the first non-finite sample: a +// maximum-error loop cannot see a NaN, because every comparison +// against one is false, which is how the shipped manufactured test +// passed on an all-NaN solve. Every solve assertion below runs after +// this check. +func assertFinite(t *testing.T, what string, vals []float64) { + t.Helper() + for i, v := range vals { + if math.IsNaN(v) || math.IsInf(v, 0) { + t.Fatalf("%s: sample %d is %g, the result must be finite", what, i, v) + } + } +} + +// stencilTap is one entry of a mirrored-stencil operator row: the +// column index and the coefficient in units of 1/h². +type stencilTap struct { + j int + v float64 +} + +// mirroredStencilRows returns the taps of the ghost-mirror operator's +// row i of an N-point axis: [2,-2] on the boundary, where the mirrored +// neighbour outside the interval stands in for zero slope, and +// [-1,2,-1] inside. No code in the library shares these taps; they are +// the definition the regression tests hold the solve against. +func mirroredStencilRows(n, i int) []stencilTap { + switch { + case i == 0: + return []stencilTap{{0, 2}, {1, -2}} + case i == n-1: + return []stencilTap{{n - 2, -2}, {n - 1, 2}} + default: + return []stencilTap{{i - 1, -1}, {i, 2}, {i + 1, -1}} + } +} + +// mirroredStencilApply applies the 2-D mirrored operator to a +// row-major grid. +func mirroredStencilApply(u []float64, rows, cols int, hx, hy float64) []float64 { + out := make([]float64, rows*cols) + for r := range rows { + for c := range cols { + s := 0.0 + for _, e := range mirroredStencilRows(cols, c) { + s += e.v * u[r*cols+e.j] / (hx * hx) + } + for _, e := range mirroredStencilRows(rows, r) { + s += e.v * u[e.j*cols+c] / (hy * hy) + } + out[r*cols+c] = s + } + } + return out +} + +// mirroredStencilMatrix builds the explicit operator matrix of the +// (rows·cols)-point grid: column j is the operator applied to a unit +// input at j, so row i holds the coefficients of equation i. +func mirroredStencilMatrix(rows, cols int, hx, hy float64) [][]float64 { + n := rows * cols + m := make([][]float64, n) + for i := range n { + m[i] = make([]float64, n) + } + for j := range n { + u := make([]float64, n) + u[j] = 1 + col := mirroredStencilApply(u, rows, cols, hx, hy) + for i := range n { + m[i][j] = col[i] + } + } + return m +} + +// neumannTrapz returns the trapezoidal-weighted sum of a row-major +// grid and the measure's total weight: the boundary rows and columns +// count half. That weight vector is the mirrored operator's left null +// vector, the measure its compatibility condition is stated in. +func neumannTrapz(u []float64, rows, cols int) (sum, weight float64) { + for r := range rows { + wr := 1.0 + if r == 0 || r == rows-1 { + wr = 0.5 + } + for c := range cols { + wc := 1.0 + if c == 0 || c == cols-1 { + wc = 0.5 + } + sum += wr * wc * u[r*cols+c] + weight += wr * wc + } + } + return sum, weight +} + +// neumannTrapzSum returns the trapezoidal-weighted sum of a row-major +// grid, the gate's measure. +func neumannTrapzSum(u []float64, rows, cols int) float64 { + sum, _ := neumannTrapz(u, rows, cols) + return sum +} + +// neumannTrapzMean returns the trapezoidal-weighted mean, the sum over +// its total weight. +func neumannTrapzMean(u []float64, rows, cols int) float64 { + sum, weight := neumannTrapz(u, rows, cols) + return sum / weight +} + +// neumannDenseSolve solves the singular mirrored system M u = f by +// replacing the last equation with the zero-trapezoidal-mean +// constraint, which makes it nonsingular (the trapezoidal weight is +// the left null vector, so it cannot lie in the row space), then +// eliminating with partial pivoting. Test scaffolding: it panics on a +// singular matrix rather than returning a wrong answer. +func neumannDenseSolve(m [][]float64, f []float64, rows, cols int) []float64 { + n := len(f) + a := make([][]float64, n) + for i := range n { + a[i] = append(append([]float64(nil), m[i]...), f[i]) + } + for j := range n { + r, c := j/cols, j%cols + wr, wc := 1.0, 1.0 + if r == 0 || r == rows-1 { + wr = 0.5 + } + if c == 0 || c == cols-1 { + wc = 0.5 + } + a[n-1][j] = wr * wc + } + a[n-1][n] = 0 + for col := range n { + piv := col + for r := col + 1; r < n; r++ { + if math.Abs(a[r][col]) > math.Abs(a[piv][col]) { + piv = r + } + } + a[col], a[piv] = a[piv], a[col] + p := a[col][col] + if p == 0 { + panic("neumannDenseSolve: singular matrix") + } + for r := col + 1; r < n; r++ { + fac := a[r][col] / p + for c := col; c <= n; c++ { + a[r][c] -= fac * a[col][c] + } + } + } + u := make([]float64, n) + for i := n - 1; i >= 0; i-- { + s := a[i][n] + for j := i + 1; j < n; j++ { + s -= a[i][j] * u[j] + } + u[i] = s / a[i][i] + } + return u +} + +// TestSolvePoissonNeumannMirroredStencil holds the solve against the +// explicit ghost-mirror matrix: the operator is built entry by entry +// for the small grids, inverted through a dense elimination with the +// zero-trapezoidal-mean constraint, and the library's solution must +// agree with it, satisfy M u = f and carry no trapezoidal mean. +func TestSolvePoissonNeumannMirroredStencil(t *testing.T) { + for _, g := range []struct { + rows, cols int + lx, ly float64 + }{ + {3, 3, 1, 1}, + {4, 4, 1, 1}, + {5, 3, 1.7, 0.9}, + {9, 9, 1, 1}, + } { + hx, hy := g.lx/float64(g.cols-1), g.ly/float64(g.rows-1) + // A compatible source: a cosine mode minus its trapezoidal + // mean, so the constant mode of the transform vanishes. + f := make([]float64, g.rows*g.cols) + for r := range g.rows { + for c := range g.cols { + f[r*g.cols+c] = math.Cos(2*math.Pi*float64(c)/float64(g.cols-1)) * + math.Cos(3*math.Pi*float64(r)/float64(g.rows-1)) + } + } + mu := neumannTrapzMean(f, g.rows, g.cols) + for i := range f { + f[i] -= mu + } + want := neumannDenseSolve(mirroredStencilMatrix(g.rows, g.cols, hx, hy), f, g.rows, g.cols) + got, err := SolvePoissonNeumann(mustFloats(t, f, g.rows, g.cols), g.lx, g.ly) + if err != nil { + t.Fatalf("%dx%d: SolvePoissonNeumann: %v", g.rows, g.cols, err) + } + vals := make([]float64, got.Len()) + for i := range vals { + vals[i] = got.FloatAt(i) + } + assertFinite(t, fmt.Sprintf("%dx%d solve", g.rows, g.cols), vals) + worst := 0.0 + for i := range f { + worst = max(worst, math.Abs(vals[i]-want[i])) + } + // The dense solve carries its own round-off; the agreement is + // at the 1e-16 level on these grids, so the bound below keeps + // orders of margin and still pins an operator error. + if worst > 1e-12 { + t.Fatalf("%dx%d: solve differs from the dense inverse by %.3g", g.rows, g.cols, worst) + } + resid := mirroredStencilApply(vals, g.rows, g.cols, hx, hy) + rus := 0.0 + for i := range f { + rus = max(rus, math.Abs(resid[i]-f[i])) + } + if rus > 1e-12 { + t.Fatalf("%dx%d: residual |M u - f| = %.3g", g.rows, g.cols, rus) + } + if mean := neumannTrapzMean(vals, g.rows, g.cols); math.Abs(mean) > 1e-14 { + t.Fatalf("%dx%d: trapezoidal mean of the solution is %.3g, want zero", g.rows, g.cols, mean) + } + t.Logf("%dx%d lx=%g ly=%g: max |solve - dense inverse| = %.3g", g.rows, g.cols, g.lx, g.ly, worst) + } +} + +// TestSolvePoissonNeumannDiscreteEigenfunction pins the two halves of +// the construction: the mirrored stencil really has the cosine +// eigenfunctions with the eigenvalues 2/h²(1-cos(πk/(N-1))) per axis, +// and the solve inverts the stencil on an eigenfunction to transform +// round-off, which is the exactness the spectral solve promises. +func TestSolvePoissonNeumannDiscreteEigenfunction(t *testing.T) { + for _, g := range []struct { + rows, cols, kx, ky int + }{ + {9, 9, 2, 3}, + {17, 17, 4, 5}, + {12, 8, 3, 5}, + } { + hx, hy := 1/float64(g.cols-1), 1/float64(g.rows-1) + mode := make([]float64, g.rows*g.cols) + for r := range g.rows { + for c := range g.cols { + mode[r*g.cols+c] = math.Cos(math.Pi*float64(g.kx)*float64(c)/float64(g.cols-1)) * + math.Cos(math.Pi*float64(g.ky)*float64(r)/float64(g.rows-1)) + } + } + lambda := 2/(hx*hx)*(1-math.Cos(math.Pi*float64(g.kx)/float64(g.cols-1))) + + 2/(hy*hy)*(1-math.Cos(math.Pi*float64(g.ky)/float64(g.rows-1))) + applied := mirroredStencilApply(mode, g.rows, g.cols, hx, hy) + for i := range mode { + if e := math.Abs(applied[i] - lambda*mode[i]); e > 1e-12*lambda { + t.Fatalf("%dx%d k=(%d,%d): M v - lambda v = %.3g at %d, the cosine mode is not an eigenfunction", + g.rows, g.cols, g.kx, g.ky, e, i) + } + } + got, err := SolvePoissonNeumann(mustFloats(t, applied, g.rows, g.cols), 1, 1) + if err != nil { + t.Fatalf("%dx%d: SolvePoissonNeumann: %v", g.rows, g.cols, err) + } + mu := neumannTrapzMean(mode, g.rows, g.cols) + gotVals := make([]float64, got.Len()) + for i := range gotVals { + gotVals[i] = got.FloatAt(i) + } + assertFinite(t, fmt.Sprintf("%dx%d solve", g.rows, g.cols), gotVals) + worst := 0.0 + for i := range mode { + worst = max(worst, math.Abs(gotVals[i]-(mode[i]-mu))) + } + // The acceptance level is the transform's round-off, about + // 1e-14; the measured worst across these grids is 6.1e-15. + if worst > 1e-13 { + t.Fatalf("%dx%d k=(%d,%d): eigenfunction recovered to %.3g only", g.rows, g.cols, g.kx, g.ky, worst) + } + t.Logf("%dx%d k=(%d,%d): eigenfunction error %.3g", g.rows, g.cols, g.kx, g.ky, worst) + } +} + +// TestSolvePoissonNeumannSecondOrder runs the manufactured solution +// u = cos(πx)cos(πy), f = 2π²u (which the trapezoidal measure already +// sees as zero-mean) on three grids: every sample must be finite, the +// recentred error must sit inside the stencil's O(h²) budget, the +// error must fall by about four per grid doubling, and the result must +// carry no trapezoidal mean. The shipped test of the same construction +// compares NaN samples as "not larger", so this one asserts finiteness +// first. +func TestSolvePoissonNeumannSecondOrder(t *testing.T) { + worstOf := map[int]float64{} + for _, n := range []int{9, 17, 33} { + h := 1 / float64(n-1) + flatF := make([]float64, n*n) + exact := make([]float64, n*n) + for r := range n { + for c := range n { + x := float64(c) * h + y := float64(r) * h + exact[r*n+c] = math.Cos(math.Pi*x) * math.Cos(math.Pi*y) + flatF[r*n+c] = 2 * math.Pi * math.Pi * exact[r*n+c] + } + } + u, err := SolvePoissonNeumann(mustFloats(t, flatF, n, n), 1, 1) + if err != nil { + t.Fatalf("grid %d: SolvePoissonNeumann: %v", n, err) + } + vals := make([]float64, u.Len()) + for i := range vals { + vals[i] = u.FloatAt(i) + } + assertFinite(t, fmt.Sprintf("grid %d", n), vals) + if mean := neumannTrapzMean(vals, n, n); math.Abs(mean) > 1e-14 { + t.Fatalf("grid %d: trapezoidal mean of the solution is %.3g, want zero", n, mean) + } + // Recentring: the solve fixes the constant mode, the exact + // solution's own trapezoidal mean is zero, so this is the + // contract's comparison and not a fudge. + got := neumannTrapzMean(vals, n, n) + worst := 0.0 + for i := range exact { + worst = max(worst, math.Abs(vals[i]-got-exact[i])) + } + if worst > 3*h*h { + t.Fatalf("grid %d: worst error %.3g above the O(h²) budget %.3g", n, worst, 3*h*h) + } + worstOf[n] = worst + } + // O(h²) means the error falls by about four per doubling; the + // measured ratios sit near 4, so anything under 3 fails the claim. + if r := worstOf[9] / worstOf[17]; r < 3 { + t.Fatalf("error fell by only %.3f from 9 to 17 points, not the O(h²) factor", r) + } + if r := worstOf[17] / worstOf[33]; r < 3 { + t.Fatalf("error fell by only %.3f from 17 to 33 points, not the O(h²) factor", r) + } + t.Logf("worst errors (h² = %.3g, %.3g, %.3g): %.3g, %.3g, %.3g", + 1/64.0/64.0, 1/16.0/16.0, 1/32.0/32.0, worstOf[9], worstOf[17], worstOf[33]) +} + +// TestSolvePoissonNeumannTrapzRefusal checks the compatibility gate on +// its own measure: a source whose plain mean is zero to rounding but +// whose trapezoidal-weighted sum is not has no mirrored-stencil +// solution, and the solve must refuse it by name rather than divide +// the constant mode by its zero eigenvalue. The corners-only source is +// the smallest such case: the boundary samples carry half the +// trapezoidal weight, so removing the plain mean leaves a large +// trapezoidal component. +func TestSolvePoissonNeumannTrapzRefusal(t *testing.T) { + const n = 9 + flat := make([]float64, n*n) + for _, i := range []int{0, n - 1, (n - 1) * n, n*n - 1} { + flat[i] = 1 + } + plain := 0.0 + for _, v := range flat { + plain += v + } + plain /= float64(n * n) + for i := range flat { + flat[i] -= plain + } + // The construction is only meaningful if the plain mean really + // vanished: otherwise the old gate would have caught it. + residual := 0.0 + for _, v := range flat { + residual += v + } + if math.Abs(residual) > 1e-14 { + t.Fatalf("test construction: residual plain sum %g", residual) + } + if sum := neumannTrapzSum(flat, n, n); math.Abs(sum) < 1 { + t.Fatalf("test construction: trapezoidal sum %.3g is too small to exercise the gate", sum) + } + f := mustFloats(t, flat, n, n) + out, err := SolvePoissonNeumann(f, 1, 1) + if err == nil { + t.Fatalf("zero-plain-mean but trapezoidally incompatible source accepted, solution = %v of %d samples", + out.FloatAt(0), out.Len()) + } + if !strings.Contains(err.Error(), "SolvePoissonNeumann") || !strings.Contains(err.Error(), "trapezoidal") { + t.Fatalf("refusal %q does not name the function and the measure", err) + } + // The all-ones source has no solution on either measure, and the + // shipped refusal test's shape guards stay intact. + if _, err := SolvePoissonNeumann(mustFloats(t, []float64{1, 1, 1, 1, 1, 1, 1, 1, 1}, 3, 3), 1, 1); err == nil { + t.Fatal("nonzero-mean source accepted") + } +} + +// The view contract of the new entry points: a rebased view's payload +// is longer than its element count and starts past offset zero, so a +// kernel that reaches for the raw payload instead of the elements +// reads the wrong window of the backing array. Each test compares the +// view answer against the same values in a fresh array. +func TestEntryPointsOnView(t *testing.T) { + ramp := make([]float64, 20) + for i := range ramp { + ramp[i] = float64(i) + } + back, err := core.FromFloats(ramp, 20) + if err != nil { + t.Fatal(err) + } + view, err := core.Slice(back, 0, 3, 19) + if err != nil { + t.Fatal(err) + } + fresh, err := core.FromFloats(append([]float64{}, ramp[3:19]...), 16) + if err != nil { + t.Fatal(err) + } + median, err := MedianFilter(view, 3) + if err != nil { + t.Fatalf("MedianFilter: %v", err) + } + wantMedian, err := MedianFilter(fresh, 3) + if err != nil { + t.Fatalf("MedianFilter: %v", err) + } + for i := range 16 { + if median.FloatAt(i) != wantMedian.FloatAt(i) { + t.Fatalf("MedianFilter(view)[%d] = %v, want %v", i, median.FloatAt(i), wantMedian.FloatAt(i)) + } + } + rank, err := RankFilter(view, 5, 0) + if err != nil { + t.Fatalf("RankFilter: %v", err) + } + wantRank, err := RankFilter(fresh, 5, 0) + if err != nil { + t.Fatalf("RankFilter: %v", err) + } + for i := range 16 { + if rank.FloatAt(i) != wantRank.FloatAt(i) { + t.Fatalf("RankFilter(view)[%d] = %v, want %v", i, rank.FloatAt(i), wantRank.FloatAt(i)) + } + } + b, a, err := ButterworthLowPass(2, 2.0, 0.4) + if err != nil { + t.Fatal(err) + } + filtered, err := Filtfilt(b, a, view) + if err != nil { + t.Fatalf("Filtfilt: %v", err) + } + wantFiltered, err := Filtfilt(b, a, fresh) + if err != nil { + t.Fatalf("Filtfilt: %v", err) + } + for i := range 16 { + if math.Float64bits(filtered.FloatAt(i)) != math.Float64bits(wantFiltered.FloatAt(i)) { + t.Fatalf("Filtfilt(view)[%d] = %v, want %v", i, filtered.FloatAt(i), wantFiltered.FloatAt(i)) + } + } + coef, err := DaubechiesDWT(view, DB2, 1, DWTPeriodic) + if err != nil { + t.Fatalf("DaubechiesDWT: %v", err) + } + wantCoef, err := DaubechiesDWT(fresh, DB2, 1, DWTPeriodic) + if err != nil { + t.Fatalf("DaubechiesDWT: %v", err) + } + for i := range coef.Len() { + if math.Float64bits(coef.FloatAt(i)) != math.Float64bits(wantCoef.FloatAt(i)) { + t.Fatalf("DaubechiesDWT(view)[%d] = %v, want %v", i, coef.FloatAt(i), wantCoef.FloatAt(i)) + } + } +} + +// TestWindowKaiserRefusesNonFiniteBeta pins the refusal of a beta the +// Bessel ratio cannot carry: +Inf used to answer an all-NaN window. +func TestWindowKaiserRefusesNonFiniteBeta(t *testing.T) { + for _, beta := range []float64{math.Inf(1), math.Inf(-1), math.NaN()} { + if _, err := WindowKaiser(8, beta, false); err == nil { + t.Fatalf("WindowKaiser beta %g: expected an error, got none", beta) + } + } +} diff --git a/signal/correlate.go b/signal/correlate.go new file mode 100644 index 0000000..79b290f --- /dev/null +++ b/signal/correlate.go @@ -0,0 +1,250 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Correlation functions: the time-series staples on top of +// the FFT. The autocorrelation and cross-correlation run as spectral +// products (one transform pair whatever the lag count), the partial +// autocorrelation follows by Durbin-Levinson recursion over the ACF, +// the algorithm that defined the PACF in the first place. + +// Autocorrelate returns the sample autocorrelation of a rank-1 real +// signal at lags 0..maxLag: the mean is removed first and every value +// is normalised by the lag-zero sum (the biased estimator, the one +// whose Durbin-Levinson and Bartlett-band theory is written for). +// maxLag must lie in [0, n−1]; the result has maxLag+1 entries and +// starts at 1, unless the signal is constant: its zero power makes +// every lag 0, lag zero included. +func Autocorrelate(x *core.Array, maxLag int) (*core.Array, error) { + const name = "Autocorrelate" + if x.NDim() != 1 { + return nil, base.Errf("%s: needs a rank-1 signal, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, base.Errf("%s: complex signals are not supported", name) + } + n := x.Len() + if n == 0 { + return nil, base.Errf("%s: an empty signal has no correlation", name) + } + if maxLag < 0 || maxLag >= n { + return nil, base.Errf("%s: maxLag must lie in [0, %d], got %d", name, n-1, maxLag) + } + // Demean. The payload walk reads the same values FloatAt would + // (widenFloats widens exactly), in the same order, so the mean + // and every padded sample keep their bits. + xF := x.RawFloats() + if xF == nil || x.Strided() { + xF = widenFloats(x) + } + // A non-finite sample would drive the normaliser NaN and publish + // NaN correlations with no error, so it is refused up front (the + // guard SolvePoissonPeriodic applies to its source). + for i := range n { + if v := xF[i]; math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: x holds the non-finite value %g at %d", name, v, i) + } + } + mean := 0.0 + for i := range n { + mean += xF[i] + } + mean /= float64(n) + // The correlation runs as one transform pair over the zero-padded + // signal, in place on a private buffer: the values that reach the + // butterflies and the arithmetic that follows are the FFT/IFFT + // path's unchanged, minus the copies into and out of arrays that + // nobody else ever read. + padded := make([]complex128, 2*n) + for i := range n { + padded[i] = complex(xF[i]-mean, 0) + } + transform(padded, -1) + // |F|²: multiply by the conjugate, then back. The zero-padded + // transform makes the circular correlation linear. + for i := range padded { + padded[i] = padded[i] * conj128(padded[i]) + } + transform(padded, +1) + // The inverse scaling, applied with IFFT's own operation order. + m := float64(2 * n) + for i := range padded { + padded[i] /= complex(m, 0) + } + den := real(padded[0]) + out := make([]float64, maxLag+1) + for k := range maxLag + 1 { + if den == 0 { + out[k] = 0 + continue + } + out[k] = real(padded[k]) / den + } + return floatsFromArrayMust(out, []int{maxLag + 1}), nil +} + +// conj128 conjugates one complex number (the local spelling keeps the +// dependency surface to the packages already imported). +func conj128(z complex128) complex128 { return complex(real(z), -imag(z)) } + +// CrossCorrelate returns the raw cross-correlation of two equal-length +// real signals at every lag: the result has 2n−1 entries, index i +// carrying lag i−(n−1), value r(lag) = Σ_t x[t+lag]·y[t] over the +// overlapping t. No mean removal, no normalisation: correlation as +// the inner product of shifted copies, the definition convolution and +// matched filtering share. +func CrossCorrelate(x, y *core.Array) (*core.Array, error) { + const name = "CrossCorrelate" + if x.NDim() != 1 || y.NDim() != 1 { + return nil, base.Errf("%s: needs two rank-1 signals, got %s and %s", + name, base.ShapeText(x.Shape()), base.ShapeText(y.Shape())) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, base.Errf("%s: complex signals are not supported", name) + } + n := x.Len() + if n == 0 || y.Len() != n { + return nil, base.Errf("%s: the signals must be equal-length non-empty, got %d and %d", + name, n, y.Len()) + } + // Zero-pad both signals through the raw payloads where possible; + // the widened reads equal the FloatAt values bit for bit. + xF := x.RawFloats() + if xF == nil || x.Strided() { + xF = widenFloats(x) + } + yF := y.RawFloats() + if yF == nil || y.Strided() { + yF = widenFloats(y) + } + // A non-finite sample would publish NaN correlations with no + // error, so both signals are refused up front. + for i := range n { + if v := xF[i]; math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: x holds the non-finite value %g at %d", name, v, i) + } + if v := yF[i]; math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: y holds the non-finite value %g at %d", name, v, i) + } + } + // The spectra multiply in place on private buffers: the values + // reaching the transforms and the per-element products are the + // FFT/mulSpectraConj/IFFT path's unchanged, minus the copies into + // and out of arrays nobody else ever read. + px := make([]complex128, 2*n) + py := make([]complex128, 2*n) + for i := range n { + px[i] = complex(xF[i], 0) + py[i] = complex(yF[i], 0) + } + transform(px, -1) + transform(py, -1) + // Fx·conj(Fy) transforms back to Σ_t x[t+lag]·y[t], the documented + // orientation; the index mapping lands the positive lags in the + // first half. + // The product runs in mulSpectraConj, the array-based helper the + // FFT path always used: the compiler's FMA fusion of the complex + // multiply depends on the compiled shape of its surroundings, and + // the helper's shape is the one whose product bits this result is + // pinned to. The wrappers take the buffers without copying; the + // product array is the one allocation of the old path this keeps. + prod, perr := mulSpectraConj(complexFromArrayMust(px, []int{2 * n}), + complexFromArrayMust(py, []int{2 * n})) + if perr != nil { + return nil, base.Errf("%s: %w", name, perr) + } + pc := prod.RawComplexes() + transform(pc, +1) + m := float64(2 * n) + for i := range pc { + pc[i] /= complex(m, 0) + } + // Lag 0 sits at index 0 of the transformed product, lag k at k and + // negative lag −k at 2n−k; the output centres them on index n−1. + out := make([]float64, 2*n-1) + for k := range 2*n - 1 { + lag := k - (n - 1) + if lag >= 0 { + out[k] = real(pc[lag]) + } else { + out[k] = real(pc[2*n+lag]) + } + } + return floatsFromArrayMust(out, []int{2*n - 1}), nil +} + +// mulSpectraConj multiplies the first spectrum by the conjugate of the +// second, element-wise. Both spectra come from the FFT and stream +// their complex payloads directly; the accessor loop stays as the +// fallback for anything else. The function is kept byte-for-byte as +// the FFT path always defined it: its compiled FMA pattern is what the +// correlation results are pinned to. +func mulSpectraConj(a, b *core.Array) (*core.Array, error) { + as := a.RawComplexes() + bs := b.RawComplexes() + if as != nil && bs != nil && !a.Strided() && !b.Strided() { + out := make([]complex128, len(as)) + for i := range out { + out[i] = as[i] * conj128(bs[i]) + } + return core.ComplexFromArray(out, a.Shape()...) + } + out := make([]complex128, a.Len()) + for i := range out { + out[i] = a.ComplexAt(i) * conj128(b.ComplexAt(i)) + } + return core.ComplexFromArray(out, a.Shape()...) +} + +// PartialAutocorrelate returns the partial autocorrelation of a rank-1 +// real signal at lags 1..maxLag, by the Durbin-Levinson recursion over +// the biased autocorrelation: pacf[k] is the last coefficient of the +// order-k AR fit, the "correlation with the intermediate lags +// regressed out" that box-jenkins identification reads. maxLag must +// lie in [1, n−2] for the recursion to stay meaningful. +func PartialAutocorrelate(x *core.Array, maxLag int) (*core.Array, error) { + const name = "PartialAutocorrelate" + n := x.Len() + if maxLag < 1 || maxLag > n-2 { + return nil, base.Errf("%s: maxLag must lie in [1, %d], got %d", name, n-2, maxLag) + } + acf, err := Autocorrelate(x, maxLag) + if err != nil { + return nil, err + } + // r[k] = acf at lag k (1-based recursion indexing). + r := make([]float64, maxLag+1) + r[0] = 1 + for k := range maxLag { + r[k+1] = acf.FloatAt(k + 1) + } + phi := make([]float64, maxLag+1) // current order's coefficients + phiPrev := make([]float64, maxLag+1) + out := make([]float64, maxLag) + for k := 1; k <= maxLag; k++ { + num := r[k] + den := 1.0 + for j := 1; j < k; j++ { + num -= phiPrev[j] * r[k-j] + den -= phiPrev[j] * r[j] + } + if den == 0 { + return nil, base.Errf("%s: the recursion broke down at lag %d", name, k) + } + phi[k] = num / den + for j := 1; j < k; j++ { + phi[j] = phiPrev[j] - phi[k]*phiPrev[k-j] + } + copy(phiPrev, phi) + out[k-1] = phi[k] + } + return floatsFromArrayMust(out, []int{maxLag}), nil +} diff --git a/signal/correlate_test.go b/signal/correlate_test.go new file mode 100644 index 0000000..a4789a5 --- /dev/null +++ b/signal/correlate_test.go @@ -0,0 +1,152 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestAutocorrelateAR1 pins the ACF of an AR(1) process against its +// geometric theoretical decay. +func TestAutocorrelateAR1(t *testing.T) { + const ( + phi = 0.6 + n = 20000 + ) + g := core.NewGenerator(23) + x := make([]float64, n) + noise := 0.0 + for i := range n { + noise = phi*noise + g.NormalUnit() + x[i] = noise + } + xArr, err := core.FromFloats(x, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + acf, err := Autocorrelate(xArr, 5) + if err != nil { + t.Fatalf("Autocorrelate: %v", err) + } + if math.Abs(acf.FloatAt(0)-1) > 1e-12 { + t.Fatalf("acf[0] = %.12f, want 1", acf.FloatAt(0)) + } + for k := 1; k <= 5; k++ { + want := math.Pow(phi, float64(k)) + if math.Abs(acf.FloatAt(k)-want) > 0.06 { + t.Fatalf("acf[%d] = %.4f, want %.4f", k, acf.FloatAt(k), want) + } + } +} + +// TestAutocorrelateBrute pins the ACF against the direct sum. +func TestAutocorrelateBrute(t *testing.T) { + x := []float64{3, -1, 4, -1.5, 2, -0.5, 1, 2} + n := len(x) + xArr, _ := core.FromFloats(x, n) + acf, err := Autocorrelate(xArr, n-1) + if err != nil { + t.Fatalf("Autocorrelate: %v", err) + } + mean := 0.0 + for _, v := range x { + mean += v + } + mean /= float64(n) + num := 0.0 + for _, v := range x { + num += (v - mean) * (v - mean) + } + for k := range n { + s := 0.0 + for t := 0; t+k < n; t++ { + s += (x[t] - mean) * (x[t+k] - mean) + } + want := s / num + if math.Abs(acf.FloatAt(k)-want) > 1e-12 { + t.Fatalf("acf[%d] = %.12f, direct sum %.12f", k, acf.FloatAt(k), want) + } + } +} + +// TestCrossCorrelateBrute pins the XCF against the direct sum on the +// documented lag convention. +func TestCrossCorrelateBrute(t *testing.T) { + x := []float64{2, -1, 3, 0.5, -2} + y := []float64{1, 4, -2, 0.5, 3} + n := len(x) + xArr, _ := core.FromFloats(x, n) + yArr, _ := core.FromFloats(y, n) + xcf, err := CrossCorrelate(xArr, yArr) + if err != nil { + t.Fatalf("CrossCorrelate: %v", err) + } + if xcf.Len() != 2*n-1 { + t.Fatalf("length %d, want %d", xcf.Len(), 2*n-1) + } + for k := range 2*n - 1 { + lag := k - (n - 1) + s := 0.0 + for t := 0; t+lag < n && t < n; t++ { + if t+lag < 0 { + continue + } + s += x[t+lag] * y[t] + } + if math.Abs(xcf.FloatAt(k)-s) > 1e-12 { + t.Fatalf("xcf[lag %d] = %.12f, direct sum %.12f", lag, xcf.FloatAt(k), s) + } + } +} + +// TestPartialAutocorrelateAR1 pins the PACF signature of an AR(1): +// lag 1 estimates phi, the rest are noise around zero. +func TestPartialAutocorrelateAR1(t *testing.T) { + const ( + phi = 0.7 + n = 30000 + ) + g := core.NewGenerator(31) + x := make([]float64, n) + e := 0.0 + for i := range n { + e = phi*e + g.NormalUnit() + x[i] = e + } + xArr, _ := core.FromFloats(x, n) + pacf, err := PartialAutocorrelate(xArr, 6) + if err != nil { + t.Fatalf("PartialAutocorrelate: %v", err) + } + if math.Abs(pacf.FloatAt(0)-phi) > 0.05 { + t.Fatalf("pacf[1] = %.4f, want %.2f", pacf.FloatAt(0), phi) + } + for k := 2; k <= 6; k++ { + if math.Abs(pacf.FloatAt(k-1)) > 0.05 { + t.Fatalf("pacf[%d] = %.4f, an AR(1) cuts off after lag 1", k, pacf.FloatAt(k-1)) + } + } +} + +// TestCorrelateErrors pins the input gates. +func TestCorrelateErrors(t *testing.T) { + x, _ := core.FromFloats([]float64{1, 2, 3}, 3) + if _, err := Autocorrelate(x, 3); err == nil { + t.Error("maxLag ≥ n accepted") + } + m, _ := core.FromFloats([]float64{1, 2}, 1, 2) + if _, err := Autocorrelate(m, 0); err == nil { + t.Error("rank-2 signal accepted") + } + y, _ := core.FromFloats([]float64{1, 2}, 2) + if _, err := CrossCorrelate(x, y); err == nil { + t.Error("length mismatch accepted") + } + if _, err := PartialAutocorrelate(x, 5); err == nil { + t.Error("oversized PACF maxLag accepted") + } +} diff --git a/signal/dct.go b/signal/dct.go new file mode 100644 index 0000000..f2c8cf5 --- /dev/null +++ b/signal/dct.go @@ -0,0 +1,209 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// Discrete cosine and sine transforms, types I to IV, in the +// orthonormal convention. Every transform is one padded inverse FFT: +// the input carries a per-sample complex phase, the result is read at +// a shifted frequency index and rotated by a per-frequency phase, +// taking the real half for the cosine families and the imaginary half +// for the sine ones. The orthonormal scaling makes each transform its +// own inverse or the transpose of its partner, so the inverse entry +// points are aliases: IDCT(x, 2) is DCT(x, 3), IDCT(x, 4) is +// DCT(x, 4), and the sine family follows the same rule. + +// DCT computes the orthonormal discrete cosine transform of type +// kind (1 to 4) of the vector x. A kind outside 1 to 4, a rank other +// than 1, an empty vector, or type I on fewer than two points is an +// error. +func DCT(x *core.Array, kind int) (*core.Array, error) { + return dctdst(x, kind, true) +} + +// IDCT computes the inverse orthonormal DCT: the transpose partner of +// the forward transform of the same kind. +func IDCT(x *core.Array, kind int) (*core.Array, error) { + switch kind { + case 1, 4: + return dctdst(x, kind, true) + case 2: + return dctdst(x, 3, true) + case 3: + return dctdst(x, 2, true) + } + return nil, base.Errf("IDCT: kind must be 1 to 4, got %d", kind) +} + +// DST computes the orthonormal discrete sine transform of type kind +// (1 to 4) of the vector x. +func DST(x *core.Array, kind int) (*core.Array, error) { + return dctdst(x, kind, false) +} + +// IDST computes the inverse orthonormal DST. +func IDST(x *core.Array, kind int) (*core.Array, error) { + switch kind { + case 1, 4: + return dctdst(x, kind, false) + case 2: + return dctdst(x, 3, false) + case 3: + return dctdst(x, 2, false) + } + return nil, base.Errf("IDST: kind must be 1 to 4, got %d", kind) +} + +// dctdst dispatches the eight transforms onto the shared core. The +// table per type and direction: the input phase ramp (half weights +// where the convention halves an endpoint), the frequency shift into +// the padded spectrum, the per-frequency rotation including the +// output-side half weights, the picked half and the normalisation. +func dctdst(x *core.Array, kind int, cosine bool) (*core.Array, error) { + const name = "DCT/DST" + if kind < 1 || kind > 4 { + return nil, base.Errf("%s: kind must be 1 to 4, got %d", name, kind) + } + if x.Dtype() == core.Complex { + return nil, base.Errf("%s: complex arrays are not supported", name) + } + if x.NDim() != 1 || x.Len() == 0 { + return nil, base.Errf("%s: the input must be a non-empty vector, got shape %s", name, base.ShapeText(x.Shape())) + } + n := x.Len() + if kind == 1 && n < 2 { + return nil, base.Errf("%s: type I needs at least two points, got %d", name, n) + } + vals := make([]float64, n) + for i := range n { + vals[i] = x.FloatAt(i) + } + part := func(z complex128) float64 { return real(z) } + if !cosine { + part = func(z complex128) float64 { return imag(z) } + } + // phase weights the samples, shift picks the spectrum index, + // outer rotates the picked bin, norm scales the answer. + phase := func(j int) complex128 { return 1 } + outer := func(k int) complex128 { return 1 } + norm := func(k int) float64 { return math.Sqrt(2 / float64(n)) } + shift, m := 0, 2*n + switch kind { + case 1: + if cosine { + m = 2*n - 2 + phase = func(j int) complex128 { + if j == 0 || j == n-1 { + return math.Sqrt2 / 2 + } + return 1 + } + outer = func(k int) complex128 { + w := 1.0 + if k == 0 || k == n-1 { + w = math.Sqrt2 / 2 + } + return complex(w, 0) + } + norm = func(int) float64 { return math.Sqrt(2 / float64(n-1)) } + } else { + // sin(π(j+1)(k+1)/(n+1)) on the length-2(n+1) grid: the + // shifted, rotated read carries the whole phase, and the + // sine family carries no endpoint weights at all. + m = 2*n + 2 + outer = func(k int) complex128 { + return base.CmplxPolar(1, math.Pi*float64(k+1)/float64(n+1)) + } + norm = func(int) float64 { return math.Sqrt(2 / float64(n+1)) } + shift = 1 + } + case 2: + if cosine { + // cos(π(2j+1)k/2n): the frequency rotation carries the + // (2j+1) odd shift; the k = 0 row halves. + outer = func(k int) complex128 { + return base.CmplxPolar(1, math.Pi*float64(k)/float64(2*n)) + } + norm = func(k int) float64 { + if k == 0 { + return math.Sqrt(1 / float64(n)) + } + return math.Sqrt(2 / float64(n)) + } + } else { + // sin(π(2j+1)(k+1)/2n): read the shifted bin. + outer = func(k int) complex128 { + return base.CmplxPolar(1, math.Pi*float64(k+1)/float64(2*n)) + } + norm = func(k int) float64 { + if k == n-1 { + return math.Sqrt(1 / float64(n)) + } + return math.Sqrt(2 / float64(n)) + } + shift = 1 + } + case 3: + // The transpose of type 2: cos(π(2k+1)j/2n) and + // sin(π(j+1)(2k+1)/2n), with the normalisation on the input + // index and none on the output; the transpose of the type-2 + // weights lands on the samples. + phase = func(j int) complex128 { + w := 1.0 + if (cosine && j == 0) || (!cosine && j == n-1) { + w = math.Sqrt(1 / float64(n)) + } else { + w = math.Sqrt(2 / float64(n)) + } + return base.CmplxPolar(w, math.Pi*float64(j)/float64(2*n)) + } + if cosine { + outer = func(int) complex128 { return 1 } + } else { + outer = func(k int) complex128 { + return base.CmplxPolar(1, math.Pi*float64(2*k+1)/float64(2*n)) + } + } + norm = func(int) float64 { return 1 } + case 4: + phase = func(j int) complex128 { + return base.CmplxPolar(1, math.Pi*float64(2*j+1)/float64(4*n)) + } + outer = func(k int) complex128 { + return base.CmplxPolar(1, math.Pi*float64(k)/float64(2*n)) + } + } + z := make([]complex128, m) + for j := range n { + p := phase(j) + z[j] = complex(vals[j]*real(p), vals[j]*imag(p)) + } + w := paddedIFFT(z) + out := make([]float64, n) + for k := range n { + out[k] = norm(k) * part(w[k+shift]*outer(k)) + } + return floatsFromArrayMust(out, []int{n}), nil +} + +// paddedIFFT computes the unscaled inverse DFT of z: w[k] = +// Σ_j z[j]·e^{+2πijk/m} for m = len(z), by conjugating around the +// existing transform. The divide-then-multiply pair reproduces the +// exact rounding of the old IFFT-then-rescale path, on one buffer and +// one copy fewer. z is overwritten and returned. +func paddedIFFT(z []complex128) []complex128 { + m := len(z) + transform(z, +1) + for i := range z { + z[i] /= complex(float64(m), 0) + z[i] *= complex(float64(m), 0) + } + return z +} diff --git a/signal/dct_test.go b/signal/dct_test.go new file mode 100644 index 0000000..3f10430 --- /dev/null +++ b/signal/dct_test.go @@ -0,0 +1,225 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "testing" +) + +// dctNaive evaluates the orthonormal DCT/DST of the given type by the +// defining sum, the reference the padded-FFT route must match. +func dctNaive(x []float64, kind int, cosine bool) []float64 { + n := len(x) + y := make([]float64, n) + part := math.Cos + if !cosine { + part = math.Sin + } + norm := func(k int) float64 { return math.Sqrt(2 / float64(n)) } + wIn := func(j int) float64 { return 1 } + wOut := func(k int) float64 { return 1 } + switch kind { + case 1: + if cosine { + norm = func(int) float64 { return math.Sqrt(2 / float64(n-1)) } + wIn = func(j int) float64 { + if j == 0 || j == n-1 { + return math.Sqrt2 / 2 + } + return 1 + } + wOut = wIn + } else { + norm = func(int) float64 { return math.Sqrt(2 / float64(n+1)) } + } + case 2: + if cosine { + norm = func(k int) float64 { + if k == 0 { + return math.Sqrt(1 / float64(n)) + } + return math.Sqrt(2 / float64(n)) + } + } else { + norm = func(k int) float64 { + if k == n-1 { + return math.Sqrt(1 / float64(n)) + } + return math.Sqrt(2 / float64(n)) + } + } + case 3: + wIn = func(j int) float64 { + if (cosine && j == 0) || (!cosine && j == n-1) { + return math.Sqrt(1 / float64(n)) + } + return math.Sqrt(2 / float64(n)) + } + wOut = func(int) float64 { return 1 } + norm = func(int) float64 { return 1 } + } + for k := range n { + sum := 0.0 + for j := range n { + var arg float64 + switch kind { + case 1: + if cosine { + arg = math.Pi * float64(j*k) / float64(n-1) + } else { + arg = math.Pi * float64((j+1)*(k+1)) / float64(n+1) + } + case 2: + if cosine { + arg = math.Pi * float64((2*j+1)*k) / float64(2*n) + } else { + arg = math.Pi * float64((2*j+1)*(k+1)) / float64(2*n) + } + case 3: + if cosine { + arg = math.Pi * float64((2*k+1)*j) / float64(2*n) + } else { + arg = math.Pi * float64((2*k+1)*(j+1)) / float64(2*n) + } + case 4: + arg = math.Pi * float64((2*j+1)*(2*k+1)) / float64(4*n) + } + sum += wIn(j) * x[j] * part(arg) + } + y[k] = norm(k) * wOut(k) * sum + } + return y +} + +// TestDCTDSTAgainstDefinitions checks every type and direction against +// the defining sums on a deterministic vector. +func TestDCTDSTAgainstDefinitions(t *testing.T) { + n := 8 + x := make([]float64, n) + for i := range n { + x[i] = math.Sin(float64(3*i+1)) + 0.25*math.Cos(float64(5*i)) + } + xa := mustFloats(t, x, n) + for kind := 1; kind <= 4; kind++ { + if kind == 1 && n < 2 { + continue + } + for _, cosine := range []bool{true, false} { + got, err := dctdst(xa, kind, cosine) + if err != nil { + t.Fatalf("kind %d cosine %v: %v", kind, cosine, err) + } + want := dctNaive(x, kind, cosine) + for k := range n { + if math.Abs(got.FloatAt(k)-want[k]) > 1e-12 { + t.Fatalf("kind %d cosine %v: y[%d] = %.14g, want %.14g", + kind, cosine, k, got.FloatAt(k), want[k]) + } + } + } + } +} + +// TestDCTDSTInverses checks the aliasing rule: forward then inverse +// returns the input for all eight pairs. +func TestDCTDSTInverses(t *testing.T) { + n := 11 + x := make([]float64, n) + for i := range n { + x[i] = math.Cos(float64(2*i + 1)) + } + xa := mustFloats(t, x, n) + pairs := []struct { + fwd, inv func(*core.Array, int) (*core.Array, error) + kind int + }{{DCT, IDCT, 1}, {DCT, IDCT, 2}, {DCT, IDCT, 3}, {DCT, IDCT, 4}, + {DST, IDST, 1}, {DST, IDST, 2}, {DST, IDST, 3}, {DST, IDST, 4}} + for _, p := range pairs { + mid, err := p.fwd(xa, p.kind) + if err != nil { + t.Fatalf("kind %d forward: %v", p.kind, err) + } + back, err := p.inv(mid, p.kind) + if err != nil { + t.Fatalf("kind %d inverse: %v", p.kind, err) + } + for i := range n { + if math.Abs(back.FloatAt(i)-x[i]) > 1e-11 { + t.Fatalf("kind %d round trip [%d] = %.14g, want %.14g", + p.kind, i, back.FloatAt(i), x[i]) + } + } + } +} + +// TestDCTDSTOrthogonality pins the orthonormal claim on the +// self-inverse types: applying DCT-I or DCT-IV to the identity's +// columns returns a matrix whose Gram is the identity. +func TestDCTDSTOrthogonality(t *testing.T) { + n := 6 + for _, kind := range []int{1, 4} { + col := make([]float64, n) + gram := make([]float64, n*n) + for j := range n { + for i := range n { + col[i] = 0 + if i == j { + col[i] = 1 + } + } + y, err := DCT(mustFloats(t, col, n), kind) + if err != nil { + t.Fatalf("DCT kind %d: %v", kind, err) + } + for i := range n { + gram[i*n+j] = y.FloatAt(i) + } + } + for i := range n { + for j := range n { + s := 0.0 + for l := range n { + s += gram[l*n+i] * gram[l*n+j] + } + want := 0.0 + if i == j { + want = 1 + } + if math.Abs(s-want) > 1e-12 { + t.Fatalf("DCT-%d Gram[%d][%d] = %.14g, want %.14g", kind, i, j, s, want) + } + } + } + } +} + +// TestDCTDSTErrors pins the validation contract. +func TestDCTDSTErrors(t *testing.T) { + xa := mustFloats(t, []float64{1, 2, 3}, 3) + if _, err := DCT(xa, 5); err == nil { + t.Fatal("expected an error for kind 5") + } + if _, err := IDCT(xa, 0); err == nil { + t.Fatal("expected an error for kind 0") + } + if _, err := DST(xa, 7); err == nil { + t.Fatal("expected an error for kind 7") + } + if _, err := IDST(xa, -1); err == nil { + t.Fatal("expected an error for kind -1") + } + if _, err := DCT(mustFloats(t, []float64{1}, 1), 1); err == nil { + t.Fatal("expected an error for type I on one point") + } + rank2, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := DCT(rank2, 2); err == nil { + t.Fatal("expected an error for a rank-2 input") + } + if _, err := DST(mustFloats(t, nil), 2); err == nil { + t.Fatal("expected an error for an empty input") + } +} diff --git a/signal/doc.go b/signal/doc.go new file mode 100644 index 0000000..3f65991 --- /dev/null +++ b/signal/doc.go @@ -0,0 +1,55 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package signal carries the transforms, spectra, filters and stencils +// of the library: the Fourier, cosine, sine and wavelet transforms, the +// spectral estimators, IIR filter design and application, sample-rate +// conversion, the Hilbert envelope, convolution and pooling, rank +// filters, spatial stencils, state-space estimation and time-series +// models. +// +// # Shape of the API +// +// Every call takes whole arrays and returns whole arrays. Arrays are +// immutable values, so no call modifies its input: a transform that +// needs scratch space runs on a private copy. +// +// The package holds no state between calls. There is no streaming +// filter object and nothing to plan or reuse: a design returns the +// direct-form coefficients b and a, and the caller hands those to +// FilterApply for one pass or to Filtfilt for a zero-phase forward and +// backward sweep. A wavelet family is named at every call for the same +// reason. +// +// # What the package does not do +// +// - No second-order-section cascade. Every design returns direct-form +// coefficients, and direct-form filtering loses digits as the order +// climbs; past roughly order eight the caller is expected to split +// the design into second-order sections. +// - No complex input to the estimators and resamplers. WelchPSD, +// Spectrogram, STFT, LombScargle, Decimate, Resample, +// AnalyticSignal and the Kalman filters refuse a complex array. The +// Fourier transforms are the complex path: FFT, IFFT, FFTN and +// their relatives accept a real or a complex array, while RFFT and +// IRFFT are the real-input pair. +// - No batched or multi-channel spectral work. WelchPSD, STFT, +// Spectrogram, the resamplers, the filter application and the +// wavelet transforms take one rank-1 signal per call, so a stack of +// channels is looped by the caller. +// - No statistical distributions, hypothesis tests or regressions. +// Those live in the stats package; the models this package fits are +// the autoregression, the ARMA model and the state-space model. +// +// # Conventions +// +// A sample rate is named fs and given in hertz. A cutoff or a band edge +// lies strictly inside (0, fs/2). A design returns the coefficients in +// the u = z⁻¹ convention with the denominator leading a one, so a[0] is +// 1 and a[1] is the first feedback tap. A spectral estimate is +// one-sided unless the call says otherwise, with the doubling applied +// between DC and the Nyquist bin. A wavelet transform packs its +// coefficients as [A_levels, D_levels, …, D_1], the deepest +// approximation first and the finest detail last, the layout DWT and +// DaubechiesDWT share. +package signal diff --git a/signal/dtypes_census_test.go b/signal/dtypes_census_test.go new file mode 100644 index 0000000..0562176 --- /dev/null +++ b/signal/dtypes_census_test.go @@ -0,0 +1,723 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The dtype census for signal: every array-taking public entry probed +// with Bool, the narrow integers and the Int anchor against a float64 +// baseline carrying exactly the widened probe values. The signal +// entries widen non-float inputs through widenFloats or the complex +// readers by design; the spectral Poisson family refuses every +// non-float dtype by name; nothing panics or silently misreads. + +var sgDtypes = []core.Dtype{core.Bool, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32, core.Int} + +type sgMaker func(vals []float64, shape ...int) *core.Array + +func sgCast(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 sgMakers(t *testing.T, dt core.Dtype) (probe, base sgMaker) { + t.Helper() + castOf := func(vals []float64) []float64 { + out := make([]float64, len(vals)) + for i, v := range vals { + out[i] = sgCast(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 sgElems(t *testing.T, a *core.Array) []complex128 { + t.Helper() + out := make([]complex128, a.Len()) + for i := range out { + if a.Dtype() == core.Complex { + out[i] = a.ComplexAt(i) + continue + } + out[i] = complex(a.FloatAt(i), 0) + } + return out +} + +func sgArrays(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 := sgElems(t, p), sgElems(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 sgFloat(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 pv != bv { + t.Fatalf("%s(%s) = %v, want the baseline %v", label, dt, pv, bv) + } +} + +func sgFloats(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 sgWantErr(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) + } + } +} + +// TestDtypesCensusSignal probes every array-taking public entry. +func TestDtypesCensusSignal(t *testing.T) { + sig8 := []float64{3, 1, 4, 1, 5, 2, 6, 2} + sig16 := []float64{3, 1, 4, 1, 5, 2, 6, 5, 3, 5, 8, 9, 7, 9, 3, 2} + img16 := []float64{3, 1, 4, 1, 5, 2, 6, 2, 7, 1, 8, 2, 8, 1, 8, 2} + rows := []struct { + name string + run func(t *testing.T, probe, base sgMaker, dt core.Dtype) + }{ + {"FFT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := FFT(probe(sig8, 8)) + b, berr := FFT(base(sig8, 8)) + sgArrays(t, "FFT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"IFFT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := IFFT(probe(sig8, 8)) + b, berr := IFFT(base(sig8, 8)) + sgArrays(t, "IFFT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"FFT2", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := FFT2(probe(img16[:8], 2, 4)) + b, berr := FFT2(base(img16[:8], 2, 4)) + sgArrays(t, "FFT2", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"FFT3", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := FFT3(probe(sig8, 2, 2, 2)) + b, berr := FFT3(base(sig8, 2, 2, 2)) + sgArrays(t, "FFT3", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"IFFT2", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := IFFT2(probe(img16[:8], 2, 4)) + b, berr := IFFT2(base(img16[:8], 2, 4)) + sgArrays(t, "IFFT2", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"IFFT3", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := IFFT3(probe(sig8, 2, 2, 2)) + b, berr := IFFT3(base(sig8, 2, 2, 2)) + sgArrays(t, "IFFT3", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"FFTN", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := FFTN(probe(img16[:8], 2, 4), nil) + b, berr := FFTN(base(img16[:8], 2, 4), nil) + sgArrays(t, "FFTN", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"IFFTN", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := IFFTN(probe(img16[:8], 2, 4), nil) + b, berr := IFFTN(base(img16[:8], 2, 4), nil) + sgArrays(t, "IFFTN", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"RFFT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := RFFT(probe(sig8, 8)) + b, berr := RFFT(base(sig8, 8)) + sgArrays(t, "RFFT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"DCT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := DCT(probe(sig8, 8), 2) + b, berr := DCT(base(sig8, 8), 2) + sgArrays(t, "DCT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"IDCT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := IDCT(probe(sig8, 8), 2) + b, berr := IDCT(base(sig8, 8), 2) + sgArrays(t, "IDCT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"DST", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := DST(probe(sig8, 8), 2) + b, berr := DST(base(sig8, 8), 2) + sgArrays(t, "DST", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"IDST", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := IDST(probe(sig8, 8), 2) + b, berr := IDST(base(sig8, 8), 2) + sgArrays(t, "IDST", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"DWT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := DWT(probe(sig8, 8), 2) + b, berr := DWT(base(sig8, 8), 2) + sgArrays(t, "DWT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"IDWT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + pc, cerr := DWT(probe(sig8, 8), 2) + bc, bcerr := DWT(base(sig8, 8), 2) + if cerr != nil || bcerr != nil { + sgArrays(t, "IDWT setup", dt, nil, cerr, nil, bcerr) + return + } + p, perr := IDWT(pc, 2) + b, berr := IDWT(bc, 2) + sgArrays(t, "IDWT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"CWT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := CWT(probe(sig8, 8), Morlet, []float64{2, 4}, 1.0) + b, berr := CWT(base(sig8, 8), Morlet, []float64{2, 4}, 1.0) + sgArrays(t, "CWT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"DaubechiesDWT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := DaubechiesDWT(probe(sig8, 8), DB2, 1, DWTPeriodic) + b, berr := DaubechiesDWT(base(sig8, 8), DB2, 1, DWTPeriodic) + sgArrays(t, "DaubechiesDWT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"DaubechiesIDWT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + pc, cerr := DaubechiesDWT(probe(sig8, 8), DB2, 1, DWTPeriodic) + bc, bcerr := DaubechiesDWT(base(sig8, 8), DB2, 1, DWTPeriodic) + if cerr != nil || bcerr != nil { + sgArrays(t, "DaubechiesIDWT setup", dt, nil, cerr, nil, bcerr) + return + } + p, perr := DaubechiesIDWT(pc, DB2, 1, DWTPeriodic) + b, berr := DaubechiesIDWT(bc, DB2, 1, DWTPeriodic) + sgArrays(t, "DaubechiesIDWT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SavitzkyGolay", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := SavitzkyGolay(probe(sig8, 8), 5, 2) + b, berr := SavitzkyGolay(base(sig8, 8), 5, 2) + sgArrays(t, "SavitzkyGolay", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"FilterApply", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := FilterApply([]float64{0.5}, []float64{1}, probe(sig8, 8)) + b, berr := FilterApply([]float64{0.5}, []float64{1}, base(sig8, 8)) + sgArrays(t, "FilterApply", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Filtfilt", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Filtfilt([]float64{0.5, 0.5}, []float64{1, -0.25}, probe(sig16, 16)) + b, berr := Filtfilt([]float64{0.5, 0.5}, []float64{1, -0.25}, base(sig16, 16)) + sgArrays(t, "Filtfilt", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Autocorrelate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Autocorrelate(probe(sig8, 8), 3) + b, berr := Autocorrelate(base(sig8, 8), 3) + sgArrays(t, "Autocorrelate", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"CrossCorrelate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := CrossCorrelate(probe(sig8, 8), probe([]float64{1, 2, 3, 4, 5, 6, 7, 8}, 8)) + b, berr := CrossCorrelate(base(sig8, 8), base([]float64{1, 2, 3, 4, 5, 6, 7, 8}, 8)) + sgArrays(t, "CrossCorrelate", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"PartialAutocorrelate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := PartialAutocorrelate(probe(sig16, 16), 3) + b, berr := PartialAutocorrelate(base(sig16, 16), 3) + sgArrays(t, "PartialAutocorrelate", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MedianFilter", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := MedianFilter(probe(sig8, 8), 3) + b, berr := MedianFilter(base(sig8, 8), 3) + sgArrays(t, "MedianFilter", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"RankFilter", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := RankFilter(probe(sig8, 8), 3, 1) + b, berr := RankFilter(base(sig8, 8), 3, 1) + sgArrays(t, "RankFilter", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MedianFilter2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := MedianFilter2D(probe(img16, 4, 4), 3) + b, berr := MedianFilter2D(base(img16, 4, 4), 3) + sgArrays(t, "MedianFilter2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"RankFilter2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := RankFilter2D(probe(img16, 4, 4), 3, 2) + b, berr := RankFilter2D(base(img16, 4, 4), 3, 2) + sgArrays(t, "RankFilter2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Conv1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Conv1D(probe(sig8, 1, 1, 8), probe([]float64{1, 0, 1, 0, 1, 0}, 2, 1, 3), probe([]float64{0, 0}, 2), 1, 1, 1) + b, berr := Conv1D(base(sig8, 1, 1, 8), base([]float64{1, 0, 1, 0, 1, 0}, 2, 1, 3), base([]float64{0, 0}, 2), 1, 1, 1) + sgArrays(t, "Conv1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Conv2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + in := probe(img16[:8], 1, 1, 2, 4) + ker := probe([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 3, 3) + bias := probe([]float64{0, 0}, 2) + p, perr := Conv2D(in, ker, bias, 1, 1) + b, berr := Conv2D(base(img16[:8], 1, 1, 2, 4), base([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 3, 3), base([]float64{0, 0}, 2), 1, 1) + sgArrays(t, "Conv2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Conv2DGroups", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Conv2DGroups(probe(img16[:8], 1, 1, 2, 4), probe([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 3, 3), probe([]float64{0, 0}, 2), 1, 1, 1) + b, berr := Conv2DGroups(base(img16[:8], 1, 1, 2, 4), base([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 3, 3), base([]float64{0, 0}, 2), 1, 1, 1) + sgArrays(t, "Conv2DGroups", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Conv3D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Conv3D(probe(sig8, 1, 1, 2, 2, 2), probe([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 2, 2, 2), probe([]float64{0, 0}, 2), 1, [3]int{0, 0, 0}, [3]int{1, 1, 1}) + b, berr := Conv3D(base(sig8, 1, 1, 2, 2, 2), base([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 2, 2, 2), base([]float64{0, 0}, 2), 1, [3]int{0, 0, 0}, [3]int{1, 1, 1}) + sgArrays(t, "Conv3D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"ConvTranspose2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := ConvTranspose2D(probe(img16[:9], 1, 1, 3, 3), probe([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 3, 3), probe([]float64{0, 0}, 2), 1, 1) + b, berr := ConvTranspose2D(base(img16[:9], 1, 1, 3, 3), base([]float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0}, 2, 1, 3, 3), base([]float64{0, 0}, 2), 1, 1) + sgArrays(t, "ConvTranspose2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MaxPool2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := MaxPool2D(probe(img16, 1, 1, 4, 4), 2, 2, 0) + b, berr := MaxPool2D(base(img16, 1, 1, 4, 4), 2, 2, 0) + sgArrays(t, "MaxPool2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AvgPool2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AvgPool2D(probe(img16, 1, 1, 4, 4), 2, 2, 0, false) + b, berr := AvgPool2D(base(img16, 1, 1, 4, 4), 2, 2, 0, false) + sgArrays(t, "AvgPool2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MaxPool1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := MaxPool1D(probe(sig8, 1, 1, 8), 2, 2, 0) + b, berr := MaxPool1D(base(sig8, 1, 1, 8), 2, 2, 0) + sgArrays(t, "MaxPool1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AvgPool1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AvgPool1D(probe(sig8, 1, 1, 8), 2, 2, 0, false) + b, berr := AvgPool1D(base(sig8, 1, 1, 8), 2, 2, 0, false) + sgArrays(t, "AvgPool1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"MaxPool3D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := MaxPool3D(probe(sig8, 1, 1, 2, 2, 2), [3]int{2, 2, 2}, [3]int{2, 2, 2}, [3]int{0, 0, 0}) + b, berr := MaxPool3D(base(sig8, 1, 1, 2, 2, 2), [3]int{2, 2, 2}, [3]int{2, 2, 2}, [3]int{0, 0, 0}) + sgArrays(t, "MaxPool3D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AvgPool3D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AvgPool3D(probe(sig8, 1, 1, 2, 2, 2), [3]int{2, 2, 2}, [3]int{2, 2, 2}, [3]int{0, 0, 0}, false) + b, berr := AvgPool3D(base(sig8, 1, 1, 2, 2, 2), [3]int{2, 2, 2}, [3]int{2, 2, 2}, [3]int{0, 0, 0}, false) + sgArrays(t, "AvgPool3D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AdaptiveMaxPool2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AdaptiveMaxPool2D(probe(img16, 1, 1, 4, 4), 2, 2) + b, berr := AdaptiveMaxPool2D(base(img16, 1, 1, 4, 4), 2, 2) + sgArrays(t, "AdaptiveMaxPool2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AdaptiveAvgPool2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AdaptiveAvgPool2D(probe(img16, 1, 1, 4, 4), 2, 2) + b, berr := AdaptiveAvgPool2D(base(img16, 1, 1, 4, 4), 2, 2) + sgArrays(t, "AdaptiveAvgPool2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AdaptiveMaxPool1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AdaptiveMaxPool1D(probe(sig8, 1, 1, 8), 4) + b, berr := AdaptiveMaxPool1D(base(sig8, 1, 1, 8), 4) + sgArrays(t, "AdaptiveMaxPool1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AdaptiveAvgPool1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AdaptiveAvgPool1D(probe(sig8, 1, 1, 8), 4) + b, berr := AdaptiveAvgPool1D(base(sig8, 1, 1, 8), 4) + sgArrays(t, "AdaptiveAvgPool1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AdaptiveMaxPool3D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AdaptiveMaxPool3D(probe(sig8, 1, 1, 2, 2, 2), 1, 1, 1) + b, berr := AdaptiveMaxPool3D(base(sig8, 1, 1, 2, 2, 2), 1, 1, 1) + sgArrays(t, "AdaptiveMaxPool3D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AdaptiveAvgPool3D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AdaptiveAvgPool3D(probe(sig8, 1, 1, 2, 2, 2), 1, 1, 1) + b, berr := AdaptiveAvgPool3D(base(sig8, 1, 1, 2, 2, 2), 1, 1, 1) + sgArrays(t, "AdaptiveAvgPool3D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"GlobalAvgPool2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := GlobalAvgPool2D(probe(img16, 1, 1, 4, 4)) + b, berr := GlobalAvgPool2D(base(img16, 1, 1, 4, 4)) + sgArrays(t, "GlobalAvgPool2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"GlobalMaxPool2D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := GlobalMaxPool2D(probe(img16, 1, 1, 4, 4)) + b, berr := GlobalMaxPool2D(base(img16, 1, 1, 4, 4)) + sgArrays(t, "GlobalMaxPool2D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"GlobalAvgPool1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := GlobalAvgPool1D(probe(sig8, 1, 1, 8)) + b, berr := GlobalAvgPool1D(base(sig8, 1, 1, 8)) + sgArrays(t, "GlobalAvgPool1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"GlobalMaxPool1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := GlobalMaxPool1D(probe(sig8, 1, 1, 8)) + b, berr := GlobalMaxPool1D(base(sig8, 1, 1, 8)) + sgArrays(t, "GlobalMaxPool1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"GlobalAvgPool3D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := GlobalAvgPool3D(probe(sig8, 1, 1, 2, 2, 2)) + b, berr := GlobalAvgPool3D(base(sig8, 1, 1, 2, 2, 2)) + sgArrays(t, "GlobalAvgPool3D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"GlobalMaxPool3D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := GlobalMaxPool3D(probe(sig8, 1, 1, 2, 2, 2)) + b, berr := GlobalMaxPool3D(base(sig8, 1, 1, 2, 2, 2)) + sgArrays(t, "GlobalMaxPool3D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Gradient1D", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Gradient1D(probe(sig8, 8), 1.0) + b, berr := Gradient1D(base(sig8, 8), 1.0) + sgArrays(t, "Gradient1D", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Laplacian", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Laplacian(probe(sig8, 8), 1.0) + b, berr := Laplacian(base(sig8, 8), 1.0) + sgArrays(t, "Laplacian", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SumKahan", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := SumKahan(probe(sig8, 8)) + b, berr := SumKahan(base(sig8, 8)) + sgFloat(t, "SumKahan", dt, p, perr, b, berr) + }}, + {"STFT", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + opts := STFTOptions{Segment: 8, Overlap: 4, Window: "hann"} + p, perr := STFT(probe(sig16, 16), opts) + b, berr := STFT(base(sig16, 16), opts) + sgArrays(t, "STFT", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Spectrogram", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + opts := STFTOptions{Segment: 8, Overlap: 4, Window: "hann"} + p, perr := Spectrogram(probe(sig16, 16), 1.0, opts) + b, berr := Spectrogram(base(sig16, 16), 1.0, opts) + sgArrays(t, "Spectrogram", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"WelchPSD", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + pf, pp, perr := WelchPSD(probe(sig16, 16), 1.0, 8, 4, "hann") + bf, bp, berr := WelchPSD(base(sig16, 16), 1.0, 8, 4, "hann") + sgArrays(t, "WelchPSD", dt, []*core.Array{pf, pp}, perr, []*core.Array{bf, bp}, berr) + }}, + {"LombScargle", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + times := []float64{1, 2, 3, 4, 5, 6} + vals := []float64{2, 1, 4, 3, 6, 5} + pf, pp, perr := LombScargle(probe(times, 6), probe(vals, 6), 0.1, 2.0, 8) + bf, bp, berr := LombScargle(base(times, 6), base(vals, 6), 0.1, 2.0, 8) + sgArrays(t, "LombScargle", dt, []*core.Array{pf, pp}, perr, []*core.Array{bf, bp}, berr) + }}, + {"EstimateAR", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + pr, perr := EstimateAR(probe(sig16, 16), 2) + br, berr := EstimateAR(base(sig16, 16), 2) + if berr != nil || perr != nil { + sgFloat(t, "EstimateAR", dt, 0, perr, 0, berr) + return + } + sgFloats(t, "EstimateAR AR", dt, pr.AR, nil, br.AR, nil) + sgFloat(t, "EstimateAR innovation variance", dt, pr.InnovationVariance, nil, br.InnovationVariance, nil) + // The spectrum of the estimate chains through unchanged. + pf, pp, perr := ARMASpectrum(pr, 8) + bf, bp, berr := ARMASpectrum(br, 8) + sgArrays(t, "ARMASpectrum", dt, []*core.Array{pf, pp}, perr, []*core.Array{bf, bp}, berr) + }}, + {"EstimateARMA", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + pr, perr := EstimateARMA(probe(sig16, 16), 2, 1, ARMAOptions{}) + br, berr := EstimateARMA(base(sig16, 16), 2, 1, ARMAOptions{}) + if berr != nil || perr != nil { + sgFloat(t, "EstimateARMA", dt, 0, perr, 0, berr) + return + } + sgFloats(t, "EstimateARMA AR", dt, pr.AR, nil, br.AR, nil) + sgFloats(t, "EstimateARMA MA", dt, pr.MA, nil, br.MA, nil) + }}, + {"SelectARMA", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + pr, perr := SelectARMA(probe(sig16, 16), 2, 2, ARMAOptions{}) + br, berr := SelectARMA(base(sig16, 16), 2, 2, ARMAOptions{}) + if berr != nil || perr != nil { + sgFloat(t, "SelectARMA", dt, 0, perr, 0, berr) + return + } + sgFloats(t, "SelectARMA AR", dt, pr.AR, nil, br.AR, nil) + sgFloats(t, "SelectARMA MA", dt, pr.MA, nil, br.MA, nil) + }}, + {"Resample", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Resample(probe(sig8, 8), 2, 1, 5) + b, berr := Resample(base(sig8, 8), 2, 1, 5) + sgArrays(t, "Resample", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Decimate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Decimate(probe(sig16, 16), 2, 5) + b, berr := Decimate(base(sig16, 16), 2, 5) + sgArrays(t, "Decimate", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"ResampleFourier", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := ResampleFourier(probe(sig8, 8), 16) + b, berr := ResampleFourier(base(sig8, 8), 16) + sgArrays(t, "ResampleFourier", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"NUFFTType1", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + coords := []float64{0.5, 1.5, 2.5, 3.5} + vals := []float64{1, 2, 3, 4} + p, perr := NUFFTType1(probe(coords, 4), probe(vals, 4), 16) + b, berr := NUFFTType1(base(coords, 4), base(vals, 4), 16) + sgArrays(t, "NUFFTType1", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"AnalyticSignal", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := AnalyticSignal(probe(sig8, 8)) + b, berr := AnalyticSignal(base(sig8, 8)) + sgArrays(t, "AnalyticSignal", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Envelope", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + p, perr := Envelope(probe(sig8, 8)) + b, berr := Envelope(base(sig8, 8)) + sgArrays(t, "Envelope", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"KalmanFilter", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + z := []float64{1, 2, 3, 4} + pr, perr := KalmanFilter(probe(z, 4, 1), probe([]float64{1}, 1, 1), probe([]float64{1}, 1, 1), KalmanOptions{}) + br, berr := KalmanFilter(base(z, 4, 1), base([]float64{1}, 1, 1), base([]float64{1}, 1, 1), KalmanOptions{}) + if berr != nil || perr != nil { + sgFloat(t, "KalmanFilter", dt, 0, perr, 0, berr) + return + } + sgArrays(t, "KalmanFilter states", dt, []*core.Array{pr.States, pr.Innovations}, nil, + []*core.Array{br.States, br.Innovations}, nil) + sgFloat(t, "KalmanFilter log likelihood", dt, pr.LogLikelihood, nil, br.LogLikelihood, nil) + }}, + {"ExtendedKalmanFilter", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + z := []float64{1, 2, 3, 4} + identity := func(v *core.Array) (*core.Array, error) { return core.AddF(v, 0), nil } + popts := KalmanOptions{InitialState: probe([]float64{0}, 1), InitialCovariance: probe([]float64{1}, 1, 1)} + bopts := KalmanOptions{InitialState: base([]float64{0}, 1), InitialCovariance: base([]float64{1}, 1, 1)} + pr, perr := ExtendedKalmanFilter(probe(z, 4, 1), identity, identity, popts) + br, berr := ExtendedKalmanFilter(base(z, 4, 1), identity, identity, bopts) + if berr != nil || perr != nil { + sgFloat(t, "ExtendedKalmanFilter", dt, 0, perr, 0, berr) + return + } + sgArrays(t, "ExtendedKalmanFilter states", dt, []*core.Array{pr.States}, nil, []*core.Array{br.States}, nil) + sgFloat(t, "ExtendedKalmanFilter log likelihood", dt, pr.LogLikelihood, nil, br.LogLikelihood, nil) + }}, + {"UnscentedKalmanFilter", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + z := []float64{1, 2, 3, 4} + identity := func(v *core.Array) (*core.Array, error) { return core.AddF(v, 0), nil } + popts := KalmanOptions{InitialState: probe([]float64{0}, 1), InitialCovariance: probe([]float64{1}, 1, 1)} + bopts := KalmanOptions{InitialState: base([]float64{0}, 1), InitialCovariance: base([]float64{1}, 1, 1)} + pr, perr := UnscentedKalmanFilter(probe(z, 4, 1), identity, identity, popts) + br, berr := UnscentedKalmanFilter(base(z, 4, 1), identity, identity, bopts) + if berr != nil || perr != nil { + sgFloat(t, "UnscentedKalmanFilter", dt, 0, perr, 0, berr) + return + } + sgArrays(t, "UnscentedKalmanFilter states", dt, []*core.Array{pr.States}, nil, []*core.Array{br.States}, nil) + sgFloat(t, "UnscentedKalmanFilter log likelihood", dt, pr.LogLikelihood, nil, br.LogLikelihood, nil) + }}, + // The spectral Poisson family computes in float64 end to end + // and refuses every non-float dtype by name, Int included: the + // narrow widths and bool follow Int into the same refusal, the + // existing wording preserved. + {"SolvePoissonPeriodic dtype gate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + _, perr := SolvePoissonPeriodic(probe(img16, 4, 4), 1.0, 1.0) + sgWantErr(t, "SolvePoissonPeriodic/"+dt.String(), perr, + "SolvePoissonPeriodic", "f must be a float64 array, got", dt.String()) + }}, + {"SolvePoissonDirichlet dtype gate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + _, perr := SolvePoissonDirichlet(probe(img16, 4, 4), 1.0, 1.0) + sgWantErr(t, "SolvePoissonDirichlet/"+dt.String(), perr, + "SolvePoissonDirichlet", "f must be a float64 array, got", dt.String()) + }}, + {"SolvePoissonNeumann dtype gate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + _, perr := SolvePoissonNeumann(probe(img16, 4, 4), 1.0, 1.0) + sgWantErr(t, "SolvePoissonNeumann/"+dt.String(), perr, + "SolvePoissonNeumann", "f must be a float64 array, got", dt.String()) + }}, + // IRFFT keeps its standing complex-input gate, wording verbatim + // from fft.go: every real dtype, Int included, is refused by + // name; on the complex spectrum RFFT answers for the same + // values, it computes identically whatever dtype reached RFFT. + {"IRFFT real input gate", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + _, perr := IRFFT(probe(sig8, 8), 8) + sgWantErr(t, "IRFFT/"+dt.String(), perr, "IRFFT", "input must be complex") + _, berr := IRFFT(base(sig8, 8), 8) + sgWantErr(t, "IRFFT float baseline", berr, "IRFFT", "input must be complex") + }}, + {"IRFFT complex spectrum computes", func(t *testing.T, probe, base sgMaker, dt core.Dtype) { + ps, serr := RFFT(probe(sig8, 8)) + if serr != nil { + t.Fatalf("RFFT probe setup (%s): %v", dt, serr) + } + bs, berr := RFFT(base(sig8, 8)) + if berr != nil { + t.Fatalf("RFFT baseline setup: %v", berr) + } + p, perr := IRFFT(ps, 8) + b, ierr := IRFFT(bs, 8) + sgArrays(t, "IRFFT of RFFT", dt, []*core.Array{p}, perr, []*core.Array{b}, ierr) + }}, + } + for _, row := range rows { + for _, dt := range sgDtypes { + t.Run(row.name+"/"+dt.String(), func(t *testing.T) { + probe, base := sgMakers(t, dt) + row.run(t, probe, base, dt) + }) + } + } +} + +// TestDtypesCensusSignalPoissonFloatGate pins that a float64 grid +// passes the Poisson dtype gate the narrow widths fail: the refusal is +// about the dtype, not the values. +func TestDtypesCensusSignalPoissonFloatGate(t *testing.T) { + grid, err := core.FromFloats([]float64{ + 1, 2, -1, -2, + -1, -2, 1, 2, + 2, 1, -2, -1, + -2, -1, 2, 1, + }, 4, 4) + if err != nil { + t.Fatal(err) + } + _, perr := SolvePoissonPeriodic(grid, 2*3.141592653589793, 2*3.141592653589793) + if perr != nil && strings.Contains(perr.Error(), "must be a float64 array") { + t.Fatalf("float64 grid hit the dtype gate: %v", perr) + } +} diff --git a/signal/entry_gate_pins_test.go b/signal/entry_gate_pins_test.go new file mode 100644 index 0000000..98e5e47 --- /dev/null +++ b/signal/entry_gate_pins_test.go @@ -0,0 +1,227 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// mustF builds a float array or fails the test. +func mustF(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 +} + +// TestFFTRejectsComplexMatrices pins the rank gate on the +// complex path, which used to flatten a matrix silently where the +// real path refused. +func TestFFTRejectsComplexMatrices(t *testing.T) { + c, err := core.FromComplexes([]complex128{1, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + if _, err := FFT(c); err == nil { + t.Fatal("FFT: expected a rank error for a complex matrix") + } + if _, err := IFFT(c); err == nil { + t.Fatal("IFFT: expected a rank error for a complex matrix") + } +} + +// TestNUFFTRejectsNaNCoordinate pins the named refusal: NaN +// defeats both range comparisons and used to poison the whole grid. +func TestNUFFTRejectsNaNCoordinate(t *testing.T) { + x := mustF(t, []float64{math.NaN(), 0.1}, 2) + c := mustF(t, []float64{1, 1}, 2) + if _, err := NUFFTType1(x, c, 4); err == nil { + t.Fatal("NUFFTType1: expected an error for a NaN coordinate") + } +} + +// TestInfiniteParameterGates pins that a +Inf parameter cannot +// pass a positive-only gate and turn the answer into NaN or silent +// zeros. +func TestInfiniteParameterGates(t *testing.T) { + times := mustF(t, []float64{0.1, 0.4, 0.7, 1.0, 1.3}, 5) + vals := mustF(t, []float64{1, -0.5, 0.8, -0.2, 0.6}, 5) + if _, _, err := LombScargle(times, vals, math.Inf(1), math.Inf(1), 4); err == nil { + t.Fatal("LombScargle: expected an error for an infinite frequency range") + } + x := mustF(t, make([]float64, 16), 16) + for i := range 16 { + x.SetFloatAt(i, math.Sin(float64(i)/3)) + } + if _, err := CWT(x, Morlet, []float64{1}, math.Inf(1)); err == nil { + t.Fatal("CWT: expected an error for an infinite sample spacing") + } + if _, err := CWT(x, Morlet, []float64{math.Inf(1)}, 0.1); err == nil { + t.Fatal("CWT: expected an error for an infinite scale") + } + if _, _, err := WelchPSD(x, math.Inf(1), 8, 0, "hann"); err == nil { + t.Fatal("WelchPSD: expected an error for an infinite fs") + } + if _, err := Spectrogram(x, math.Inf(1), STFTOptions{Segment: 8}); err == nil { + t.Fatal("Spectrogram: expected an error for an infinite fs") + } + if _, _, err := ButterworthLowPass(2, math.Inf(1), 100); err == nil { + t.Fatal("ButterworthLowPass: expected an error for an infinite fs") + } + if _, _, err := ChebyshevLowPass(2, math.Inf(1), 100, 1); err == nil { + t.Fatal("ChebyshevLowPass: expected an error for an infinite fs") + } +} + +// TestPoissonDirichletMinimumGrid pins the documented 3×3 +// minimum, whose length-one interior used to hit the DST-I floor. +func TestPoissonDirichletMinimumGrid(t *testing.T) { + f := mustF(t, []float64{0, 0, 0, 0, 1, 0, 0, 0, 0}, 3, 3) + u, err := SolvePoissonDirichlet(f, 1, 1) + if err != nil { + t.Fatalf("SolvePoissonDirichlet on the 3×3 minimum: %v", err) + } + // The single interior unknown solves −(4u)/h² = 1 at h = 1/2: + // u = -1/16... the sign follows the convention −Δu = f, so the + // interior value is f/16 with the operator's sign. + if v := u.FloatAt(4); math.IsNaN(v) || math.IsInf(v, 0) { + t.Fatalf("interior value = %v", v) + } + // (3, k) and (k, 3) grids take the one-point DST on one axis only. + f2 := mustF(t, []float64{0, 0, 0, 0, 0, 0, 0, 2, 0, 0, 0, 0}, 3, 4) + if _, err := SolvePoissonDirichlet(f2, 1, 1); err != nil { + t.Fatalf("SolvePoissonDirichlet on 3×4: %v", err) + } +} + +// TestMaxPoolNaNEveryRank pins the NaN-propagation promise on +// the four max-pool entry points the earlier pin left uncovered. +func TestMaxPoolNaNEveryRank(t *testing.T) { + x1 := mustF(t, []float64{1, math.NaN(), 3, 4, 5, 6}, 1, 1, 6) + got1, err := MaxPool1D(x1, 2, 1, 0) + if err != nil { + t.Fatalf("MaxPool1D: %v", err) + } + if !math.IsNaN(got1.FloatAt(0)) { + t.Fatalf("MaxPool1D[0] = %v, want NaN", got1.FloatAt(0)) + } + a1, err := AdaptiveMaxPool1D(x1, 2) + if err != nil { + t.Fatalf("AdaptiveMaxPool1D: %v", err) + } + if !math.IsNaN(a1.FloatAt(0)) { + t.Fatalf("AdaptiveMaxPool1D[0] = %v, want NaN", a1.FloatAt(0)) + } + g1, err := GlobalMaxPool1D(x1) + if err != nil { + t.Fatalf("GlobalMaxPool1D: %v", err) + } + if !math.IsNaN(g1.FloatAt(0)) { + t.Fatalf("GlobalMaxPool1D[0] = %v, want NaN", g1.FloatAt(0)) + } + x3 := make([]float64, 2*3*4*4*4) + for i := range x3 { + x3[i] = float64(i % 7) + } + x3[40] = math.NaN() + c3, err := core.FromFloats(x3, 2, 3, 4, 4, 4) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + g3, err := GlobalMaxPool3D(c3) + if err != nil { + t.Fatalf("GlobalMaxPool3D: %v", err) + } + if !math.IsNaN(g3.FloatAt(0)) { + t.Fatalf("GlobalMaxPool3D[0] = %v, want NaN", g3.FloatAt(0)) + } + m3, err := MaxPool3D(c3, [3]int{2, 2, 2}, [3]int{1, 1, 1}, [3]int{0, 0, 0}) + if err != nil { + t.Fatalf("MaxPool3D: %v", err) + } + // Flat 40 is (d, h, w) = (2, 2, 0); the window that covers it + // starts at (1, 1, 0), output index 12 of the 3×3×3 grid. + if !math.IsNaN(m3.FloatAt(12)) { + t.Fatalf("MaxPool3D[12] = %v, want NaN", m3.FloatAt(12)) + } + a3, err := AdaptiveMaxPool3D(c3, 1, 1, 1) + if err != nil { + t.Fatalf("AdaptiveMaxPool3D: %v", err) + } + if !math.IsNaN(a3.FloatAt(0)) { + t.Fatalf("AdaptiveMaxPool3D[0] = %v, want NaN", a3.FloatAt(0)) + } +} + +// TestGlobalPoolsRejectEmptySpatial pins the named refusal the +// adaptive ranks already make. +func TestGlobalPoolsRejectEmptySpatial(t *testing.T) { + empty1, err := core.Zeros(core.Float, 2, 3, 0) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + if _, err := GlobalMaxPool1D(empty1); err == nil { + t.Fatal("GlobalMaxPool1D: expected an error for an empty spatial dimension") + } + if _, err := GlobalAvgPool1D(empty1); err == nil { + t.Fatal("GlobalAvgPool1D: expected an error for an empty spatial dimension") + } + empty3, err := core.Zeros(core.Float, 2, 3, 0, 4, 4) + if err != nil { + t.Fatalf("Zeros: %v", err) + } + if _, err := GlobalMaxPool3D(empty3); err == nil { + t.Fatal("GlobalMaxPool3D: expected an error for an empty spatial dimension") + } + if _, err := GlobalAvgPool3D(empty3); err == nil { + t.Fatal("GlobalAvgPool3D: expected an error for an empty spatial dimension") + } +} + +// TestResampleMatchesDirectSum pins the polyphase window +// restriction: the bounded loop must produce bit-identical output to +// the full scan it replaced. +func TestResampleMatchesDirectSum(t *testing.T) { + n := 61 + src := make([]float64, n) + for i := range n { + src[i] = math.Sin(float64(i)*0.37) + 0.2*math.Cos(float64(i)*1.7) + } + x := mustF(t, src, n) + for _, c := range []struct{ up, down int }{{2, 3}, {3, 2}, {1, 2}, {2, 1}, {5, 7}} { + got, err := Resample(x, c.up, c.down, 13) + if err != nil { + t.Fatalf("Resample(%d/%d): %v", c.up, c.down, err) + } + // The reference: the original full scan over every input + // index, computing the same taps by hand. + taps := 13 + fc := 0.5 * math.Min(1/float64(c.up), 1/float64(c.down)) + h := kaiserSinc(taps, fc) + delay := (taps - 1) / 2 + outLen := (n*c.up + c.down - 1) / c.down + if got.Len() != outLen { + t.Fatalf("Resample(%d/%d): length %d, want %d", c.up, c.down, got.Len(), outLen) + } + for m := range outLen { + centre := delay + m*c.down + total := 0.0 + for k := range n { + j := centre - k*c.up + if j < 0 || j >= taps { + continue + } + total += h[j] * src[k] + } + if want := float64(c.up) * total; got.FloatAt(m) != want { + t.Fatalf("Resample(%d/%d)[%d] = %g, want %g", c.up, c.down, m, got.FloatAt(m), want) + } + } + } +} diff --git a/signal/example_test.go b/signal/example_test.go new file mode 100644 index 0000000..0029fca --- /dev/null +++ b/signal/example_test.go @@ -0,0 +1,195 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal_test + +// Runnable godoc examples for the flagship workflows. Every one of them +// runs on a fixed input and prints a fixed result, so `go test` checks +// the documentation against the code. + +import ( + "fmt" + "log" + "math" + + tensor "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/signal" +) + +// A two-tone signal at 4 Hz and 12 Hz, sampled at 64 Hz: Welch's +// averaged, windowed estimate puts both tones on their own bins. +func ExampleWelchPSD() { + const n, fs = 64, 64.0 + raw := make([]float64, n) + for i := range raw { + t := float64(i) / fs + raw[i] = math.Cos(2*math.Pi*4*t) + 0.5*math.Cos(2*math.Pi*12*t) + } + x, err := tensor.FromFloats(raw, n) + if err != nil { + log.Fatal(err) + } + freqs, psd, err := signal.WelchPSD(x, fs, 32, 16, "hann") + if err != nil { + log.Fatal(err) + } + peak := 0 + for k := range psd.Len() { + if psd.FloatAt(k) > psd.FloatAt(peak) { + peak = k + } + } + fmt.Printf("peak %.0f Hz, %.3f power, %d bins\n", + freqs.FloatAt(peak), psd.FloatAt(peak), psd.Len()) + // Output: peak 4 Hz, 0.167 power, 17 bins +} + +// A 5 Hz tone in a 20 Hz passband: the one-pass filter lags the tone, +// the forward-and-backward sweep leaves it where it was. +func ExampleFiltfilt() { + const n, fs, f = 200, 100.0, 5.0 + raw := make([]float64, n) + for i := range raw { + raw[i] = math.Sin(2 * math.Pi * f * float64(i) / fs) + } + x, err := tensor.FromFloats(raw, n) + if err != nil { + log.Fatal(err) + } + b, a, err := signal.ButterworthLowPass(2, fs, 20) + if err != nil { + log.Fatal(err) + } + once, err := signal.FilterApply(b, a, x) + if err != nil { + log.Fatal(err) + } + twice, err := signal.Filtfilt(b, a, x) + if err != nil { + log.Fatal(err) + } + peak := func(y *tensor.Array) int { + at := 0 + for i := range y.Len() { + if y.FloatAt(i) > y.FloatAt(at) { + at = i + } + } + return at + } + fmt.Printf("input peaks at %d, one pass at %d, filtfilt at %d\n", + peak(x), peak(once), peak(twice)) + // Output: input peaks at 5, one pass at 6, filtfilt at 5 +} + +// The DB2 transform packs the coefficients as [A, D2, D1]. The periodic +// boundary keeps the energy, so the inverse returns the signal. +func ExampleDaubechiesDWT() { + const levels = 2 + x, err := tensor.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8}, 8) + if err != nil { + log.Fatal(err) + } + coef, err := signal.DaubechiesDWT(x, signal.DB2, levels, signal.DWTPeriodic) + if err != nil { + log.Fatal(err) + } + back, err := signal.DaubechiesIDWT(coef, signal.DB2, levels, signal.DWTPeriodic) + if err != nil { + log.Fatal(err) + } + energy := func(v *tensor.Array) float64 { + total := 0.0 + for i := range v.Len() { + total += v.FloatAt(i) * v.FloatAt(i) + } + return total + } + worst := 0.0 + for i := range x.Len() { + worst = math.Max(worst, math.Abs(back.FloatAt(i)-x.FloatAt(i))) + } + fmt.Printf("%d levels, %d coefficients, round trip exact: %v\n", + levels, coef.Len(), worst < 1e-12) + fmt.Printf("energy kept to %.3f\n", energy(coef)/energy(x)) + // Output: 2 levels, 8 coefficients, round trip exact: true + // energy kept to 1.000 +} + +// The Hilbert envelope of a tone whose amplitude swings at 1 Hz: the +// modulus of the analytic signal traces the swing without smoothing lag. +func ExampleEnvelope() { + const n, fs = 256, 128.0 + raw := make([]float64, n) + for i := range raw { + t := float64(i) / fs + // A 20 Hz carrier under a 1 Hz amplitude swing. + raw[i] = (1 + 0.5*math.Cos(2*math.Pi*t)) * math.Cos(2*math.Pi*20*t) + } + x, err := tensor.FromFloats(raw, n) + if err != nil { + log.Fatal(err) + } + env, err := signal.Envelope(x) + if err != nil { + log.Fatal(err) + } + // The swing peaks at t = 0 s and t = 1 s, and bottoms at t = 0.5 s. + at := func(seconds float64) int { return int(seconds * fs) } + fmt.Printf("t=0 s %.3f, t=0.5 s %.3f, t=1 s %.3f\n", + env.FloatAt(at(0)), env.FloatAt(at(0.5)), env.FloatAt(at(1))) + // Output: t=0 s 1.500, t=0.5 s 0.500, t=1 s 1.500 +} + +// A constant position observed three times: the linear Kalman filter +// fuses each measurement into the running estimate and reports the +// filtered means, one row per measurement. +func ExampleKalmanFilter() { + z, err := tensor.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + log.Fatal(err) + } + f, err := tensor.FromFloats([]float64{1}, 1, 1) + if err != nil { + log.Fatal(err) + } + h, err := tensor.FromFloats([]float64{1}, 1, 1) + if err != nil { + log.Fatal(err) + } + res, err := signal.KalmanFilter(z, f, h, signal.KalmanOptions{}) + if err != nil { + log.Fatal(err) + } + fmt.Printf("states %.3f %.3f %.3f\n", res.States.FloatAt(0), res.States.FloatAt(1), res.States.FloatAt(2)) + fmt.Printf("log likelihood %.2f\n", res.LogLikelihood) + // Output: states 0.500 1.000 1.500 + // log likelihood -5.95 +} + +// An AR(1) series built from a fixed linear congruential driver: the +// Yule-Walker fit recovers the coefficient that generated it. +func ExampleEstimateAR() { + const n = 512 + raw := make([]float64, n) + seed := 1.0 + for i := range raw { + // The minimal-standard generator: deterministic, so the + // example never calls a random source. + seed = math.Mod(16807*seed, 2147483647) + raw[i] = seed/2147483647 - 0.5 + if i > 0 { + raw[i] += 0.5 * raw[i-1] + } + } + x, err := tensor.FromFloats(raw, n) + if err != nil { + log.Fatal(err) + } + res, err := signal.EstimateAR(x, 1) + if err != nil { + log.Fatal(err) + } + fmt.Printf("AR(%d): phi_1 %.3f\n", res.P, res.AR[0]) + // Output: AR(1): phi_1 0.495 +} diff --git a/signal/examples_test.go b/signal/examples_test.go new file mode 100644 index 0000000..a849b2f --- /dev/null +++ b/signal/examples_test.go @@ -0,0 +1,36 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "fmt" + "log" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +func ExampleConv2D() { + x, _ := core.FromFloats([]float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + }, 1, 1, 2, 4) + w, _ := core.FromFloats([]float64{1, 0, 0, 1}, 1, 1, 2, 2) + out, err := Conv2D(x, w, nil, 1, 0) + if err != nil { + log.Fatal(err) + } + fmt.Println(out) + // Output: float (1, 1, 1, 3) [7, 9, 11] +} + +func ExampleFFT() { + x, _ := core.FromFloats([]float64{1, 1, 1, 1}, 4) + spec, err := FFT(x) + if err != nil { + log.Fatal(err) + } + dc, _ := core.ComplexAt(spec, 0) + fmt.Println(real(dc)) + // Output: 4 +} diff --git a/signal/fft.go b/signal/fft.go new file mode 100644 index 0000000..a7ac937 --- /dev/null +++ b/signal/fft.go @@ -0,0 +1,874 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "sourcedock.dev/petrbalvin/tensor/internal/engine" + +import ( + "math" + "math/bits" + "sync" +) + +// Discrete Fourier transforms. FFT returns the forward transform +// and IFFT the inverse, both as complex arrays from any 1-D input; real +// elements convert. Powers of two run an iterative Cooley-Tukey whose +// stages fuse into radix-4 passes above a small floor (a lone radix-2 +// stage remains when the stage count is even); every other length runs +// Bluestein's chirp-z transform over a padded power of two, so no +// length is rejected. + +// FFT returns the forward discrete Fourier transform of a 1-D array. An +// empty array is an error, and so is any array above rank one: the +// complex payload shortcut would otherwise flatten a matrix silently +// where the real path refuses. The input is never modified: arrays are +// immutable values, so the transform runs on a private copy. +func FFT(a *core.Array) (*core.Array, error) { + if a.NDim() != 1 { + return nil, base.Errf("FFT: needs a 1-D array, got shape %s", base.ShapeText(a.Shape())) + } + if a.Len() == 0 { + return nil, base.Errf("FFT: an empty array has no FFT") + } + vals, err := a.ComplexValues("FFT") + if err != nil { + return nil, err + } + // ComplexValues shares the payload of a dense complex array and hands + // back a private slice for every other input. The transform works in + // place, so only the shared read needs the copy. + if raw := a.RawComplexes(); len(raw) > 0 && &vals[0] == &raw[0] { + vals = append([]complex128(nil), vals...) + } + transform(vals, -1) + return complexFromArrayMust(vals, []int{len(vals)}), nil +} + +// IFFT returns the inverse discrete Fourier transform of a 1-D array, +// scaled by 1/n. An empty array is an error and so is any array above +// rank one. Like FFT it leaves its input untouched. +func IFFT(a *core.Array) (*core.Array, error) { + if a.NDim() != 1 { + return nil, base.Errf("IFFT: needs a 1-D array, got shape %s", base.ShapeText(a.Shape())) + } + if a.Len() == 0 { + return nil, base.Errf("IFFT: an empty array has no FFT") + } + vals, err := a.ComplexValues("IFFT") + if err != nil { + return nil, err + } + // The same alias rule as FFT: only a shared payload is copied, the + // scaling below runs in place. + if raw := a.RawComplexes(); len(raw) > 0 && &vals[0] == &raw[0] { + vals = append([]complex128(nil), vals...) + } + transform(vals, +1) + n := float64(len(vals)) + for i := range vals { + vals[i] /= complex(n, 0) + } + return complexFromArrayMust(vals, []int{len(vals)}), nil +} + +// transform computes the DFT in place: sign -1 is the forward transform, +// +1 the inverse without scaling. +func transform(vals []complex128, sign float64) { + if isPowerOfTwo(len(vals)) { + fftPow2(vals, sign) + return + } + bluestein(vals, sign) +} + +// radix4MinN bounds the transform size below which the plain radix-2 +// walk stays: the fused stages earn nothing on a transform that fits +// the first-level cache a few times over, and the small sizes are too +// few butterflies to repay the wider stage code. +const radix4MinN = 1 << 6 + +// fftPow2 dispatches a power-of-two length between the radix-2 walk +// for small transforms and the fused radix-4 stages above the floor. +func fftPow2(vals []complex128, sign float64) { + if len(vals) < radix4MinN { + fftRadix2(vals, sign) + return + } + fftRadix4(vals, sign) +} + +// twiddleCacheMax bounds the transform sizes whose twiddle tables stay +// cached between calls (1 MiB of tables per key at the cap). Bigger +// sizes still build their tables, they just do not keep them. +const twiddleCacheMax = 1 << 16 + +// twiddle tables are keyed by size and direction sign. A row holds the +// twiddle factors of one butterfly stage: entry k is step multiplied +// into itself k times starting from 1, exactly the value the on-the-fly +// recurrence w *= step produced in butterfly k. Reading the table +// therefore cannot change a single bit of the result; it only spares +// recomputing the recurrence in every parallel block. +type twiddleKey struct { + sign float64 + n int +} + +var ( + twiddleMu sync.RWMutex + twiddleTables = map[twiddleKey][][]complex128{} +) + +// twiddlesFor returns the per-stage twiddle rows for a size-n power-of-two +// transform with the given direction sign, building them on first use. +func twiddlesFor(sign float64, n int) [][]complex128 { + key := twiddleKey{sign: sign, n: n} + twiddleMu.RLock() + tables, ok := twiddleTables[key] + twiddleMu.RUnlock() + if ok { + return tables + } + tables = make([][]complex128, bits.Len(uint(n))-1) + for length := 2; length <= n; length <<= 1 { + angle := sign * 2 * math.Pi / float64(length) + sinStep, cosStep := math.Sincos(angle) + step := complex(cosStep, sinStep) + half := length / 2 + row := make([]complex128, half) + w := complex(1, 0) + for k := range half { + row[k] = w + w *= step + } + tables[bits.TrailingZeros(uint(length))-1] = row + } + if n <= twiddleCacheMax { + twiddleMu.Lock() + twiddleTables[key] = tables + twiddleMu.Unlock() + } + return tables +} + +// bit-reversal swap tables are keyed by size alone: the permutation +// does not depend on the direction sign. A table holds the flat pairs +// (i, j) the in-place reversal loop swaps, i < j, each pair exactly +// once. Applying the pairs reproduces the loop's memory state exactly, +// so reading the table cannot change a single bit of the result. +var ( + bitReversalMu sync.RWMutex + bitReversalTables = map[int][]int{} +) + +// bitReversalFor returns the flat swap-pair table for the bit-reversal +// permutation of a size-n power-of-two transform, building it on first +// use. Pairs are recorded from the same incremental reversal loop the +// in-place pass ran, so the table is exactly the set of swaps that loop +// performs. Pairs are disjoint, which keeps their application order +// irrelevant. The twiddleCacheMax policy applies: bigger sizes still +// build their tables, they just do not keep them. +func bitReversalFor(n int) []int { + bitReversalMu.RLock() + swaps, ok := bitReversalTables[n] + bitReversalMu.RUnlock() + if ok { + return swaps + } + swaps = make([]int, 0, n) + for i, j := 1, 0; i < n; i++ { + bit := n >> 1 + for ; j&bit != 0; bit >>= 1 { + j ^= bit + } + j |= bit + if i < j { + swaps = append(swaps, i, j) + } + } + if n <= twiddleCacheMax { + bitReversalMu.Lock() + bitReversalTables[n] = swaps + bitReversalMu.Unlock() + } + return swaps +} + +// fftRadix2 runs iterative Cooley-Tukey with bit reversal. +func fftRadix2(vals []complex128, sign float64) { + n := len(vals) + if n < 2 { + return + } + twiddles := twiddlesFor(sign, n) + // Bit-reversal permutation from the cached swap table. + swaps := bitReversalFor(n) + for p := 0; p < len(swaps); p += 2 { + i, j := swaps[p], swaps[p+1] + vals[i], vals[j] = vals[j], vals[i] + } + // Butterflies over doubling block sizes; every block of a stage + // reads the same cached twiddle row. The blocks of one stage are + // mutually independent, so a wide stage of a large transform splits + // its block range across workers; disjoint blocks touch disjoint + // indices, so the split cannot change the result. Narrow stages and + // small transforms stay on the calling goroutine, where the spawn + // cost would not amortise. + for length := 2; length <= n; length <<= 1 { + row := twiddles[bits.TrailingZeros(uint(length))-1] + blocks := n / length + if n >= parallelMinN && length >= parallelMinBlock && blocks > 1 { + engine.Parallel(blocks, func(bs, be int) { + for b := bs; b < be; b++ { + butterflyStage(vals[b*length:(b+1)*length], row) + } + }) + continue + } + butterflyStage(vals, row) + } +} + +// parallelMinN bounds the transform size whose stages may split across +// workers, parallelMinBlock the stage width worth splitting. At smaller +// sizes a stage does too little work to repay spawning the workers. +// The floor is low enough that a padded correlation or convolution +// (2·n for an n-point signal) still parallelises its widest stages. +const ( + parallelMinN = 1 << 14 + parallelMinBlock = 1 << 11 +) + +// butterflyStage runs one in-place butterfly stage over the whole of +// vals: with half = len(row), every block of 2·half entries combines +// vals[start+k] with vals[start+k+half] through row[k]. It is the +// arithmetic of the original serial stage, unchanged. The two halves +// of a block are cut out as disjoint sub-slices so the tight loop +// indexes without bounds checks. +func butterflyStage(vals, row []complex128) { + half := len(row) + if half == 1 { + // The first stage pairs every even index with its odd + // successor: one flat pass pays no per-block slice setup on + // a single butterfly, which is all a block holds here. + w := row[0] + for i := 1; i < len(vals); i += 2 { + u := vals[i-1] + v := vals[i] * w + vals[i-1] = u + v + vals[i] = u - v + } + return + } + if half == 2 { + // Two butterflies per block: the flat form still wins over + // a slice pair per 4 entries. + w0, w1 := row[0], row[1] + for start := 0; start+3 < len(vals); start += 4 { + u := vals[start] + v := vals[start+2] * w0 + vals[start] = u + v + vals[start+2] = u - v + u = vals[start+1] + v = vals[start+3] * w1 + vals[start+1] = u + v + vals[start+3] = u - v + } + return + } + for start := 0; start < len(vals); start += 2 * half { + lo := vals[start : start+half : start+half] + hi := vals[start+half : start+2*half : start+2*half] + for k := range half { + u := lo[k] + v := hi[k] * row[k] + lo[k] = u + v + hi[k] = u - v + } + } +} + +// radix4Tables caches the fused-stage twiddle rows by size and +// direction sign. One entry covers every fused stage of a size-n +// transform; a stage of block 4·L holds three rows of L entries, the +// powers W^k, W^2k and W^3k of the block twiddle W. Every entry is an +// individual sine/cosine evaluation, so each carries the single +// rounding of its argument instead of the k accumulated roundings a +// recurrence step would leave. The twiddleCacheMax policy applies. +var ( + radix4Mu sync.RWMutex + radix4Tables = map[twiddleKey][][3][]complex128{} +) + +// radix4TwiddlesFor returns the per-stage twiddle rows of a size-n +// fused transform, one [3][]complex128 per stage block 8, 32, 128, …, +// building them on first use. +func radix4TwiddlesFor(sign float64, n int) [][3][]complex128 { + key := twiddleKey{sign: sign, n: n} + radix4Mu.RLock() + tables, ok := radix4Tables[key] + radix4Mu.RUnlock() + if ok { + return tables + } + tables = nil + for length := 8; length <= n; length *= 4 { + l := length / 4 + var rows [3][]complex128 + for j := range rows { + row := make([]complex128, l) + for k := range row { + angle := sign * 2 * math.Pi * float64(k*(j+1)) / float64(length) + s, c := math.Sincos(angle) + row[k] = complex(c, s) + } + rows[j] = row + } + tables = append(tables, rows) + } + if n <= twiddleCacheMax { + radix4Mu.Lock() + radix4Tables[key] = tables + radix4Mu.Unlock() + } + return tables +} + +// fftRadix4 runs a power-of-two Cooley-Tukey whose stage pairs are +// fused into radix-4 passes: one sweep over the payload does the work +// of two radix-2 sweeps and folds one of the four twiddle multiplies +// a radix-2 pair would spend per butterfly. The input ordering is the +// bit-reversal permutation the radix-2 walk uses, so the cached swap +// tables are shared. The walk opens with the flat length-2 stage and, +// when the stage count is even, closes with a lone radix-2 stage. +func fftRadix4(vals []complex128, sign float64) { + n := len(vals) + swaps := bitReversalFor(n) + for p := 0; p < len(swaps); p += 2 { + i, j := swaps[p], swaps[p+1] + vals[i], vals[j] = vals[j], vals[i] + } + butterflyStage(vals, twiddlesFor(sign, 2)[0]) + tables := radix4TwiddlesFor(sign, n) + covered := 2 + for ti, length := 0, 8; length <= n; ti, length = ti+1, length*4 { + rows := tables[ti] + covered = length + blocks := n / length + // The blocks of one stage are mutually independent, so a wide + // stage of a large transform splits its block range across + // workers on the same policy the radix-2 walk applies. + if n >= parallelMinN && length >= parallelMinBlock && blocks > 1 { + engine.Parallel(blocks, func(bs, be int) { + for b := bs; b < be; b++ { + radix4Stage(vals[b*length:(b+1)*length], rows, -sign) + } + }) + continue + } + radix4Stage(vals, rows, -sign) + } + if covered < n { + butterflyStage(vals, twiddlesFor(sign, n)[bits.TrailingZeros(uint(n))-1]) + } +} + +// radix4Stage runs one fused radix-4 stage over vals, whose blocks of +// 4·L entries (L = len(tw1)) each hold four L-point transforms. Per +// index k the butterfly forms x0, x1·W^k, x2·W^2k, x3·W^3k over the +// four decimated subsequences and combines them with the exact +// rotation ∓i carrying the odd output slots, an exchange of the pair +// (re, im) with a sign rather than a multiply. q is that sign: +1 +// rotates the forward transform's branch by −i, −1 the inverse's by +// +i. The base-2 reversal lays the four transforms down in the order +// ≡0, ≡2, ≡1, ≡3, so the odd decimations read across the middle of +// the block while the outputs land in place. +func radix4Stage(vals []complex128, rows [3][]complex128, q float64) { + tw1, tw2, tw3 := rows[0], rows[1], rows[2] + l := len(tw1) + for start := 0; start < len(vals); start += 4 * l { + s0 := vals[start : start+l : start+l] + s1 := vals[start+l : start+2*l : start+2*l] + s2 := vals[start+2*l : start+3*l : start+3*l] + s3 := vals[start+3*l : start+4*l : start+4*l] + for k := range tw1 { + x0 := s0[k] + x1 := s2[k] * tw1[k] + x2 := s1[k] * tw2[k] + x3 := s3[k] * tw3[k] + e0 := x0 + x2 + e1 := x0 - x2 + f0 := x1 + x3 + f1 := x1 - x3 + s0[k] = e0 + f0 + s2[k] = e0 - f0 + s1[k] = complex(real(e1)+q*imag(f1), imag(e1)-q*real(f1)) + s3[k] = complex(real(e1)-q*imag(f1), imag(e1)+q*real(f1)) + } + } +} + +// bluesteinPlan holds the two pieces of a chirp-z transform that depend +// only on the transform length and the direction sign: the chirp itself +// and the convolution kernel already carried into the frequency domain. +// Both are read-only once published, so any number of transforms may +// share one plan. +type bluesteinPlan struct { + chirp []complex128 + kernel []complex128 // the kernel's forward transform + m int +} + +// bluestein cache, keyed by size and direction sign like the twiddle +// tables. The plan is a pure function of the key, so serving it from the +// cache cannot change a single bit of the result; it only spares the +// chirp rebuild and one of the three transforms per call. The +// twiddleCacheMax policy applies: bigger sizes still build their plan, +// they just do not keep it. +var ( + bluesteinMu sync.RWMutex + bluesteinKeys = map[twiddleKey]*bluesteinPlan{} +) + +// bluesteinPlanFor returns the plan for a size-n transform with the +// given direction sign, building it on first use. +func bluesteinPlanFor(sign float64, n int) *bluesteinPlan { + key := twiddleKey{sign: sign, n: n} + bluesteinMu.RLock() + plan, ok := bluesteinKeys[key] + bluesteinMu.RUnlock() + if ok { + return plan + } + chirp := make([]complex128, n) + for k := range n { + k2 := (uint64(k) * uint64(k)) % uint64(2*n) + angle := sign * math.Pi * float64(k2) / float64(n) + sinA, cosA := math.Sincos(angle) + chirp[k] = complex(cosA, sinA) + } + m := 1 + for m < 2*n-1 { + m <<= 1 + } + // kernel is conj(chirp) for the non-negative offsets and its mirror + // for the wrapped negative ones. The kernel is symmetric: + // c_{−l} = c_l. + kernel := make([]complex128, m) + for i := range n { + kernel[i] = conj(chirp[i]) + if i != 0 { + kernel[m-n+i] = conj(chirp[n-i]) + } + } + fftPow2(kernel, -1) + plan = &bluesteinPlan{chirp: chirp, kernel: kernel, m: m} + if n <= twiddleCacheMax { + bluesteinMu.Lock() + bluesteinKeys[key] = plan + bluesteinMu.Unlock() + } + return plan +} + +// bluestein computes the DFT through the chirp-z transform. With the +// chirp c_l = e^{sign·πi·l²/n}, the identity 2jk = j² + k² − (j−k)² gives +// X_k = c_k · Σ_j (x_j·c_j)·conj(c_{j−k}), a circular convolution that +// runs as a multiplication in the frequency domain of a padded power of +// two m ≥ 2n−1. Squares reduce modulo 2n (the chirp's period) to keep the +// angle small and exact. +func bluestein(vals []complex128, sign float64) { + n := len(vals) + plan := bluesteinPlanFor(sign, n) + chirp, m := plan.chirp, plan.m + + // a carries the input over the first n slots, zero padded to m so the + // circular convolution has room for every offset. + a := make([]complex128, m) + for i := range vals { + a[i] = vals[i] * chirp[i] + } + + fftPow2(a, -1) + kernel := plan.kernel + for i := range a { + a[i] *= kernel[i] + } + fftPow2(a, +1) + scale := complex(1/float64(m), 0) + for i := range a { + a[i] *= scale + } + + for i := range vals { + vals[i] = a[i] * chirp[i] + } +} + +// conj returns the complex conjugate. +func conj(c complex128) complex128 { + return complex(real(c), -imag(c)) +} + +func isPowerOfTwo(n int) bool { + return n > 0 && n&(n-1) == 0 +} + +// FFT2 returns the 2-D discrete Fourier transform. Input shape +// (H, W), output shape (H, W). Separable: FFT each row, then each +// column. Empty arrays are errors. +func FFT2(a *core.Array) (*core.Array, error) { + return fftND(a, 2, false) +} + +// IFFT2 returns the inverse 2-D DFT, scaled by 1/(H·W). +func IFFT2(a *core.Array) (*core.Array, error) { + return fftND(a, 2, true) +} + +// FFT3 returns the 3-D DFT over shape (D, H, W). +func FFT3(a *core.Array) (*core.Array, error) { + return fftND(a, 3, false) +} + +// IFFT3 returns the inverse 3-D DFT. +func IFFT3(a *core.Array) (*core.Array, error) { + return fftND(a, 3, true) +} + +// FFTN returns the N-D DFT along the dimensions in dims (or all +// dimensions if dims is empty); IFFTN over the same dims is its +// inverse. Duplicate entries in dims are applied once per entry: a +// dimension listed twice is transformed twice, the second pass +// running on the result of the first. +func FFTN(a *core.Array, dims []int) (*core.Array, error) { + if len(dims) == 0 { + dims = base.RangeN(a.NDim()) + } + cur := a + for _, d := range dims { + // The copy condition mirrors fftND: copy while the array is + // still the caller's input, then chain in place. + out, err := fftAlongDimAny(cur, d, false, cur == a) + if err != nil { + return nil, err + } + cur = out + } + return cur, nil +} + +// IFFTN is the inverse N-D DFT, the inverse of FFTN over the same +// dims: it divides by the product of the transformed extents, so +// IFFTN(FFTN(a, dims), dims) restores a. +func IFFTN(a *core.Array, dims []int) (*core.Array, error) { + if len(dims) == 0 { + dims = base.RangeN(a.NDim()) + } + cur := a + product := 1 + for _, d := range dims { + if d < 0 || d >= cur.NDim() { + return nil, base.Errf("IFFTN: dimension %d out of range for shape %s", d, base.ShapeText(cur.Shape())) + } + // The shape does not change along the chain, so the extents can + // be read before each transform; the scale is their product, not + // the total length, which differs when dims covers a subset. + product *= cur.Shape()[d] + out, err := fftAlongDimAny(cur, d, true, cur == a) + if err != nil { + return nil, err + } + cur = out + } + // Scale by 1/product: each 1-D IFFT is un-scaled (matches the 1-D + // IFFT convention where the scaling is done once at the end). + scale := 1.0 / float64(product) + oc := cur.RawComplexes() + for i := range oc { + oc[i] = complex(real(oc[i])*scale, imag(oc[i])*scale) + } + return cur, nil +} + +// RFFT returns the real-input DFT: input is 1-D real-valued, output +// has length n/2+1 complex entries (the non-redundant half of the +// full complex spectrum). For real-valued input the second half is +// the complex conjugate of the first. +func RFFT(a *core.Array) (*core.Array, error) { + if a.NDim() != 1 { + return nil, base.Errf("RFFT: needs a 1-D real input, got shape %s", base.ShapeText(a.Shape())) + } + if a.Dtype() == core.Complex { + return nil, base.Errf("RFFT: input must be real") + } + n := a.Len() + if n == 0 { + return nil, base.Errf("RFFT: an empty array has no FFT") + } + vals := make([]complex128, n) + for i := range n { + vals[i] = complex(a.FloatAt(i), 0) + } + transform(vals, -1) + half := n/2 + 1 + out := make([]complex128, half) + copy(out, vals[:half]) + return complexFromArrayMust(out, []int{half}), nil +} + +// IRFFT is the inverse of RFFT: it takes the non-redundant half of a +// real spectrum and returns the original real signal of length n. n +// must be at least 1; if zero, it defaults to 2·(len−1), or to 1 when +// the spectrum holds the single DC bin. +func IRFFT(a *core.Array, n int) (*core.Array, error) { + if a.NDim() != 1 { + return nil, base.Errf("IRFFT: needs a 1-D complex input, got shape %s", base.ShapeText(a.Shape())) + } + if a.Dtype() != core.Complex { + return nil, base.Errf("IRFFT: input must be complex") + } + half := a.Len() + if half < 1 { + return nil, base.Errf("IRFFT: the spectrum must have at least one entry, got %d", half) + } + if n == 0 { + if half == 1 { + n = 1 + } else { + n = 2 * (half - 1) + } + } + if n < 1 { + return nil, base.Errf("IRFFT: n must be at least 1, got %d", n) + } + if half != n/2+1 { + return nil, base.Errf("IRFFT: spectrum length %d doesn't match n=%d", half, n) + } + full := make([]complex128, n) + for i := range half { + full[i] = a.ComplexAt(i) + } + for i := half; i < n; i++ { + full[i] = complex(real(full[n-i]), -imag(full[n-i])) + } + transform(full, +1) + scale := complex(1/float64(n), 0) + out := make([]float64, n) + for i := range n { + out[i] = real(full[i] * scale) + } + return floatsFromArrayMust(out, []int{n}), nil +} + +// FFTFreq returns the discrete Fourier transform sample frequencies +// for a signal of length n with sample spacing d. d=1.0 by default. +// Returns the positive and negative frequencies in the standard FFT +// order: [0, 1/n, 2/n, ..., -1/2, ..., -1/n] / d. +func FFTFreq(n int, d float64) *core.Array { + if n <= 0 { + return core.New(core.Float, []int{0}...) + } + if d == 0 { + d = 1 + } + vals := make([]float64, n) + // Standard FFT order: 0, 1, …, ⌈n/2⌉−1 over the positive half, + // then −⌊n/2⌋, …, −1; for even n the Nyquist slot n/2 carries + // the negative −1/2, as the doc's "…, −1/2, …" promises. + posHalf := (n-1)/2 + 1 + for i := range posHalf { + vals[i] = float64(i) / float64(n) / d + } + for i := posHalf; i < n; i++ { + vals[i] = float64(i-n) / float64(n) / d + } + return floatsFromArrayMust(vals, []int{n}) +} + +// fftND computes the n-D FFT (rank=2 or 3). Separable over each axis. +func fftND(a *core.Array, rank int, inverse bool) (*core.Array, error) { + if a.NDim() != rank { + return nil, base.Errf("fft%d: needs a %d-D array, got shape %s", rank, rank, base.ShapeText(a.Shape())) + } + if a.Len() == 0 { + return nil, base.Errf("fft%d: an empty array has no transform", rank) + } + cur := a + // Promote to complex on first iteration. + if cur.Dtype() != core.Complex { + vals := make([]complex128, cur.Len()) + for i := range vals { + vals[i] = complex(cur.FloatAt(i), 0) + } + cur = complexFromArrayMust(vals, cur.Shape()) + } + for d := rank - 1; d >= 0; d-- { + // The first pass copies while cur is still the caller's own + // input (unless the promotion above already built a private + // payload); every later pass reads the previous pass's own + // output, which no one else can see, so it transforms in + // place. Either way the caller's input keeps its bits + // untouched. + out, err := fftAlongDimComplex(cur, d, inverse, cur == a) + if err != nil { + return nil, err + } + cur = out + } + // Scale for inverse: divide by total size. + if inverse { + scale := 1.0 / float64(cur.Len()) + oc := cur.RawComplexes() + for i := range oc { + oc[i] = complex(real(oc[i])*scale, imag(oc[i])*scale) + } + } + return cur, nil +} + +// fftAlongDimAny is the public-facing entry point: it accepts real +// or complex input and promotes to complex internally. Used by +// FFTN/IFFTN, which need to chain along multiple axes. copyIn may +// transform the payload in place only when the caller owns it; a +// promotion always builds a private payload, so it copies nothing +// regardless. +func fftAlongDimAny(a *core.Array, dim int, inverse, copyIn bool) (*core.Array, error) { + if a.Dtype() != core.Complex { + vals := make([]complex128, a.Len()) + for i := range vals { + vals[i] = complex(a.FloatAt(i), 0) + } + a = complexFromArrayMust(vals, a.Shape()) + copyIn = false + } + return fftAlongDimComplex(a, dim, inverse, copyIn) +} + +// fftLineFloor returns the smallest number of line transforms worth a +// worker's spawn: a line of length n costs about n·⌈log₂n⌉ butterfly +// steps, and a strided line pays a gather and a scatter of the same +// length on top, so the floor is the count carrying a fixed budget of +// them. The stage walk inside a line keeps its own parallelMinBlock +// floor, so a split set of short lines stays serial however many lines +// it holds. +func fftLineFloor(n int, strided bool) int { + steps := n * max(1, bits.Len(uint(n-1))) + if strided { + steps *= 2 + } + return workFloorFor(steps) +} + +// fftAlongDimComplex applies a 1-D FFT along a single dimension of a +// complex-valued array. Returns a new complex array, or the input +// itself when copyIn is false and the payload is already private: the +// caller must then guarantee no other alias reads the array, because +// the transform runs on it directly. Skipping the copy lets the +// multi-dimension transforms chain without re-copying an array they +// just built. +func fftAlongDimComplex(a *core.Array, dim int, inverse, copyIn bool) (*core.Array, error) { + if dim < 0 || dim >= a.NDim() { + return nil, base.Errf("fft: dimension %d out of range for shape %s", dim, base.ShapeText(a.Shape())) + } + if a.Dtype() != core.Complex { + return nil, base.Errf("fft: input must be complex") + } + nDim := a.Shape()[dim] + if nDim == 0 { + return nil, base.Errf("fft: zero-length dimension %d", dim) + } + stride := 1 + for k := dim + 1; k < a.NDim(); k++ { + stride *= a.Shape()[k] + } + blocks := 1 + for k := range dim { + blocks *= a.Shape()[k] + } + total := a.Len() + out := a + oc := a.RawComplexes() + if copyIn || oc == nil || a.Strided() { + // The input is the caller's (or not a flat complex payload): + // work on a copy so the caller's array keeps its values. A + // strided payload is gathered in element order; a linear copy + // would keep the physical order instead. + vals := make([]complex128, total) + if a.Strided() { + for i := range total { + vals[i] = a.ComplexAt(i) + } + } else { + copy(vals, oc) + } + out = complexFromArrayMust(vals, a.Shape()) + oc = out.RawComplexes() + } + sign := -1.0 + if inverse { + sign = +1.0 + } + // Every (block, stride) line transforms independently: the lines + // split across workers. Each worker allocates one scratch + // line and reuses it for every position it serves: the line is + // fully rewritten before each transform, so no stale value can + // leak, and two lines never share an output cell. Splitting the + // flat line space (rather than blocks, then strides inside a + // block) keeps single-block layouts, a lone row of columns, as + // parallel as any other. + if stride == 1 { + // Unit stride: every line is already contiguous in the output + // copy, so gather and scatter are pure overhead. Transform each + // line in place: the initial copy above keeps the input + // untouched, and the lines of disjoint blocks never overlap. + engine.ParallelMin(blocks, fftLineFloor(nDim, false), func(bs, be int) { + for b := bs; b < be; b++ { + off := b * nDim + transform(oc[off:off+nDim], sign) + } + }) + return out, nil + } + engine.ParallelMin(blocks*stride, fftLineFloor(nDim, true), func(ls, le int) { + slice := make([]complex128, nDim) + for line := ls; line < le; line++ { + b := line / stride + s := line % stride + base := b * nDim * stride + for j := range nDim { + slice[j] = oc[base+j*stride+s] + } + transform(slice, sign) + for j := range nDim { + oc[base+j*stride+s] = slice[j] + } + } + }) + return out, nil +} + +// complexFromArrayMust and floatsFromArrayMust wrap the taking +// constructors whose lengths always match their shapes by +// construction. They are unexported on purpose: a library that panics +// on a caller's input is a defect, and the callers here cannot fail. +func complexFromArrayMust(vals []complex128, shape []int) *core.Array { + a, err := core.ComplexFromArray(vals, shape...) + if err != nil { + panic("signal: " + err.Error()) + } + return a +} + +func floatsFromArrayMust(vals []float64, shape []int) *core.Array { + a, err := core.FloatsFromArray(vals, shape...) + if err != nil { + panic("signal: " + err.Error()) + } + return a +} diff --git a/signal/fft_nd_test.go b/signal/fft_nd_test.go new file mode 100644 index 0000000..3afe7b6 --- /dev/null +++ b/signal/fft_nd_test.go @@ -0,0 +1,247 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "strings" + "testing" +) + +func TestFFT2(t *testing.T) { + // 2×2 identity-like: shift. + a := mustFromFloats(t, []float64{ + 1, 0, + 0, 0, + }, 2, 2) + got, err := FFT2(a) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 2 || got.Shape()[1] != 2 { + t.Errorf("FFT2 shape: %v", got.Shape()) + } + if got.Dtype() != core.Complex { + t.Errorf("FFT2 dtype: %s", got.Dtype()) + } + // Round-trip: IFFT2(FFT2(x)) ≈ x. + back, err := IFFT2(got) + if err != nil { + t.Fatal(err) + } + for i := range 4 { + v, _ := core.ComplexAt(back, i/2, i%2) + orig, _ := core.FloatAt(a, i/2, i%2) + if math.Abs(real(v)-orig) > 1e-9 { + t.Errorf("FFT2 round-trip [%d]: got %v, want %v", i, real(v), orig) + } + } +} + +func TestFFT3(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 0, 0, 0, + 0, 0, 0, 0, + 0, 0, 0, 0, + }, 3, 2, 2) + got, err := FFT3(a) + if err != nil { + t.Fatal(err) + } + if got.Shape()[0] != 3 || got.Shape()[1] != 2 || got.Shape()[2] != 2 { + t.Errorf("FFT3 shape: %v", got.Shape()) + } + // Round-trip. + back, err := IFFT3(got) + if err != nil { + t.Fatal(err) + } + v0, _ := core.ComplexAt(back, 0, 0, 0) + if math.Abs(real(v0)-1) > 1e-9 { + t.Errorf("FFT3 round-trip [0,0,0]: got %v, want 1", real(v0)) + } +} + +func TestFFTN(t *testing.T) { + a := mustFromFloats(t, []float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + }, 2, 4) + got, err := FFTN(a, nil) + if err != nil { + t.Fatal(err) + } + if got.Dtype() != core.Complex { + t.Errorf("FFTN dtype: %s", got.Dtype()) + } + // Round-trip. + back, err := IFFTN(got, nil) + if err != nil { + t.Fatal(err) + } + for i := range a.Len() { + v, _ := core.ComplexAt(back, i/4, i%4) + orig, _ := core.FloatAt(a, i/4, i%4) + if math.Abs(real(v)-orig) > 1e-9 { + t.Errorf("FFTN round-trip [%d]: got %v, want %v", i, real(v), orig) + } + } +} + +func TestRFFT(t *testing.T) { + // Real input of length 8 gives a spectrum of length 5. + a := mustFromFloats(t, []float64{1, 0, 0, 0, 0, 0, 0, 0}, 8) + spec, err := RFFT(a) + if err != nil { + t.Fatal(err) + } + if spec.Len() != 5 { + t.Errorf("RFFT spectrum length: %d, want 5", spec.Len()) + } + // IRFFT round-trip. + back, err := IRFFT(spec, 8) + if err != nil { + t.Fatal(err) + } + for i := range 8 { + v, _ := core.FloatAt(back, i) + orig, _ := core.FloatAt(a, i) + if math.Abs(v-orig) > 1e-9 { + t.Errorf("RFFT round-trip [%d]: got %v, want %v", i, v, orig) + } + } +} + +func TestFFTFreq(t *testing.T) { + f := FFTFreq(8, 1.0) + if f.Len() != 8 { + t.Errorf("FFTFreq length: %d", f.Len()) + } + // Frequencies: [0, 1/8, 2/8, 3/8, -4/8, -3/8, -2/8, -1/8] = [0, 0.125, 0.25, 0.375, -0.5, -0.375, -0.25, -0.125] + want := []float64{0, 0.125, 0.25, 0.375, -0.5, -0.375, -0.25, -0.125} + for i, w := range want { + v, _ := core.FloatAt(f, i) + if math.Abs(v-w) > 1e-9 { + t.Errorf("FFTFreq [%d]: got %v, want %v", i, v, w) + } + } +} + +// TestIRFFTLengthOneSpectrum pins the degenerate spectrum: a single +// DC bin defaults to n = 1 instead of n = 0, an empty spectrum is an +// error, and a non-positive explicit n is an error. +func TestIRFFTLengthOneSpectrum(t *testing.T) { + spec := mustComplexes(t, []complex128{complex(5, 0)}, 1) + back, err := IRFFT(spec, 0) + if err != nil { + t.Fatalf("IRFFT default n: %v", err) + } + if back.Len() != 1 || math.Abs(back.FloatAt(0)-5) > 1e-12 { + t.Fatalf("IRFFT of a DC-only spectrum = %v, want [5]", back) + } + back, err = IRFFT(spec, 1) + if err != nil { + t.Fatalf("IRFFT n=1: %v", err) + } + if back.Len() != 1 || math.Abs(back.FloatAt(0)-5) > 1e-12 { + t.Fatalf("IRFFT n=1 = %v, want [5]", back) + } + if _, err := IRFFT(spec, 3); err == nil { + t.Fatal("expected an error for a spectrum length that mismatches n") + } + if _, err := IRFFT(mustComplexes(t, nil, 0), 0); err == nil { + t.Fatal("expected an error for an empty spectrum") + } + if _, err := IRFFT(spec, -2); err == nil { + t.Fatal("expected an error for a negative n") + } +} + +// TestFFTEmptyComplexIsError pins the empty guard for complex input, +// which bypasses the real-to-complex conversion where the check used +// to live. +func TestFFTEmptyComplexIsError(t *testing.T) { + empty := mustComplexes(t, nil, 0) + if _, err := FFT(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("FFT of an empty complex array: %v", err) + } + if _, err := IFFT(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("IFFT of an empty complex array: %v", err) + } +} + +// TestIFFTNSubsetDimsRoundTrip pins the inverse scaling: IFFTN divides +// by the product of the transformed extents, so inverting a subset of +// the dimensions restores the input exactly (scaling by the total +// length would return a/2 here). +func TestIFFTNSubsetDimsRoundTrip(t *testing.T) { + src := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) + fwd, err := FFTN(src, []int{1}) + if err != nil { + t.Fatalf("FFTN: %v", err) + } + back, err := IFFTN(fwd, []int{1}) + if err != nil { + t.Fatalf("IFFTN: %v", err) + } + for i := range 6 { + want, _ := core.FloatAt(src, i) + got, _ := core.ComplexAt(back, i) + if math.Abs(real(got)-want) > 1e-9 || math.Abs(imag(got)) > 1e-9 { + t.Fatalf("subset round-trip [%d]: got %v, want %v", i, got, want) + } + } + // The full-dims path keeps its usual total scaling. + fwdAll, err := FFTN(src, nil) + if err != nil { + t.Fatalf("FFTN all dims: %v", err) + } + backAll, err := IFFTN(fwdAll, nil) + if err != nil { + t.Fatalf("IFFTN all dims: %v", err) + } + for i := range 6 { + want, _ := core.FloatAt(src, i) + got, _ := core.ComplexAt(backAll, i) + if math.Abs(real(got)-want) > 1e-9 || math.Abs(imag(got)) > 1e-9 { + t.Fatalf("full round-trip [%d]: got %v, want %v", i, got, want) + } + } +} + +// TestFFT2DoesNotMutateComplexInput pins the immutability contract for +// the multi-dimension chain: FFT2 with an already complex input must +// not overwrite the caller's payload. An in-place first pass used to +// slip through here because the return values stayed correct; only the +// caller's array silently held the spectrum afterwards. +func TestFFT2DoesNotMutateComplexInput(t *testing.T) { + vals := []complex128{ + 1 + 2i, 3 - 1i, + -0.5 + 0.5i, 2 + 0i, + } + a, err := core.FromComplexes(vals, 2, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + got, err := FFT2(a) + if err != nil { + t.Fatalf("FFT2: %v", err) + } + for i, want := range vals { + if a.ComplexAt(i) != want { + t.Fatalf("FFT2 mutated input[%d] = %v, want %v", i, a.ComplexAt(i), want) + } + } + back, err := IFFT2(got) + if err != nil { + t.Fatalf("IFFT2: %v", err) + } + for i, want := range vals { + if back.ComplexAt(i) != want { + t.Fatalf("IFFT2 mutated gradient input or lost precision [%d] = %v, want %v", i, back.ComplexAt(i), want) + } + } +} diff --git a/signal/fft_precision_test.go b/signal/fft_precision_test.go new file mode 100644 index 0000000..8a458bb --- /dev/null +++ b/signal/fft_precision_test.go @@ -0,0 +1,233 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "math/big" + "testing" +) + +// The fused radix-4 stages replace the radix-2 walk's arithmetic, so +// the accuracy question is answered against references this package +// does not implement: the quadratic definition in float64 (the one the +// functional tests read) and the same definition carried out in 256-bit +// floating point. The two paths under comparison are called directly, +// the radix-2 walk and the package dispatcher, so the error figures +// attribute cleanly to the stage kernels. + +// precisionFixture returns the fixed test signal of length n: small +// integer-like values whose transform is far above the rounding floor +// of either path but far below the tolerances below. +func precisionFixture(n int) []complex128 { + vals := make([]complex128, n) + for i := range vals { + vals[i] = complex(float64(i%17)-8, float64(i%5)-2) + } + return vals +} + +// TestFFTPathAccuracyAgainstNaive pins the fused radix-4 stages to the +// quadratic float64 reference at lengths the functional tests never +// reach (every length they hold sits below the radix-4 floor). The +// tolerance leaves the quadratic reference's own accumulated rounding +// an order of magnitude of headroom. +func TestFFTPathAccuracyAgainstNaive(t *testing.T) { + for _, n := range []int{64, 128, 256, 1024} { + vals := precisionFixture(n) + want := naiveDFT(vals, -1) + got := append([]complex128(nil), vals...) + fftPow2(got, -1) + for i := range want { + if math.Abs(real(got[i])-real(want[i])) > 1e-6 { + t.Fatalf("FFT(%d)[%d]: real error %g against the quadratic reference", n, i, math.Abs(real(got[i])-real(want[i]))) + } + if math.Abs(imag(got[i])-imag(want[i])) > 1e-6 { + t.Fatalf("FFT(%d)[%d]: imaginary error %g against the quadratic reference", n, i, math.Abs(imag(got[i])-imag(want[i]))) + } + } + } +} + +func newBig(prec uint) *big.Float { return new(big.Float).SetPrec(prec) } + +// machinPi computes π at the given precision by Machin's formula +// π = 16·atan(1/5) − 4·atan(1/239), the arctangents by their Taylor +// series. +func machinPi(prec uint) *big.Float { + atan := func(x *big.Float) *big.Float { + x2 := newBig(prec).Mul(x, x) + term := newBig(prec).Set(x) + sum := newBig(prec).Set(x) + tiny := newBig(prec).SetMantExp(big.NewFloat(1), -int(prec)-12) + for k := 1; k < 200; k++ { + term.Mul(term, x2) + // The denominator moves from 2k−1 to 2k+1 across the step. + term.Mul(term, big.NewFloat(float64(2*k-1))) + term.Quo(term, big.NewFloat(float64(2*k+1))) + if k%2 == 1 { + sum.Sub(sum, term) + } else { + sum.Add(sum, term) + } + if abs := newBig(prec).Abs(term); abs.Cmp(tiny) < 0 { + break + } + } + return sum + } + fifth := newBig(prec).Quo(big.NewFloat(1), big.NewFloat(5)) + t239th := newBig(prec).Quo(big.NewFloat(1), big.NewFloat(239)) + pi := atan(fifth) + pi.Mul(pi, big.NewFloat(16)) + t := atan(t239th) + t.Mul(t, big.NewFloat(4)) + return pi.Sub(pi, t) +} + +// sincosBig evaluates cosine and sine of a finite angle at 256-bit +// precision: the angle is range-reduced into [−π, π] and both series +// are summed to well below the precision floor. The pair comes back +// real part first. +func sincosBig(theta *big.Float, prec uint) [2]*big.Float { + pi := machinPi(prec) + twoPi := newBig(prec).Mul(pi, big.NewFloat(2)) + piNeg := newBig(prec).Neg(pi) + for theta.Cmp(pi) > 0 { + theta.Sub(theta, twoPi) + } + for theta.Cmp(piNeg) < 0 { + theta.Add(theta, twoPi) + } + return [2]*big.Float{sincosSeries(theta, false, prec), sincosSeries(theta, true, prec)} +} + +// sincosSeries sums the Taylor series of sine (odd powers) or cosine +// (even powers) of a reduced angle. +func sincosSeries(theta *big.Float, sine bool, prec uint) *big.Float { + sum := newBig(prec) + term := newBig(prec) + if sine { + term.Set(theta) + } else { + term.SetInt64(1) + } + theta2 := newBig(prec).Mul(theta, theta) + tiny := newBig(prec).SetMantExp(big.NewFloat(1), -int(prec)-12) + for k := range 120 { + if k%2 == 1 { + sum.Sub(sum, term) + } else { + sum.Add(sum, term) + } + d1, d2 := 2*k+1, 2*k+2 + if sine { + d1, d2 = 2*k+2, 2*k+3 + } + term.Mul(term, theta2) + term.Quo(term, big.NewFloat(float64(d1)*float64(d2))) + if abs := newBig(prec).Abs(term); abs.Cmp(tiny) < 0 { + break + } + } + return sum +} + +// bigFloatDFT evaluates the DFT by its quadratic definition at 256-bit +// precision: the twiddle of every residue is an exact-angle series +// evaluation, the accumulation runs in big.Float complex arithmetic, +// and the result is rounded once into complex128. +func bigFloatDFT(vals []complex128, sign float64) []complex128 { + const prec = 256 + n := len(vals) + pi := machinPi(prec) + twoPi := newBig(prec).Mul(pi, big.NewFloat(2)) + tw := make([][2]*big.Float, n) + for r := range n { + angle := newBig(prec).SetInt64(int64(r)) + angle.Mul(angle, twoPi) + angle.Quo(angle, big.NewFloat(float64(n))) + if sign < 0 { + angle.Neg(angle) + } + tw[r] = sincosBig(angle, prec) + } + out := make([][2]*big.Float, n) + for k := range out { + out[k] = [2]*big.Float{newBig(prec), newBig(prec)} + } + vr, vi := newBig(prec), newBig(prec) + pr, pii := newBig(prec), newBig(prec) + tr, ti := newBig(prec), newBig(prec) + for k := range n { + re, im := out[k][0], out[k][1] + for j := range n { + t := tw[j*k%n] + vr.SetFloat64(real(vals[j])) + vi.SetFloat64(imag(vals[j])) + tr.Set(t[0]) + ti.Set(t[1]) + pr.Mul(vr, tr) + pii.Mul(vi, ti) + pr.Sub(pr, pii) + pii.Mul(vr, ti) + ti.Mul(vi, tr) + pii.Add(pii, ti) + re.Add(re, pr) + im.Add(im, pii) + } + } + res := make([]complex128, n) + for k := range n { + f64re, _ := out[k][0].Float64() + f64im, _ := out[k][1].Float64() + res[k] = complex(f64re, f64im) + } + return res +} + +// TestFFTPrecisionReport compares the radix-2 walk and the fused +// radix-4 dispatch against the 256-bit quadratic reference and reports +// both maximum errors. The assertion only demands that the fused path +// hold the radix-2 path's accuracy within a small factor; the logged +// figures are the evidence the report quotes. +func TestFFTPrecisionReport(t *testing.T) { + for _, n := range []int{64, 256} { + vals := precisionFixture(n) + ref := bigFloatDFT(vals, -1) + oldPath := append([]complex128(nil), vals...) + fftRadix2(oldPath, -1) + newPath := append([]complex128(nil), vals...) + fftPow2(newPath, -1) + errOld, errNew := 0.0, 0.0 + for i := range ref { + eo := math.Max(math.Abs(real(oldPath[i])-real(ref[i])), math.Abs(imag(oldPath[i])-imag(ref[i]))) + en := math.Max(math.Abs(real(newPath[i])-real(ref[i])), math.Abs(imag(newPath[i])-imag(ref[i]))) + errOld = math.Max(errOld, eo) + errNew = math.Max(errNew, en) + } + t.Logf("n=%d: radix-2 max error %.3e, fused radix-4 max error %.3e, ratio %.3f", n, errOld, errNew, errNew/errOld) + if errNew > 4*errOld+1e-13 { + t.Fatalf("n=%d: fused radix-4 error %.3e is worse than 4x the radix-2 error %.3e", n, errNew, errOld) + } + } +} + +// TestBigFloatReferenceAgreesWithQuadratic guards the precision +// harness itself: at lengths this small the 256-bit reference and the +// float64 quadratic definition must agree to the float64 path's own +// rounding, so a broken series evaluation cannot pass unnoticed into +// the error figures above. +func TestBigFloatReferenceAgreesWithQuadratic(t *testing.T) { + for _, n := range []int{4, 8} { + vals := precisionFixture(n) + ref := bigFloatDFT(vals, -1) + naive := naiveDFT(vals, -1) + for i := range ref { + if math.Abs(real(ref[i])-real(naive[i])) > 1e-12 || math.Abs(imag(ref[i])-imag(naive[i])) > 1e-12 { + t.Fatalf("n=%d [%d]: 256-bit reference %v, quadratic %v", n, i, ref[i], naive[i]) + } + } + } +} diff --git a/signal/fft_test.go b/signal/fft_test.go new file mode 100644 index 0000000..098fcc9 --- /dev/null +++ b/signal/fft_test.go @@ -0,0 +1,124 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "strings" + "testing" +) + +// naiveDFT computes the transform quadratically, the reference the fast +// paths are verified against. +func naiveDFT(vals []complex128, sign float64) []complex128 { + n := len(vals) + out := make([]complex128, n) + for k := range n { + var s complex128 + for j, x := range vals { + angle := sign * 2 * math.Pi * float64(j*k%n) / float64(n) + s += x * complex(math.Cos(angle), math.Sin(angle)) + } + out[k] = s + } + return out +} + +func assertCplxApprox(t *testing.T, name string, got []complex128, want []complex128) { + t.Helper() + for i := range want { + if math.Abs(real(got[i])-real(want[i])) > 1e-9 || math.Abs(imag(got[i])-imag(want[i])) > 1e-9 { + t.Fatalf("%s[%d]: got %v, want %v", name, i, got[i], want[i]) + } + } +} + +func fftValues(t *testing.T, a *core.Array) []complex128 { + t.Helper() + out, err := FFT(a) + if err != nil { + t.Fatalf("FFT: %v", err) + } + vals := make([]complex128, out.Len()) + for i := range vals { + vals[i], _ = core.ComplexAt(out, i) + } + return vals +} + +func TestFFTKnownValues(t *testing.T) { + // FFT of [1, 2, 3, 4] is the textbook [10, -2+2i, -2, -2-2i]. + a := mustFromInts(t, []int64{1, 2, 3, 4}, 4) + got := fftValues(t, a) + assertCplxApprox(t, "FFT", got, []complex128{10, complex(-2, 2), -2, complex(-2, -2)}) + + // A constant input concentrates in the DC bin. + c := mustFromFloats(t, []float64{2.5, 2.5, 2.5, 2.5}, 4) + got = fftValues(t, c) + assertCplxApprox(t, "FFT DC", got[:1], []complex128{10}) + for i := 1; i < 4; i++ { + if math.Abs(real(got[i])) > 1e-9 || math.Abs(imag(got[i])) > 1e-9 { + t.Fatalf("FFT constant input leaked into bin %d: %v", i, got[i]) + } + } +} + +func TestFFTAgainstNaive(t *testing.T) { + // Powers of two take the radix-2 path; the odd and mixed lengths take + // Bluestein. Both must agree with the quadratic reference. + lengths := []int{1, 2, 4, 8, 16, 3, 5, 6, 7, 12, 13, 20} + for _, n := range lengths { + vals := make([]complex128, n) + g := core.NewGenerator(int64(n)) + for i := range vals { + f, _ := core.Floats(g, 2) + re, _ := core.FloatAt(f, 0) + im, _ := core.FloatAt(f, 1) + vals[i] = complex(re*10-5, im*10-5) + } + src, err := core.FromComplexes(vals, n) + if err != nil { + t.Fatalf("FromComplexes(%d): %v", n, err) + } + assertCplxApprox(t, "FFT radix/bluestein", fftValues(t, src), naiveDFT(vals, -1)) + } +} + +func TestIFFTRoundTrip(t *testing.T) { + all := []float64{0.5, -1.25, 3, -0.5, 2, 1, -2, 4, 0.25, -3, 1.5, 2.75} + for _, n := range []int{8, 5, 12} { + src := mustFromFloats(t, all[:n], n) + fwd, err := FFT(src) + if err != nil { + t.Fatalf("FFT(%d): %v", n, err) + } + back, err := IFFT(fwd) + if err != nil { + t.Fatalf("IFFT(%d): %v", n, err) + } + for i := range n { + want, _ := core.FloatAt(src, i) + got, _ := core.ComplexAt(back, i) + if math.Abs(real(got)-want) > 1e-9 || math.Abs(imag(got)) > 1e-9 { + t.Fatalf("IFFT(%d)[%d]: got %v, want %v", n, i, got, want) + } + } + } +} + +func TestFFTErrors(t *testing.T) { + empty := mustFromInts(t, nil, 0) + if _, err := FFT(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("FFT empty: %v", err) + } + if _, err := IFFT(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("IFFT empty: %v", err) + } + m := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) + if _, err := FFT(m); err == nil || !strings.Contains(err.Error(), "needs a 1-D array") { + t.Fatalf("FFT 2-D: %v", err) + } +} diff --git a/signal/filter.go b/signal/filter.go new file mode 100644 index 0000000..e5cb2f0 --- /dev/null +++ b/signal/filter.go @@ -0,0 +1,515 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Savitzky-Golay smoothing: every sample is replaced by the value at +// its position of a polynomial of the given order fitted by least +// squares over a symmetric window of samples. Unlike a moving average +// this preserves the moments of the signal up to the fit order, so +// peaks keep their heights and widths instead of being flattened, +// which is why spectroscopy and chromatography standardised on it. +// +// The interior uses one set of precomputed weights, the polynomial +// fit's hat row for the window centre. Samples within half a window +// of an edge use a window truncated at the boundary, fitted at the +// sample's own position, so the output is complete without padding +// and the ends keep the full polynomial treatment at reduced support. + +// SavitzkyGolay smooths a rank-1 signal with a Savitzky-Golay +// polynomial filter: window is the (odd) number of samples per fit, +// order the polynomial degree. The standard constraint order ≤ +// window/2 applies: a higher degree leaves the truncated edge windows +// with fewer samples than coefficients, so the fit is underdetermined +// and an order above window/2 is an error. A polynomial of degree ≤ +// order passes through unchanged. +func SavitzkyGolay(data *core.Array, window, order int) (*core.Array, error) { + if data.Dtype() == core.Complex { + return nil, base.Errf("SavitzkyGolay: complex signals are not supported") + } + if data.NDim() != 1 { + return nil, base.Errf("SavitzkyGolay: needs a rank-1 signal, got shape %s", base.ShapeText(data.Shape())) + } + n := data.Len() + if n == 0 { + return nil, base.Errf("SavitzkyGolay: the signal must not be empty") + } + if window < 3 || window%2 == 0 { + return nil, base.Errf("SavitzkyGolay: window must be an odd number ≥ 3, got %d", window) + } + if window > n { + return nil, base.Errf("SavitzkyGolay: window %d exceeds the signal length %d", window, n) + } + if order < 0 { + return nil, base.Errf("SavitzkyGolay: order must not be negative, got %d", order) + } + if order > window/2 { + return nil, base.Errf("SavitzkyGolay: order %d exceeds window/2 (%d), the edge fits would be underdetermined", + order, window/2) + } + h := window / 2 + + // Weights for a window with offsets offs, evaluating the fit at + // offset 0 (the centre sample). + weights := func(offs []int) ([]float64, error) { + k := order + 1 + m := len(offs) + // Normal equations of the Vandermonde system. + ata := make([][]float64, k) + for i := range k { + ata[i] = make([]float64, k) + for j := range k { + s := 0.0 + for _, o := range offs { + s += math.Pow(float64(o), float64(i+j)) + } + ata[i][j] = s + } + } + // Huge windows at high order overflow the power sums: a + // non-finite entry would flow through the solver into NaN + // weights published as a smooth signal, so it is refused here. + for i := range k { + for j := range k { + if math.IsInf(ata[i][j], 0) || math.IsNaN(ata[i][j]) { + return nil, base.Errf("SavitzkyGolay: the normal equations overflow for window %d at order %d", 2*h+1, order) + } + } + } + // Right-hand side: the coefficient vector of evaluating at 0 + // is e_0, so b_i = 0^i, which is 1 at i = 0 and 0 after. + b := make([]float64, k) + b[0] = 1 + sol, err := base.SolveSystem("SavitzkyGolay", ata, [][]float64{b}) + if err != nil { + return nil, base.Errf("SavitzkyGolay: %w", err) + } + coef := sol[0] + w := make([]float64, m) + for i, o := range offs { + s := 0.0 + for c := range k { + s += coef[c] * math.Pow(float64(o), float64(c)) + } + if math.IsInf(s, 0) || math.IsNaN(s) { + return nil, base.Errf("SavitzkyGolay: the weights overflow for window %d at order %d", 2*h+1, order) + } + w[i] = s + } + return w, nil + } + + // The interior weights and every edge point's truncated window are + // a pure function of the window and the order, so the whole fit set + // is served from a cache keyed by the two: the power sums, the + // solver and the weight evaluations run once per window-and-order + // instead of once per call. + fits, ferr := savitzkyGolayFits(window, order, h, weights) + if ferr != nil { + return nil, ferr + } + edge := h + out, oerr := core.Zeros(core.Float, []int{n}...) + if oerr != nil { + return nil, oerr + } + outRow := out.RawFloats() + src := widenFloats(data) + // Every output sample owns its own slot and, within half a window + // of an edge, its own precomputed weight fit, so the points split + // across workers once a worker's chunk clears the floor below; the + // per-point summation over the window keeps the serial order. The + // regions split so the interior walk carries no per-point branch: + // each region sees its weights and its slice bounds settled once. + // Every loop is clamped to the worker's own [ps, pe): an edge wider + // than one chunk would otherwise let the edge loops cross a chunk + // boundary, and two workers would write the same output slot at + // once, a data race whose bits happen to agree but whose access the + // memory model and the race detector both refuse. + engine.ParallelMin(n, sgMinPointsPerWorker, func(ps, pe int) { + sweep := func(i int) { + var w []float64 + start := max(i-h, 0) + switch { + case i < edge: + w = fits.left[i] + case i >= n-edge: + w = fits.right[n-1-i] + default: + w = fits.full + } + sum := 0.0 + row := src[start : start+len(w)] + for k, weight := range w { + sum += weight * row[k] + } + outRow[i] = sum + } + midLo := max(ps, edge) + midHi := min(pe, n-edge) + for i := ps; i < min(midLo, pe); i++ { + sweep(i) + } + if midHi > midLo { + // Interior: the full weights always apply and the window + // never truncates, so the loop runs without the fit choice. + for i := midLo; i < midHi; i++ { + start := i - h + sum := 0.0 + row := src[start : start+len(fits.full)] + for k, weight := range fits.full { + sum += weight * row[k] + } + outRow[i] = sum + } + } + for i := max(midHi, ps); i < pe; i++ { + sweep(i) + } + }) + return out, nil +} + +// sgFitKey names one cached fit set: the window and order that +// determine every weight row in it. +type sgFitKey struct { + window, order int +} + +// sgFits holds the interior hat row and the truncated edge fits: for +// left edge point j the window runs over the offsets [-j, h], for the +// mirrored right edge point over [-h, j]. The rows are read-only once +// published. +type sgFits struct { + full []float64 + left, right [][]float64 +} + +// sgFitCacheMax bounds how many window-and-order fit sets stay cached; +// each set is a few hundred floats, so the map stays negligible +// whatever the callers do. +const sgFitCacheMax = 256 + +var ( + sgFitMu sync.RWMutex + sgFitKeys = map[sgFitKey]*sgFits{} +) + +// savitzkyGolayFits returns the fit set for a window and order, +// building it through the caller's weights evaluation on first use. +// Every edge point owns a truncated window, fitted at the sample's own +// position, so the parallel sweep cannot fail mid-kernel. +func savitzkyGolayFits(window, order, h int, weights func(offs []int) ([]float64, error)) (*sgFits, error) { + key := sgFitKey{window: window, order: order} + sgFitMu.RLock() + fits, ok := sgFitKeys[key] + sgFitMu.RUnlock() + if ok { + return fits, nil + } + full, err := weights(rangeInts(-h, h)) + if err != nil { + return nil, err + } + edge := h + fits = &sgFits{full: full, left: make([][]float64, edge), right: make([][]float64, edge)} + for j := range edge { + lw, lerr := weights(rangeInts(-j, h)) + if lerr != nil { + return nil, lerr + } + fits.left[j] = lw + rw, rerr := weights(rangeInts(-h, j)) + if rerr != nil { + return nil, rerr + } + fits.right[j] = rw + } + sgFitMu.Lock() + if len(sgFitKeys) < sgFitCacheMax { + sgFitKeys[key] = fits + } + sgFitMu.Unlock() + return fits, nil +} + +// sgMinPointsPerWorker is the smallest per-worker chunk of output +// points the Savitzky-Golay sweep splits for: below it a chunk's +// window arithmetic no longer pays the worker spawn cost. +const sgMinPointsPerWorker = 1 << 10 + +// widenFloats returns the array's elements widened to float64 exactly +// as FloatAt widens them: a contiguous float64 payload aliases the +// array, every other dtype (and any strided array, whose payload order +// is not its element order) converts into a fresh slice. Treat an +// aliased result as read-only. +func widenFloats(a *core.Array) []float64 { + if a.Dtype() == core.Float && !a.Strided() { + return a.RawFloats() + } + out := make([]float64, a.Len()) + for i := range out { + out[i] = a.FloatAt(i) + } + return out +} + +// rangeInts returns the integers from a to b inclusive. +func rangeInts(a, b int) []int { + out := make([]int, b-a+1) + for i := range out { + out[i] = a + i + } + return out +} + +// Digital filter design. The Butterworth family is the flat +// magnitude response: |H|² is monotone in frequency with no ripple on +// either side, which is what a data-cleaning low-pass wants. The +// design runs the classical route: prewarp the cutoff with the +// bilinear transform's tangent, place the analog Butterworth poles on +// the left half-plane circle, map them inside the unit circle, and +// read the coefficients off the factored transfer function. Conjugate +// pole pairs keep every coefficient real; odd orders carry one real +// pole. + +// ButterworthLowPass designs an order-th low-pass filter for sample +// rate fs, returning the direct-form coefficients b (numerator) and a +// (denominator, a[0] = 1). The order must be at least 1, fs positive, +// and the cutoff strictly inside (0, fs/2). The response is −3 dB at +// the cutoff and unity at DC. +func ButterworthLowPass(order int, fs, cutoff float64) (b, a []float64, err error) { + return butterworth(order, fs, cutoff, false) +} + +// ButterworthHighPass designs the high-pass mirror: the same flat +// response turned around, −3 dB at the same cutoff, unity gain at +// Nyquist. It shares the low-pass poles and carries its zeros at +// z = 1, the bilinear image of the analog prototype's origin zeros. +func ButterworthHighPass(order int, fs, cutoff float64) (b, a []float64, err error) { + return butterworth(order, fs, cutoff, true) +} + +// butterworth carries the shared design: the digital low-pass poles, +// zeros and gain as coefficient polynomials, sign-flipped for the +// high-pass. +func butterworth(order int, fs, cutoff float64, highpass bool) (b, a []float64, err error) { + const name = "Butterworth" + if order < 1 { + return nil, nil, base.Errf("%s: the order must be at least 1, got %d", name, order) + } + if !(fs > 0) || math.IsInf(fs, 0) { + return nil, nil, base.Errf("%s: fs must be positive and finite, got %g", name, fs) + } + if !(cutoff > 0) || cutoff >= fs/2 { + return nil, nil, base.Errf("%s: the cutoff must lie in (0, fs/2), got %g for fs %g", name, cutoff, fs) + } + // Prewarped analog cutoff: the bilinear transform maps the digital + // cutoff onto it exactly, so the −3 dB point lands where asked. + w := math.Tan(math.Pi * cutoff / fs) + // Analog poles s_k = w·e^{iθ_k} on the left half-plane circle; + // odd orders place one pole at θ = π, which is real. + a = []float64{1} + for k := range order { + num := 2*k + order + 1 + var s complex128 + if num == 2*order { + // The odd order's real pole at θ = π: built exactly, not + // through sin(π), whose float noise would detour it into + // the pair branch and fabricate a quadratic factor. + s = complex(-w, 0) + } else { + theta := math.Pi * float64(num) / float64(2*order) + // One Sincos serves both factors, the same pair the + // separate calls produced, so the pole keeps its bits. + sinT, cosT := math.Sincos(theta) + s = complex(w*cosT, w*sinT) + } + z := (1 + s) / (1 - s) + if imag(s) == 0 { + a = mulPolyReal(a, []float64{1, -real(z)}) + continue + } + if imag(s) < 0 { + continue // the conjugate partner was consumed on its pass + } + zc := (1 + conj(s)) / (1 - conj(s)) + // (1 − z·u)(1 − z̄·u) = 1 − 2Re(z)·u + |z|²·u², real by + // construction; the rounding noise in the imaginary part drops. + a = mulPolyReal(a, []float64{1, -2 * real(z), real(z * zc)}) + } + // Zeros: the low-pass puts them all at z = −1 (numerator + // (1+u)^order, DC-normalised); the high-pass shares the very same + // poles and moves the zeros to z = +1 (numerator (1−u)^order, + // Nyquist-normalised). Reflecting the poles by negating z instead + // would answer a different question at the wrong frequencies. + if highpass { + b = binomialCoeffs(order) + for i := range b { + if i%2 == 1 { + b[i] = -b[i] + } + } + k := polyEvalAtMinusOne(a) / math.Pow(2, float64(order)) + for i := range b { + b[i] *= k + } + return b, a, nil + } + b = binomialCoeffs(order) + k := polyEvalAtOne(a) / math.Pow(2, float64(order)) + for i := range b { + b[i] *= k + } + return b, a, nil +} + +// mulPolyReal multiplies two real polynomials in u = z^{-1} (index m +// is the coefficient of u^m). +func mulPolyReal(p, q []float64) []float64 { + out := make([]float64, len(p)+len(q)-1) + for i, pv := range p { + for j, qv := range q { + out[i+j] += pv * qv + } + } + return out +} + +// binomialCoeffs returns (1+u)^order as a coefficient list, built by +// exact rational steps so the symmetry C(n,m) = C(n,n−m) survives +// rounding. +func binomialCoeffs(order int) []float64 { + out := make([]float64, order+1) + out[0] = 1 + for k := range order { + out[k+1] = out[k] * float64(order-k) / float64(k+1) + } + return out +} + +// polyEvalAtMinusOne evaluates a u-polynomial at z = −1: the +// alternating sum of its coefficients. +func polyEvalAtMinusOne(p []float64) float64 { + s := 0.0 + for m, v := range p { + if m%2 == 1 { + s -= v + } else { + s += v + } + } + return s +} + +// polyEvalAtOne evaluates a u-polynomial at z = 1: the sum of its +// coefficients. +func polyEvalAtOne(p []float64) float64 { + s := 0.0 + for _, v := range p { + s += v + } + return s +} + +// FilterApply runs a direct-form-II-transposed filter over a rank-1 +// real signal: y[n] = Σ b_m·x[n−m] − Σ a_m·y[n−m]. The coefficient +// lists may differ in length (the shorter is zero-padded), a[0] must +// be non-zero and is normalised away, and the output has the same +// length as the input. Float32 and float64 inputs keep their dtype; +// int input widens to float64, because a filtered signal is not +// integral and truncating it would quantise the answer to zero +// wherever the response is small. This is the evaluator the +// Butterworth designs hand their coefficients to, and the sweep +// Filtfilt runs in both directions. +func FilterApply(b, a []float64, x *core.Array) (*core.Array, error) { + const name = "FilterApply" + bc, ac, out, err := filterPrepare(name, b, a, x) + if err != nil { + return nil, err + } + src := widenFloats(x) + dst := make([]float64, x.Len()) + filterSweep(bc, ac, src, dst) + outF := out.RawFloats() + if outF == nil || out.Strided() { + for i := range dst { + out.SetFloatAt(i, dst[i]) + } + } else { + copy(outF, dst) + } + return out, nil +} + +// filterPrepare runs the shared contract of the filter evaluators: +// rank-1 real input, non-empty coefficient lists with a non-zero +// leading denominator, coefficients normalised by a[0] and padded to +// one common length, and an output array keeping the input's float +// dtype. The normalised coefficients and the output come back +// together. +func filterPrepare(name string, b, a []float64, x *core.Array) (bc, ac []float64, out *core.Array, err error) { + if x.NDim() != 1 { + return nil, nil, nil, base.Errf("%s: needs a rank-1 signal, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, nil, nil, base.Errf("%s: complex signals are not supported", name) + } + if len(b) == 0 || len(a) == 0 { + return nil, nil, nil, base.Errf("%s: the coefficient lists must not be empty", name) + } + if a[0] == 0 { + return nil, nil, nil, base.Errf("%s: a[0] must be non-zero", name) + } + n := max(len(b), len(a)) + bc = make([]float64, n) + ac = make([]float64, n) + for i := range b { + bc[i] = b[i] / a[0] + } + ac[0] = 1 + for i := 1; i < len(a); i++ { + ac[i] = a[i] / a[0] + } + outDT := x.Dtype() + if outDT != core.Float && outDT != core.Float32 { + outDT = core.Float + } + out, oerr := core.Zeros(outDT, x.Len()) + if oerr != nil { + return nil, nil, nil, oerr + } + return bc, ac, out, nil +} + +// filterSweep is the direct-form-II-transposed kernel both filter +// evaluators run: one pass of the normalised coefficients over src, +// written to dst, which must not overlap src. Transposed delays: +// coefficient indices run 0..n−1, so the recurrence carries exactly +// n−1 registers; the last one takes no +state term, which is where a +// phantom extra delay would sneak in. +func filterSweep(bc, ac, src, dst []float64) { + n := len(bc) + state := make([]float64, max(n-1, 1)) + for i := range src { + xv := src[i] + yv := bc[0] * xv + if n > 1 { + yv += state[0] + for j := 1; j < n-1; j++ { + state[j-1] = bc[j]*xv - ac[j]*yv + state[j] + } + state[n-2] = bc[n-1]*xv - ac[n-1]*yv + } + dst[i] = yv + } +} diff --git a/signal/filter_dtype_pin_test.go b/signal/filter_dtype_pin_test.go new file mode 100644 index 0000000..1ba4d9f --- /dev/null +++ b/signal/filter_dtype_pin_test.go @@ -0,0 +1,92 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Dtype and extent pins for the filter path. + +// TestFilterApplyWidensIntInput pins the dtype contract: a filtered +// signal is not integral, so an int input must widen rather than have +// every sample truncated towards zero. +func TestFilterApplyWidensIntInput(t *testing.T) { + vals := make([]int64, 64) + for i := range vals { + vals[i] = 1000 + } + x, err := core.FromInts(vals, 64) + if err != nil { + t.Fatal(err) + } + out, err := FilterApply([]float64{0.5}, []float64{1}, x) + if err != nil { + t.Fatalf("FilterApply: %v", err) + } + if out.Dtype() != core.Float { + t.Fatalf("dtype = %s, want Float", out.Dtype()) + } + for i := range 64 { + if got := out.FloatAt(i); math.Abs(got-500) > 1e-12 { + t.Fatalf("out[%d] = %v, want 500: an int output would truncate", i, got) + } + } + // A float32 signal keeps its dtype. + x32, err := core.FromFloat32s([]float32{1, 2, 3, 4}, 4) + if err != nil { + t.Fatal(err) + } + out32, err := FilterApply([]float64{1}, []float64{1}, x32) + if err != nil { + t.Fatal(err) + } + if out32.Dtype() != core.Float32 { + t.Fatalf("float32 input gave %s, want Float32", out32.Dtype()) + } +} + +// TestFilterApplyUnitGain pins a filter that passes the signal +// through, the simplest check that the widening did not change the +// arithmetic. +func TestFilterApplyUnitGain(t *testing.T) { + src := []float64{1, -2, 3.5, 4, 0, -1} + x, err := core.FromFloats(src, len(src)) + if err != nil { + t.Fatal(err) + } + out, err := FilterApply([]float64{1}, []float64{1}, x) + if err != nil { + t.Fatal(err) + } + for i, want := range src { + if got := out.FloatAt(i); got != want { + t.Fatalf("out[%d] = %v, want %v", i, got, want) + } + } +} + +// TestSumKahanOnView pins the extent: the aliased payload of a rebased +// view is longer than its element count, and the accumulator must not +// walk past it. +func TestSumKahanOnView(t *testing.T) { + x, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 6) + if err != nil { + t.Fatal(err) + } + view, err := core.Slice(x, 0, 1, 4) + if err != nil { + t.Fatal(err) + } + got, err := SumKahan(view) + if err != nil { + t.Fatalf("SumKahan: %v", err) + } + if got != 9 { // 2 + 3 + 4 + t.Fatalf("SumKahan(view) = %v, want 9", got) + } +} diff --git a/signal/filter_test.go b/signal/filter_test.go new file mode 100644 index 0000000..40a203b --- /dev/null +++ b/signal/filter_test.go @@ -0,0 +1,349 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// TestSavitzkyGolayPreservesPolynomials checks the defining property: +// a signal that is a polynomial of degree ≤ order must pass through +// unchanged, edges included. +func TestSavitzkyGolayPreservesPolynomials(t *testing.T) { + const n = 40 + poly := make([]float64, n) + for i := range n { + x := float64(i) - 20 + poly[i] = 0.3*x*x*x - 1.5*x*x + 2*x - 7 + } + got, err := SavitzkyGolay(mustFloats(t, poly), 7, 3) + if err != nil { + t.Fatalf("SavitzkyGolay: %v", err) + } + for i := range n { + if math.Abs(got.FloatAt(i)-poly[i]) > 1e-9 { + t.Fatalf("cubic sample %d: %.12g, want %.12g", i, got.FloatAt(i), poly[i]) + } + } + quad := make([]float64, n) + for i := range n { + quad[i] = float64(i*i - 9*i + 4) + } + got, err = SavitzkyGolay(mustFloats(t, quad), 5, 2) + if err != nil { + t.Fatalf("SavitzkyGolay quadratic: %v", err) + } + for i := range n { + if math.Abs(got.FloatAt(i)-quad[i]) > 1e-9 { + t.Fatalf("quadratic sample %d: %.12g, want %.12g", i, got.FloatAt(i), quad[i]) + } + } + got, err = SavitzkyGolay(mustFloats(t, quad), 9, 4) + if err != nil { + t.Fatalf("SavitzkyGolay quadratic under order 4: %v", err) + } + for i := range n { + if math.Abs(got.FloatAt(i)-quad[i]) > 1e-9 { + t.Fatalf("quadratic under order 4, sample %d: %.12g, want %.12g", + i, got.FloatAt(i), quad[i]) + } + } +} + +// TestSavitzkyGolayConstant checks a flat signal stays exactly flat. +func TestSavitzkyGolayConstant(t *testing.T) { + flat := make([]float64, 15) + for i := range flat { + flat[i] = 3.25 + } + got, err := SavitzkyGolay(mustFloats(t, flat), 5, 1) + if err != nil { + t.Fatalf("SavitzkyGolay: %v", err) + } + for i := range flat { + if math.Abs(got.FloatAt(i)-3.25) > 1e-12 { + t.Fatalf("flat sample %d: %.16g, want 3.25", i, got.FloatAt(i)) + } + } +} + +// TestSavitzkyGolayDenoises checks the smoothed signal tracks a clean +// sine better than the noisy input in RMS, with deterministic noise. +func TestSavitzkyGolayDenoises(t *testing.T) { + const n = 200 + clean := make([]float64, n) + noisy := make([]float64, n) + for i := range n { + x := 2 * math.Pi * float64(i) / 40 + clean[i] = math.Sin(x) + 0.3*math.Sin(3*x) + // A deterministic hash-like noise, no generator needed. + noise := math.Sin(float64(i)*12.9898) * 43758.5453 + noise -= math.Floor(noise) + noisy[i] = clean[i] + 0.4*(2*noise-1) + } + got, err := SavitzkyGolay(mustFloats(t, noisy), 11, 3) + if err != nil { + t.Fatalf("SavitzkyGolay: %v", err) + } + rms := func(v []float64) float64 { + s := 0.0 + for i := range n { + d := v[i] - clean[i] + s += d * d + } + return math.Sqrt(s / n) + } + before := rms(noisy) + after := rms(got.RawFloats()) + if after >= before { + t.Fatalf("smoothing did not help: RMS %g after vs %g before", after, before) + } + // The signal's own 3× harmonic sits partly in the filter's + // passband, so halving the RMS is a fair bar, not more. + if after > 0.6*before { + t.Fatalf("RMS only fell to %g from %g, want lower", after, before) + } +} + +func TestSavitzkyGolayErrors(t *testing.T) { + ok := mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7}) + cases := []struct { + name string + data *core.Array + window, order int + }{ + {"even window", ok, 4, 1}, + {"tiny window", ok, 1, 0}, + {"window over length", ok, 9, 1}, + {"order at window", ok, 3, 3}, + {"order over half the window", ok, 5, 3}, + {"order just over half the window", ok, 7, 4}, + {"negative order", ok, 3, -1}, + {"empty signal", mustFloats(t, nil), 3, 1}, + {"rank 2", mustFloats(t, []float64{1, 2, 3, 4}, 2, 2), 3, 1}, + } + for _, c := range cases { + if _, err := SavitzkyGolay(c.data, c.window, c.order); err == nil { + t.Fatalf("%s: want an error", c.name) + } + } +} + +// TestSavitzkyGolayOrderAtHalfWindow pins the boundary: order == +// window/2 is the standard SG constraint and still fits, since the +// truncated edge windows then hold exactly order+1 samples. +func TestSavitzkyGolayOrderAtHalfWindow(t *testing.T) { + const n = 24 + quad := make([]float64, n) + for i := range n { + x := float64(i) - 12 + quad[i] = 0.5*x*x - 2*x + 1 + } + got, err := SavitzkyGolay(mustFloats(t, quad), 5, 2) + if err != nil { + t.Fatalf("SavitzkyGolay(5, 2): %v", err) + } + for i := range n { + if math.Abs(got.FloatAt(i)-quad[i]) > 1e-9 { + t.Fatalf("quadratic sample %d: %.12g, want %.12g", i, got.FloatAt(i), quad[i]) + } + } +} + +// cmag is |z| for the response checks. +func cmag(z complex128) float64 { return math.Hypot(real(z), imag(z)) } + +// freqResponse evaluates H(e^{iω}) from direct-form coefficients. +func freqResponse(b, a []float64, w float64) complex128 { + var num, den complex128 + for m, c := range b { + num += complex(c, 0) * complex(math.Cos(-w*float64(m)), math.Sin(-w*float64(m))) + } + for m, c := range a { + den += complex(c, 0) * complex(math.Cos(-w*float64(m)), math.Sin(-w*float64(m))) + } + return num / den +} + +// TestButterworthHalfPowerPoint pins the defining property: |H| at the +// cutoff is 1/√2 for every order, both directions. +func TestButterworthHalfPowerPoint(t *testing.T) { + const fs = 100.0 + const cutoff = 20.0 + wc := 2 * math.Pi * cutoff / fs + for _, order := range []int{1, 2, 3, 4, 5, 8} { + b, a, err := ButterworthLowPass(order, fs, cutoff) + if err != nil { + t.Fatalf("order %d: %v", order, err) + } + got := cmag(freqResponse(b, a, wc)) + if math.Abs(got-1/math.Sqrt2) > 1e-9 { + t.Fatalf("LP order %d: |H(wc)| = %.12f, want 1/sqrt2", order, got) + } + if math.Abs(cmag(freqResponse(b, a, 0))-1) > 1e-12 { + t.Fatalf("LP order %d: DC gain is not 1", order) + } + bh, ah, err := ButterworthHighPass(order, fs, cutoff) + if err != nil { + t.Fatalf("HP order %d: %v", order, err) + } + got = cmag(freqResponse(bh, ah, wc)) + if math.Abs(got-1/math.Sqrt2) > 1e-9 { + t.Fatalf("HP order %d: |H(wc)| = %.12f, want 1/sqrt2", order, got) + } + if math.Abs(cmag(freqResponse(bh, ah, math.Pi))-1) > 1e-9 { + t.Fatalf("HP order %d: Nyquist gain is not 1", order) + } + if cmag(freqResponse(bh, ah, 0)) > 1e-12 { + t.Fatalf("HP order %d: DC leakage %g", order, cmag(freqResponse(bh, ah, 0))) + } + } +} + +// TestButterworthMonotoneStopband pins the flatness contract: higher +// order suppresses the stopband strictly more at the same frequency. +func TestButterworthMonotoneStopband(t *testing.T) { + const fs, cutoff = 100.0, 10.0 + wStop := 2 * math.Pi * 40.0 / fs + prev := 1.0 + for _, order := range []int{1, 2, 4, 8} { + b, a, err := ButterworthLowPass(order, fs, cutoff) + if err != nil { + t.Fatalf("order %d: %v", order, err) + } + got := cmag(freqResponse(b, a, wStop)) + if got >= prev { + t.Fatalf("order %d stopband %.3e not below the previous %.3e", order, got, prev) + } + prev = got + } +} + +// TestFilterApplySteadyState pins the evaluator: a constant signal +// reaches the DC gain after the transient, and a stopband sine is +// attenuated by |H| there. +func TestFilterApplySteadyState(t *testing.T) { + const fs, cutoff = 200.0, 20.0 + b, a, err := ButterworthLowPass(4, fs, cutoff) + if err != nil { + t.Fatalf("ButterworthLowPass: %v", err) + } + const n = 400 + consts := make([]float64, n) + for i := range n { + consts[i] = 2.5 + } + ca, _ := core.FromFloats(consts, n) + out, err := FilterApply(b, a, ca) + if err != nil { + t.Fatalf("FilterApply: %v", err) + } + if math.Abs(out.FloatAt(n-1)-2.5) > 1e-9 { + t.Fatalf("steady state = %g, want 2.5", out.FloatAt(n-1)) + } + f := 60.0 + w := 2 * math.Pi * f / fs + want := cmag(freqResponse(b, a, w)) + sine := make([]float64, n) + for i := range n { + sine[i] = math.Sin(w * float64(i)) + } + sa, _ := core.FromFloats(sine, n) + out2, err := FilterApply(b, a, sa) + if err != nil { + t.Fatalf("FilterApply: %v", err) + } + // Amplitude by RMS over the tail: phase-agnostic, so a sampling + // rate of barely three points per period is fine. + ms := 0.0 + for i := 200; i < n; i++ { + ms += out2.FloatAt(i) * out2.FloatAt(i) + } + amp := math.Sqrt(2 * ms / float64(n-200)) + if math.Abs(amp-want) > 0.05*want+1e-6 { + t.Fatalf("stopband amplitude %.5f, |H| says %.5f", amp, want) + } +} + +// TestButterworthValidation pins the input gates. +func TestButterworthValidation(t *testing.T) { + if _, _, err := ButterworthLowPass(0, 100, 10); err == nil { + t.Error("order 0 accepted") + } + if _, _, err := ButterworthLowPass(2, 0, 10); err == nil { + t.Error("fs 0 accepted") + } + if _, _, err := ButterworthLowPass(2, 100, 0); err == nil { + t.Error("cutoff 0 accepted") + } + if _, _, err := ButterworthLowPass(2, 100, 50); err == nil { + t.Error("cutoff at Nyquist accepted") + } + x, _ := core.FromFloats([]float64{1, 2, 3}, 3) + if _, err := FilterApply(nil, []float64{1}, x); err == nil { + t.Error("empty b accepted") + } + if _, err := FilterApply([]float64{1}, []float64{0, 1}, x); err == nil { + t.Error("a[0] = 0 accepted") + } + cx, _ := core.FromComplexes([]complex128{1, 2}, 2) + if _, err := FilterApply([]float64{1}, []float64{1}, cx); err == nil { + t.Error("complex signal accepted") + } +} + +// TestSavitzkyGolayRefusesOverflowingNormalEquations pins the overflow +// gate: a window at the order boundary whose power sums overflow into +// Inf used to flow through the solver into NaN weights and come back +// as a smooth-looking all-NaN signal with no error. Window 167 at +// order 83 is the smallest boundary case whose top entry (83^166, +// about 1e319) is past MaxFloat64. +func TestSavitzkyGolayRefusesOverflowingNormalEquations(t *testing.T) { + x := mustFromFloats(t, make([]float64, 167), 167) + if _, err := SavitzkyGolay(x, 167, 83); err == nil || !strings.Contains(err.Error(), "overflow") { + t.Fatalf("SavitzkyGolay with window 167 at order 83: %v", err) + } +} + +// TestSavitzkyGolayWideWindowChunkEdges pins the chunk containment of +// the parallel sweep: with a window wider than two worker chunks the +// edge regions are wider than a chunk, and the edge loops used to run +// past their worker's own [ps, pe) into the neighbouring chunk, so two +// workers wrote the same output slot at once. The values agreed by +// determinism, but the access is a data race the memory model and the +// race detector both refuse: this test drives the geometry that makes +// the detector fire, and checks the values a quadratic-preserving fit +// must answer everywhere, edges included. +func TestSavitzkyGolayWideWindowChunkEdges(t *testing.T) { + const ( + n = 8192 + window = 4097 // edge = 2048, past the 1024-point worker chunk + order = 2 + ) + vals := make([]float64, n) + for i := range n { + x := float64(i) - n/2 + vals[i] = 0.5*x*x - 3*x + 7 + } + // Eight workers give chunks of 1024 points, so the chunk boundaries + // at 1024 and 6144 land inside the 2048-sample edge regions: the + // geometry the clamped loops must survive. + prev := engine.SetNumWorkers(8) + defer engine.SetNumWorkers(prev) + got, err := SavitzkyGolay(mustFloats(t, vals), window, order) + if err != nil { + t.Fatalf("SavitzkyGolay(%d, %d): %v", window, order, err) + } + for i := range n { + if math.Abs(got.FloatAt(i)-vals[i]) > 1e-6 { + t.Fatalf("quadratic sample %d: %.12g, want %.12g", i, got.FloatAt(i), vals[i]) + } + } +} diff --git a/signal/filterdesign.go b/signal/filterdesign.go new file mode 100644 index 0000000..01eafd3 --- /dev/null +++ b/signal/filterdesign.go @@ -0,0 +1,635 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "math/cmplx" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// IIR designs beyond the Butterworth pair: the Chebyshev equiripple +// family, the inverse Chebyshev, the elliptic (Cauer) designs, and +// the band shapes for every prototype. Each design walks the same +// road: an analog low-pass prototype at unit passband edge, the +// shape transformation of its pole-zero set with both band edges +// prewarped to the bilinear axis, and the bilinear mapping of each +// pole and zero. The mapping is exact, so the prototype's passband +// peak and edge attenuations land on the digital side at the mapped +// frequencies; that is what the tests pin. +// +// The coefficients come back in the same u = z⁻¹ convention as the +// Butterworth pair, with the denominator leading a one. Direct-form +// filtering loses digits as the order climbs, so past roughly order +// eight a design should be split into second-order sections by the +// caller; where that line sits is deliberately left to taste. + +// prototype is one analog low-pass at unit passband edge. Poles and +// finite zeroes carry exact conjugate symmetry; the zeroes a low-pass +// holds at infinity are counted, not listed. +type prototype struct { + poles []complex128 + zeros []complex128 + zerosAtInfinity int + gain float64 +} + +// chebyshev1 returns the type I prototype: the Butterworth circle +// squashed by the ripple's hyperbolic factor, so the passband swings +// between one and 1/sqrt(1+eps²), the peak normalised to one. The +// edge sits where the gain first drops to -rippleDB. +func chebyshev1(order int, rippleDB float64) prototype { + eps := math.Sqrt(math.Pow(10, rippleDB/10) - 1) + mu := math.Asinh(1/eps) / float64(order) + poles := make([]complex128, 0, order) + for k := range order { + theta := math.Pi * float64(2*k+1) / float64(2*order) + // One Sincos serves the pole's both factors: the pair is the + // one the separate Sin and Cos calls produced, so the pole + // keeps its bits. + sinT, cosT := math.Sincos(theta) + if math.Abs(cosT) < 1e-15 { + // The odd order's real pole, built exactly. + poles = append(poles, complex(-math.Sinh(mu), 0)) + continue + } + if cosT < 0 { + continue // the mirror angle's conjugate, already stored + } + p := complex(-math.Sinh(mu)*sinT, math.Cosh(mu)*cosT) + poles = append(poles, p, cmplx.Conj(p)) + } + gain := 1.0 + if order%2 == 0 { + gain = 1 / math.Sqrt(1+eps*eps) + } + return prototype{poles: poles, zerosAtInfinity: order, gain: gain} +} + +// chebyshev2 returns the inverse Chebyshev: zeroes on the imaginary +// axis at the reciprocals of the type I ripple points, poles at the +// reciprocals of type I poles built with the stopband's factor, so +// the stopband floor is exactly the demanded attenuation. The root +// lattice runs over the integers of one parity between −(n−1) and +// n−1: for odd orders it passes through zero, holding the real pole +// there, and the zero at infinity keeps m = 0 out of the zero list. +func chebyshev2(order int, stopbandDB float64) prototype { + de := 1 / math.Sqrt(math.Pow(10, stopbandDB/10)-1) + mu := math.Asinh(1/de) / float64(order) + poles := make([]complex128, 0, order) + zeros := make([]complex128, 0, order) + for m := -order + 1; m <= order-1; m += 2 { + theta := math.Pi * float64(m) / float64(2*order) + if m >= 0 { + p := -1 / cmplx.Sinh(complex(mu, theta)) + if m == 0 { + poles = append(poles, p) + } else { + poles = append(poles, p, cmplx.Conj(p)) + } + if m > 0 { + z := complex(0, 1/math.Sin(theta)) + zeros = append(zeros, z, cmplx.Conj(z)) + } + } + } + return prototype{poles: poles, zeros: zeros, + zerosAtInfinity: order - len(zeros), gain: 1} +} + +// cauer returns the elliptic prototype. The zeroes are cd of the +// quarter-period fractions, read through the library's Jacobi +// functions; the poles ride the addition theorem through the +// amplitude v0 solved from the stopband's discrimination; and the +// degree equation closes through the nome, so ripple and attenuation +// both land on spec at the first transition edge. +func cauer(order int, rippleDB, stopbandDB float64) prototype { + epsSq := math.Pow(10, rippleDB/10) - 1 + eps := math.Sqrt(epsSq) + m1 := epsSq / (math.Pow(10, stopbandDB/10) - 1) + m := ellipdeg(order, m1) + capk := core.EllipticKScalar(m) + + // Amplitudes along the quarter period: u_j = j·K/n for odd j when + // the order is even, even j when it is odd, where u = 0 holds the + // zero at infinity. + zeroes := make([]complex128, 0, order) + amplitudes := make([][3]float64, 0, order) + for j := 1 - order%2; j < order; j += 2 { + u := float64(j) * capk / float64(order) + s := jacobiScalar(core.JacobiSN, u, m) + c := jacobiScalar(core.JacobiCN, u, m) + d := jacobiScalar(core.JacobiDN, u, m) + if math.Abs(s) > 1e-12 { + z := complex(0, 1/(math.Sqrt(m)*s)) + zeroes = append(zeroes, z, cmplx.Conj(z)) + } + amplitudes = append(amplitudes, [3]float64{s, c, d}) + } + + // v0: the amplitude solving sc(v0, 1−m1) = 1/ε on the + // complementary parameter, through the identity sc(u, m) = + // tan(am(u, m)): v0 = F(atan(1/ε), 1−m1), scaled by the nome + // ratio the degree equation provides. The pole formula then reads + // its Jacobi amplitudes at v0 on the prototype's own parameter. + v0 := capk * core.EllipticFScalar(math.Atan(1/eps), 1-m1) / + (float64(order) * core.EllipticKScalar(m1)) + sv := jacobiScalar(core.JacobiSN, v0, 1-m) + cv := jacobiScalar(core.JacobiCN, v0, 1-m) + dv := jacobiScalar(core.JacobiDN, v0, 1-m) + + poles := make([]complex128, 0, order) + for _, a := range amplitudes { + s, c, d := a[0], a[1], a[2] + p := -complex(c*d*sv*cv, s*dv) / (1 - complex((d*sv)*(d*sv), 0)) + if math.Abs(imag(p)) < 1e-10 { + poles = append(poles, complex(real(p), 0)) + continue + } + poles = append(poles, p, cmplx.Conj(p)) + } + gain := 1.0 + if order%2 == 0 { + gain = 1 / math.Sqrt(1+epsSq) + } + return prototype{poles: poles, zeros: zeroes, + zerosAtInfinity: order - len(zeroes), gain: gain} +} + +// ellipdeg solves the degree equation n·K(m)/K'(m) = K(m1)/K'(m1) +// for m through the nome q = exp(−π·K'/K), whose theta product is +// accurate to double precision within the first eight powers. +func ellipdeg(n int, m1 float64) float64 { + k1 := core.EllipticKScalar(m1) + k1p := core.EllipticKScalar(1 - m1) + q := math.Pow(math.Exp(-math.Pi*k1p/k1), 1.0/float64(n)) + num, den := 1.0, 1.0 + for k := 1; k <= 7; k++ { + num += math.Pow(q, float64(k*(k+1))) + } + for k := 1; k <= 8; k++ { + den += 2 * math.Pow(q, float64(k*k)) + } + return 16 * q * math.Pow(num/den, 4) +} + +// jacobiScalar reads one Jacobi function at one point through the +// array implementation, which inverts the amplitude by bracketed +// Newton against Carlson's incomplete integral. +func jacobiScalar(f func(u *core.Array, m float64) (*core.Array, error), u, m float64) float64 { + arr, err := core.FromFloats([]float64{u}, 1) + if err != nil { + return math.NaN() + } + out, err := f(arr, m) + if err != nil { + return math.NaN() + } + return out.FloatAt(0) +} + +// mapEdge transforms one analog root of the unit-edge prototype into +// the prewarped target band. Band roots come back as the pair of a +// quadratic, the pairing surviving because the map sends conjugate +// pairs to conjugate pairs. +func mapEdge(s complex128, sh shape, w1, w2 float64) []complex128 { + switch sh { + case lowPass: + return []complex128{s * complex(w1, 0)} + case highPass: + return []complex128{complex(w1, 0) / s} + case bandPass: + // Scale to the half bandwidth, then split about the centre: + // the pair solves s² − BW·root·s + w0² = 0. + half := s * complex((w2-w1)/2, 0) + disc := cmplx.Sqrt(half*half - complex(w1*w2, 0)) + return []complex128{half + disc, half - disc} + default: // bandStop: invert to the half-bandwidth high-pass, then + // split about the centre the same way. + half := complex((w2-w1)/2, 0) / s + disc := cmplx.Sqrt(half*half - complex(w1*w2, 0)) + return []complex128{half + disc, half - disc} + } +} + +// shape picks the band the design passes. +type shape int + +const ( + lowPass shape = iota + highPass + bandPass + bandStop +) + +// design runs the shared road from prototype to coefficients: shape +// transformation of every root, bilinear mapping, assembly into real +// u = z⁻¹ polynomials, and the gain taken from the prototype itself +// at its DC, which every shape reaches through the mapping (the +// low-pass at DC, the high-pass at Nyquist, the band pair at the +// band centre and DC respectively). +func design(sh shape, proto prototype, w1, w2 float64) (b, a []float64, err error) { + const name = "filter design" + // Denominator roots: every pole, shape-transformed and mapped. + var poles []complex128 + for _, p := range proto.poles { + for _, s := range mapEdge(p, sh, w1, w2) { + if real(s) > 1e-7*(1+cmplx.Abs(s)) { + return nil, nil, base.Errf("%s: a pole escaped the left half-plane", name) + } + poles = append(poles, bilinear(s)) + } + } + // Numerator roots: every finite zero's image, then the zeroes at + // infinity: (1+u) factors for the low-pass, whose infinity maps + // to u = −1; (1−u) pairs for the band-pass, whose infinity maps + // to s = 0, u = 1; and one ±j·w0 conjugate pair per zero for the + // band-stop. The infinity factors are digital already; only the + // finite zeroes pass through the bilinear map. + var zeros []complex128 + for _, z := range proto.zeros { + zeros = append(zeros, mapEdge(z, sh, w1, w2)...) + } + zeros = bilinearAll(zeros) + zeros = append(zeros, infinityFactors(sh, proto.zerosAtInfinity, math.Sqrt(w1*w2))...) + b, err = assembleRoots(zeros) + if err != nil { + return nil, nil, err + } + a, err = assembleRoots(poles) + if err != nil { + return nil, nil, err + } + // Gain: the prototype's gain field is the response the design + // promises at its reference point (the passband peak, one or + // 1/sqrt(1+eps²) by order parity), so the numerator scales until + // the digital response at the mapped reference equals it. + var uRef complex128 + switch sh { + case lowPass, bandStop: + uRef = 1 + case highPass: + uRef = -1 + default: // bandPass: the band centre, the geometric mean edge; the + // conjugate side keeps the polynomial evaluation real. + uRef = cmplx.Conj(bilinear(complex(0, math.Sqrt(w1*w2)))) + } + hd := polyEvalC(b, uRef) / polyEvalC(a, uRef) + if hd == 0 { + return nil, nil, base.Errf("%s: the reference point carries no gain", name) + } + // The magnitude is what the gain field promises; the phase at the + // reference follows from the roots and is no business of the + // scaling. + scale := proto.gain / cmplx.Abs(hd) + for i := range b { + b[i] *= scale + } + return b, a, nil +} + +// infinityFactors names the numerator roots the prototype's zeroes at +// infinity turn into after the shape transformation: the low-pass +// zeroes at s = ∞ land at u = −1; the high-pass zeroes at s = 0 land +// at u = +1; the band-pass substitution squares its frequency, so +// every infinity zero becomes a double zero at s = 0, two (1−u) +// factors; and the band-stop turns every infinity zero into the +// conjugate pair ±j·w0, the roots of s² + w0² = 0. +func infinityFactors(sh shape, atInfinity int, w0 float64) []complex128 { + switch sh { + case lowPass: + roots := make([]complex128, 0, atInfinity) + for range atInfinity { + roots = append(roots, -1) + } + return roots + case bandPass: + // The substitution's s = (1−u)/(1+u) leaves the numerator as + // (1−u²)^N: half the roots at u = 1, half at u = −1. + roots := make([]complex128, 0, 2*atInfinity) + for range atInfinity { + roots = append(roots, 1, -1) + } + return roots + case bandStop: + roots := make([]complex128, 0, 2*atInfinity) + for range atInfinity { + d := bilinear(complex(0, w0)) + roots = append(roots, d, cmplx.Conj(d)) + } + return roots + default: // highPass + roots := make([]complex128, 0, atInfinity) + for range atInfinity { + roots = append(roots, 1) + } + return roots + } +} + +// assembleRoots factors digital roots into a real polynomial in u = +// z⁻¹: conjugate pairs become the real quadratic +// (1 − z·u)(1 − z̄·u) = 1 − 2Re(z)·u + |z|²·u², real roots the linear +// factor (1 − z·u). The pairing matches each root against the +// remaining roots' conjugates, so it does not depend on the order +// the roots arrived in. +func assembleRoots(roots []complex128) ([]float64, error) { + const name = "filter design" + poly := []float64{1} + used := make([]bool, len(roots)) + for i, r := range roots { + if used[i] { + continue + } + if math.Abs(imag(r)) < 1e-9 { + used[i] = true + poly = mulPolyReal(poly, []float64{1, -real(r)}) + continue + } + // The nearest conjugate partner among the unused roots. + partner := -1 + best := math.Inf(1) + for j := i + 1; j < len(roots); j++ { + if used[j] { + continue + } + if d := cmplx.Abs(roots[j] - cmplx.Conj(r)); d < best { + best, partner = d, j + } + } + if partner < 0 || best > 1e-6*cmplx.Abs(r) { + return nil, base.Errf("%s: the roots lost their conjugate symmetry", name) + } + used[partner] = true + poly = mulPolyReal(poly, []float64{1, -2 * real(r), real(r * cmplx.Conj(r))}) + } + return poly, nil +} + +// bilinearAll maps a root list through the bilinear transform. +func bilinearAll(roots []complex128) []complex128 { + out := make([]complex128, len(roots)) + for i, s := range roots { + out[i] = bilinear(s) + } + return out +} + +// bilinear maps an analog root to its digital image, the T = 2 +// sampling the prewarp assumes. +func bilinear(s complex128) complex128 { + return (1 + s) / (1 - s) +} + +// polyEvalC evaluates a u = z⁻¹ polynomial at one complex point. +func polyEvalC(poly []float64, u complex128) complex128 { + total := complex(0, 0) + for _, p := range slices.Backward(poly) { + total = total*u + complex(p, 0) + } + return total +} + +// The public designs. Each validates its arguments, prewarps the +// edges, and hands the shared road its prototype. + +// designArgs bundles the validated prewarped edges for one call. +func designArgs(name string, order int, fs, edge1, edge2 float64, sh shape) (w1, w2 float64, err error) { + if order < 1 { + return 0, 0, base.Errf("%s: the order must be at least 1, got %d", name, order) + } + if !(fs > 0) || math.IsInf(fs, 0) { + return 0, 0, base.Errf("%s: fs must be positive and finite, got %g", name, fs) + } + if !(edge1 > 0) || edge1 >= fs/2 { + return 0, 0, base.Errf("%s: the edge must lie in (0, fs/2), got %g for fs %g", name, edge1, fs) + } + w1 = math.Tan(math.Pi * edge1 / fs) + if sh == bandPass || sh == bandStop { + if !(edge2 > edge1) || edge2 >= fs/2 { + return 0, 0, base.Errf("%s: the band must span (edge1, edge2) inside (0, fs/2), got %g, %g for fs %g", + name, edge1, edge2, fs) + } + w2 = math.Tan(math.Pi * edge2 / fs) + } + return w1, w2, nil +} + +// ChebyshevLowPass designs an order-N type I Chebyshev low-pass at fs +// hertz with its ripple in decibels: the passband oscillates between +// 0 and -rippleDB, the edge is the last touch of -rippleDB, and the +// stopband rolls off as fast as that budget allows. +func ChebyshevLowPass(order int, fs, cutoff, rippleDB float64) (b, a []float64, err error) { + const name = "ChebyshevLowPass" + if rippleDB <= 0 { + return nil, nil, base.Errf("%s: the ripple must be positive decibels, got %g", name, rippleDB) + } + w, _, err := designArgs(name, order, fs, cutoff, 0, lowPass) + if err != nil { + return nil, nil, err + } + return design(lowPass, chebyshev1(order, rippleDB), w, 0) +} + +// ChebyshevHighPass is the type I mirror: the same equiripple +// passband above the edge, rolling off below it. +func ChebyshevHighPass(order int, fs, cutoff, rippleDB float64) (b, a []float64, err error) { + const name = "ChebyshevHighPass" + if rippleDB <= 0 { + return nil, nil, base.Errf("%s: the ripple must be positive decibels, got %g", name, rippleDB) + } + w, _, err := designArgs(name, order, fs, cutoff, 0, highPass) + if err != nil { + return nil, nil, err + } + return design(highPass, chebyshev1(order, rippleDB), w, 0) +} + +// InverseChebyshevLowPass designs the type II low-pass: a flat +// passband through the edge, with the stopband bottoming out at +// -stopbandDB and equiripple beyond it. +func InverseChebyshevLowPass(order int, fs, cutoff, stopbandDB float64) (b, a []float64, err error) { + const name = "InverseChebyshevLowPass" + if stopbandDB <= 0 { + return nil, nil, base.Errf("%s: the stopband attenuation must be positive decibels, got %g", name, stopbandDB) + } + w, _, err := designArgs(name, order, fs, cutoff, 0, lowPass) + if err != nil { + return nil, nil, err + } + return design(lowPass, chebyshev2(order, stopbandDB), w, 0) +} + +// InverseChebyshevHighPass is the type II mirror above the edge. +func InverseChebyshevHighPass(order int, fs, cutoff, stopbandDB float64) (b, a []float64, err error) { + const name = "InverseChebyshevHighPass" + if stopbandDB <= 0 { + return nil, nil, base.Errf("%s: the stopband attenuation must be positive decibels, got %g", name, stopbandDB) + } + w, _, err := designArgs(name, order, fs, cutoff, 0, highPass) + if err != nil { + return nil, nil, err + } + return design(highPass, chebyshev2(order, stopbandDB), w, 0) +} + +// CauerLowPass designs the elliptic low-pass: equiripple in the +// passband within rippleDB and equiripple stopband not above +// -stopbandDB, with the narrowest transition of any design at the +// order. The zeroes sit in the stopband, finite and on the unit +// circle after mapping. +func CauerLowPass(order int, fs, cutoff, rippleDB, stopbandDB float64) (b, a []float64, err error) { + const name = "CauerLowPass" + if rippleDB <= 0 || stopbandDB <= rippleDB { + return nil, nil, base.Errf("%s: the ripple must be positive and the attenuation larger, got %g and %g", + name, rippleDB, stopbandDB) + } + w, _, err := designArgs(name, order, fs, cutoff, 0, lowPass) + if err != nil { + return nil, nil, err + } + return design(lowPass, cauer(order, rippleDB, stopbandDB), w, 0) +} + +// CauerHighPass is the elliptic mirror above the edge. +func CauerHighPass(order int, fs, cutoff, rippleDB, stopbandDB float64) (b, a []float64, err error) { + const name = "CauerHighPass" + if rippleDB <= 0 || stopbandDB <= rippleDB { + return nil, nil, base.Errf("%s: the ripple must be positive and the attenuation larger, got %g and %g", + name, rippleDB, stopbandDB) + } + w, _, err := designArgs(name, order, fs, cutoff, 0, highPass) + if err != nil { + return nil, nil, err + } + return design(highPass, cauer(order, rippleDB, stopbandDB), w, 0) +} + +// ChebyshevBandPass designs the type I band-pass spanning edge1 to +// edge2: the prototype's order doubles through the band move. +func ChebyshevBandPass(order int, fs, edge1, edge2, rippleDB float64) (b, a []float64, err error) { + const name = "ChebyshevBandPass" + if rippleDB <= 0 { + return nil, nil, base.Errf("%s: the ripple must be positive decibels, got %g", name, rippleDB) + } + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandPass) + if err != nil { + return nil, nil, err + } + return design(bandPass, chebyshev1(order, rippleDB), w1, w2) +} + +// ChebyshevBandStop designs the type I band-stop. +func ChebyshevBandStop(order int, fs, edge1, edge2, rippleDB float64) (b, a []float64, err error) { + const name = "ChebyshevBandStop" + if rippleDB <= 0 { + return nil, nil, base.Errf("%s: the ripple must be positive decibels, got %g", name, rippleDB) + } + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandStop) + if err != nil { + return nil, nil, err + } + return design(bandStop, chebyshev1(order, rippleDB), w1, w2) +} + +// InverseChebyshevBandPass designs the type II band-pass. +func InverseChebyshevBandPass(order int, fs, edge1, edge2, stopbandDB float64) (b, a []float64, err error) { + const name = "InverseChebyshevBandPass" + if stopbandDB <= 0 { + return nil, nil, base.Errf("%s: the stopband attenuation must be positive decibels, got %g", name, stopbandDB) + } + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandPass) + if err != nil { + return nil, nil, err + } + return design(bandPass, chebyshev2(order, stopbandDB), w1, w2) +} + +// InverseChebyshevBandStop designs the type II band-stop. +func InverseChebyshevBandStop(order int, fs, edge1, edge2, stopbandDB float64) (b, a []float64, err error) { + const name = "InverseChebyshevBandStop" + if stopbandDB <= 0 { + return nil, nil, base.Errf("%s: the stopband attenuation must be positive decibels, got %g", name, stopbandDB) + } + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandStop) + if err != nil { + return nil, nil, err + } + return design(bandStop, chebyshev2(order, stopbandDB), w1, w2) +} + +// CauerBandPass designs the elliptic band-pass. +func CauerBandPass(order int, fs, edge1, edge2, rippleDB, stopbandDB float64) (b, a []float64, err error) { + const name = "CauerBandPass" + if rippleDB <= 0 || stopbandDB <= rippleDB { + return nil, nil, base.Errf("%s: the ripple must be positive and the attenuation larger, got %g and %g", + name, rippleDB, stopbandDB) + } + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandPass) + if err != nil { + return nil, nil, err + } + return design(bandPass, cauer(order, rippleDB, stopbandDB), w1, w2) +} + +// CauerBandStop designs the elliptic band-stop. +func CauerBandStop(order int, fs, edge1, edge2, rippleDB, stopbandDB float64) (b, a []float64, err error) { + const name = "CauerBandStop" + if rippleDB <= 0 || stopbandDB <= rippleDB { + return nil, nil, base.Errf("%s: the ripple must be positive and the attenuation larger, got %g and %g", + name, rippleDB, stopbandDB) + } + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandStop) + if err != nil { + return nil, nil, err + } + return design(bandStop, cauer(order, rippleDB, stopbandDB), w1, w2) +} + +// ButterworthBandPass designs the maximally flat band-pass. +func ButterworthBandPass(order int, fs, edge1, edge2 float64) (b, a []float64, err error) { + const name = "ButterworthBandPass" + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandPass) + if err != nil { + return nil, nil, err + } + return design(bandPass, butterworthPrototype(order), w1, w2) +} + +// ButterworthBandStop designs the maximally flat band-stop. +func ButterworthBandStop(order int, fs, edge1, edge2 float64) (b, a []float64, err error) { + const name = "ButterworthBandStop" + w1, w2, err := designArgs(name, order, fs, edge1, edge2, bandStop) + if err != nil { + return nil, nil, err + } + return design(bandStop, butterworthPrototype(order), w1, w2) +} + +// butterworthPrototype rebuilds the maximally flat poles in the +// prototype shape, so the band shapes share the same road; the +// existing ButterworthLowPass and ButterworthHighPass keep their own +// pinned implementations untouched. +func butterworthPrototype(order int) prototype { + poles := make([]complex128, 0, order) + for k := range order { + theta := math.Pi * float64(2*k+1) / float64(2*order) + // One Sincos serves the pole's both factors, the same pair the + // separate calls produced. + sinT, cosT := math.Sincos(theta) + if math.Abs(cosT) < 1e-15 { + poles = append(poles, complex(-1, 0)) + continue + } + if cosT < 0 { + continue // the mirror angle's conjugate, already stored + } + p := complex(-sinT, cosT) + poles = append(poles, p, cmplx.Conj(p)) + } + return prototype{poles: poles, zerosAtInfinity: order, gain: 1} +} diff --git a/signal/filterdesign_test.go b/signal/filterdesign_test.go new file mode 100644 index 0000000..f56e977 --- /dev/null +++ b/signal/filterdesign_test.go @@ -0,0 +1,217 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" +) + +// designResponse evaluates |H| of a transfer function at frequency f +// of fs, through the u = e^{−jω} polynomial the coefficients are. +func designResponse(b, a []float64, f, fs float64) float64 { + u := cmplxExp(-2 * math.Pi * f / fs) + h := polyEvalC(b, u) / polyEvalC(a, u) + return cmplxAbs(h) +} + +func cmplxExp(w float64) complex128 { + return complex(math.Cos(w), math.Sin(w)) +} + +func cmplxAbs(c complex128) float64 { + return math.Hypot(real(c), imag(c)) +} + +// TestFilterDesignsAgainstReferenceValues pins the responses against +// independently computed reference values (the same prototypes at the +// same orders, edges and ripple budgets, evaluated by direct +// polynomial division), covering every prototype in +// low-pass, high-pass, band-pass and band-stop shapes. The agreement +// proves the pole sets, the shape transformations and the gain +// handling at once. +func TestFilterDesignsAgainstReferenceValues(t *testing.T) { + const fs = 1000.0 + cases := []struct { + name string + run func() ([]float64, []float64, error) + freqs []float64 + want []float64 + }{ + {"cheby1-4-1db-100-low", func() ([]float64, []float64, error) { + return ChebyshevLowPass(4, fs, 100, 1) + }, []float64{10, 50, 100, 150, 300}, + []float64{9.046325269689e-01, 9.748546119130e-01, 8.912509381337e-01, 6.601335374641e-02, 8.075980703015e-04}}, + {"cheby2-4-40db-150-low", func() ([]float64, []float64, error) { + return InverseChebyshevLowPass(4, fs, 150, 40) + }, []float64{10, 100, 150, 200, 400}, + []float64{9.999999835157e-01, 2.847673225125e-01, 1.000000000000e-02, 9.994652720796e-03, 7.867574785234e-03}}, + {"ellip-5-1-60-120-low", func() ([]float64, []float64, error) { + return CauerLowPass(5, fs, 120, 1, 60) + }, []float64{20, 100, 120, 130, 250}, + []float64{9.483272809069e-01, 8.937539072858e-01, 8.912509381337e-01, 3.213042804039e-01, 2.027437669190e-04}}, + {"cheby1-6-0.5db-80-200-bp", func() ([]float64, []float64, error) { + return ChebyshevBandPass(6, fs, 80, 200, 0.5) + }, []float64{20, 90, 140, 190, 300, 450}, + []float64{1.674879454545e-06, 9.924739465269e-01, 9.800207063177e-01, 9.441608644607e-01, 3.307968386000e-04, 1.574651671228e-08}}, + {"butter-4-100-300-bs", func() ([]float64, []float64, error) { + return ButterworthBandStop(4, fs, 100, 300) + }, []float64{20, 90, 200, 310, 450}, + []float64{9.999998769440e-01, 8.934996852832e-01, 1.242247996472e-04, 8.354477748278e-01, 9.999996762445e-01}}, + {"ellip-4-2-50-120-high", func() ([]float64, []float64, error) { + return CauerHighPass(4, fs, 120, 2, 50) + }, []float64{30, 60, 120, 200, 400}, + []float64{4.401377613614e-05, 2.422743236491e-03, 7.943282347243e-01, 9.263786543625e-01, 8.261049684527e-01}}, + } + for _, c := range cases { + b, a, err := c.run() + if err != nil { + t.Fatalf("%s: %v", c.name, err) + } + for i, f := range c.freqs { + got := designResponse(b, a, f, fs) + want := c.want[i] + // Points past a hundred dB down sit in the direct form's + // cancellation noise: both answers are "essentially zero" + // there, and only the shallow points carry information. + if want < 1e-5 { + if got > 1e-4 { + t.Fatalf("%s at %g Hz: |H| = %.12g, want below 1e-4", c.name, f, got) + } + continue + } + if math.Abs(got-want) > 1e-7*math.Max(1, want) { + t.Fatalf("%s at %g Hz: |H| = %.12g, want %.12g", c.name, f, got, want) + } + } + } + // Property checks written from the design parameters themselves: + // the gain the prototype promises at the mapped edges, and the + // stopband bound the attenuation demands. Nothing is pinned to + // the implementation's output; every expectation follows from the + // decibel budgets. + // The inverse Chebyshev high-pass bottoms out at exactly + // −stopbandDB at its edge: a 40 dB demand means a gain of + // 10^(−40/20) = 0.01 there. + b, a, err := InverseChebyshevHighPass(4, fs, 120, 40) + if err != nil { + t.Fatalf("InverseChebyshevHighPass: %v", err) + } + if got := designResponse(b, a, 120, fs); math.Abs(got-0.01) > 1e-7 { + t.Fatalf("InverseChebyshevHighPass at the edge: |H| = %.12g, want the −40 dB floor 0.01", got) + } + // The type I band-stop touches −rippleDB at both band edges: a + // 1 dB ripple means 10^(−1/20) there. + edgeGain := math.Pow(10, -1.0/20) + b, a, err = ChebyshevBandStop(4, fs, 200, 400, 1) + if err != nil { + t.Fatalf("ChebyshevBandStop: %v", err) + } + for _, f := range []float64{200, 400} { + if got := designResponse(b, a, f, fs); math.Abs(got-edgeGain) > 1e-7 { + t.Fatalf("ChebyshevBandStop at %g Hz: |H| = %.12g, want the −1 dB edge gain %.12g", f, got, edgeGain) + } + } + // The elliptic band-pass keeps the whole stopband at or under the + // demanded −50 dB. The probes sit clear of the transitions and of + // the mirrored passband past Nyquist: for fs = 1000 the true + // stopband runs below the lower edge and between the upper edge + // and 500 Hz, where the equiripple floor touches the bound. + stopGain := math.Pow(10, -50.0/20) + b, a, err = CauerBandPass(4, fs, 200, 400, 1, 50) + if err != nil { + t.Fatalf("CauerBandPass: %v", err) + } + for _, f := range []float64{30, 50, 100, 450, 480} { + if got := designResponse(b, a, f, fs); got > stopGain { + t.Fatalf("CauerBandPass at %g Hz: |H| = %.12g exceeds the −50 dB bound %.12g", f, got, stopGain) + } + } +} + +// TestFilterDesignStability checks that every design's poles sit +// inside the unit circle, the property direct-form filtering lives +// and dies by. +func TestFilterDesignStability(t *testing.T) { + const fs = 2000.0 + designs := []struct { + name string + run func() ([]float64, []float64, error) + }{ + {"cheby1-low", func() ([]float64, []float64, error) { return ChebyshevLowPass(6, fs, 200, 2) }}, + {"cheby1-high", func() ([]float64, []float64, error) { return ChebyshevHighPass(5, fs, 300, 1) }}, + {"cheby2-low", func() ([]float64, []float64, error) { return InverseChebyshevLowPass(6, fs, 400, 50) }}, + {"cheby2-bp", func() ([]float64, []float64, error) { return InverseChebyshevBandPass(4, fs, 200, 600, 45) }}, + {"cheby2-high", func() ([]float64, []float64, error) { return InverseChebyshevHighPass(6, fs, 300, 40) }}, + {"cheby1-bs", func() ([]float64, []float64, error) { return ChebyshevBandStop(4, fs, 200, 500, 1) }}, + {"cauer-low", func() ([]float64, []float64, error) { return CauerLowPass(6, fs, 250, 1, 55) }}, + {"cauer-bs", func() ([]float64, []float64, error) { return CauerBandStop(3, fs, 300, 700, 1, 50) }}, + {"cauer-bp", func() ([]float64, []float64, error) { return CauerBandPass(4, fs, 200, 500, 1, 50) }}, + {"butter-bs", func() ([]float64, []float64, error) { return ButterworthBandStop(5, fs, 150, 500) }}, + {"butter-bp", func() ([]float64, []float64, error) { return ButterworthBandPass(4, fs, 100, 400) }}, + } + // The denominator's roots sit inside the unit circle exactly when + // the reflection coefficients of the Schur-Cohn recursion stay + // under one in magnitude; the recursion needs no root finding. + for _, d := range designs { + _, a, err := d.run() + if err != nil { + t.Fatalf("%s: %v", d.name, err) + } + if !schurCohnStable(a) { + t.Fatalf("%s: denominator is not minimum phase", d.name) + } + } +} + +// schurCohnStable reports whether the denominator polynomial's roots +// all lie strictly inside the unit circle, by the Schur-Cohn +// recursion: the coefficients already lead with the monic one. +func schurCohnStable(a []float64) bool { + n := len(a) - 1 + if n < 1 || a[0] != 1 { + return false + } + d := append([]float64(nil), a...) + for n > 0 { + last := d[n] + if math.Abs(last) >= 1 { + return false + } + next := make([]float64, n) + next[0] = 1 + for i := 1; i < n; i++ { + next[i] = (d[i] - last*d[n-i]) / (1 - last*last) + } + d = next + n-- + } + return true +} + +// TestFilterDesignRefusals checks the loud rejections on orders, +// edges and decibel budgets. +func TestFilterDesignRefusals(t *testing.T) { + if _, _, err := ChebyshevLowPass(0, 1000, 100, 1); err == nil { + t.Fatal("order zero accepted") + } + if _, _, err := ChebyshevLowPass(4, 1000, 600, 1); err == nil { + t.Fatal("edge past Nyquist accepted") + } + if _, _, err := ChebyshevLowPass(4, 1000, 100, -1); err == nil { + t.Fatal("negative ripple accepted") + } + if _, _, err := CauerLowPass(4, 1000, 100, 3, 2); err == nil { + t.Fatal("attenuation below ripple accepted") + } + if _, _, err := ChebyshevBandPass(4, 1000, 300, 200, 1); err == nil { + t.Fatal("reversed band accepted") + } + if _, _, err := InverseChebyshevBandStop(4, 1000, 100, 900, 40); err == nil { + t.Fatal("edge past Nyquist accepted for the band-stop") + } + if _, _, err := ButterworthBandPass(2, 1000, 100, 0); err == nil { + t.Fatal("zero edge accepted") + } +} diff --git a/signal/filtfilt.go b/signal/filtfilt.go new file mode 100644 index 0000000..e3e77d9 --- /dev/null +++ b/signal/filtfilt.go @@ -0,0 +1,87 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Zero-phase filtering. A causal IIR filter delays every +// feature by its group delay; running the same filter backwards over +// its own output doubles the magnitude response and cancels the phase +// exactly, because the reversed pass carries the conjugated transfer +// function. The cost is the edge question: the first pass starts from +// rest against a signal that did not, and its start-up transient would +// otherwise bleed into the answer's first samples. + +// Filtfilt filters the rank-1 real signal x forwards and backwards +// with the same transfer function the filter designs hand out (b over +// a, a[0] non-zero and normalised away), and returns a result the +// length of x whose magnitude response is the square of the one-pass +// filter's and whose phase is zero: a sinusoid in the passband comes +// out aligned with its input, not lagged, and the group delay at +// every frequency is 0 samples. The dtype contract is FilterApply's: +// float32 and float64 keep their dtype, other real dtypes widen to +// float64. +// +// Edge initialisation: both ends of the signal are extended by an +// even reflection of pad = 3·(nfilt−1) samples (nfilt the longer +// coefficient list), the edge sample not repeated, so the extension +// is continuous in value at both seams. The pad length is the +// standard three times the filter's memory: for a stable filter the +// start-up transient decays like the impulse response tail, which +// 3·(nfilt−1) samples of a direct-form kernel drive below the rounding +// floor for every design this package produces. The forward pass runs +// over the padded signal, the signal is time-reversed, the second +// pass runs, and the pad region of both ends is cropped away, so the +// residual seam error, a slope kink the reflection cannot hide from a +// filter that differentiates, stays outside the returned range. A +// signal no longer than the pad (length ≤ 3·(nfilt−1)) leaves nothing +// to return once both transient regions are excluded and is refused. +func Filtfilt(b, a []float64, x *core.Array) (*core.Array, error) { + const name = "Filtfilt" + bc, ac, out, err := filterPrepare(name, b, a, x) + if err != nil { + return nil, err + } + n := x.Len() + nfilt := max(len(b), len(a)) + pad := 3 * (nfilt - 1) + if n <= pad { + return nil, base.Errf("%s: the length-%d signal must exceed the pad length 3·(%d−1) = %d this filter needs", + name, n, nfilt, pad) + } + src := widenFloats(x) + // The reflected extension: left pad holds x[pad], x[pad−1], …, + // x[1] in reverse order, the right pad mirrors it, and the seam at + // either end repeats no sample (whole-sample symmetric even + // reflection). + padded := make([]float64, n+2*pad) + copy(padded[pad:pad+n], src) + for k := range pad { + padded[k] = src[pad-k] + padded[pad+n+k] = src[n-2-k] + } + p := len(padded) + fwd := make([]float64, p) + filterSweep(bc, ac, padded, fwd) + rev := make([]float64, p) + for i := range p { + rev[i] = fwd[p-1-i] + } + back := make([]float64, p) + filterSweep(bc, ac, rev, back) + outF := out.RawFloats() + if outF == nil || out.Strided() { + for t := range n { + out.SetFloatAt(t, back[p-1-pad-t]) + } + } else { + for t := range n { + outF[t] = back[p-1-pad-t] + } + } + return out, nil +} diff --git a/signal/filtfilt_test.go b/signal/filtfilt_test.go new file mode 100644 index 0000000..1958664 --- /dev/null +++ b/signal/filtfilt_test.go @@ -0,0 +1,273 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// filtfiltTestFilter is the shared design of the pins: a fourth-order +// Butterworth low-pass at a fifth of the sampling rate, stable and +// mildly selective. +func filtfiltTestFilter(t *testing.T) (b, a []float64) { + t.Helper() + b, a, err := ButterworthLowPass(4, 100, 20) + if err != nil { + t.Fatalf("ButterworthLowPass: %v", err) + } + return b, a +} + +// TestFiltfiltZeroPhase pins the defining property from two sides. +// First, the response to an interior impulse is exactly symmetric +// about it, which is the time-domain face of zero phase. Second, a +// passband sinusoid comes out with no phase shift against its input, +// the frequency-domain face, where the one-pass filter visibly lags. +func TestFiltfiltZeroPhase(t *testing.T) { + b, a := filtfiltTestFilter(t) + const ( + n = 501 + m = 250 + ) + vals := make([]float64, n) + vals[m] = 1 + got, err := Filtfilt(b, a, mustFloats(t, vals)) + if err != nil { + t.Fatalf("Filtfilt: %v", err) + } + // The pad regions (twice the filter's pad at each end, once per + // pass) host the seam transients; everything inside stays clean. + const clean = 100 + for d := 1; d <= clean; d++ { + left := got.FloatAt(m - d) + right := got.FloatAt(m + d) + if math.Abs(left-right) > 1e-12*math.Max(1, math.Abs(left)) { + t.Fatalf("impulse response asymmetric at offset %d: %g vs %g", d, left, right) + } + } + // Phase of a passband sinusoid: fit A·sin + B·cos over the + // interior and read the phase back. Zero phase means B ≈ 0. + // wTone is the angular frequency of the 5 hertz test tone at + // 100 hertz. + const wTone = 2 * math.Pi * 5.0 / 100.0 + fitPhase := func(y *core.Array, from, to int) (amp, phase float64) { + var ss, sc, cc, ys, yc float64 + for i := from; i < to; i++ { + s, c := math.Sincos(wTone * float64(i)) + v := y.FloatAt(i) + ss += s * s + sc += s * c + cc += c * c + ys += v * s + yc += v * c + } + a1 := (cc*ys - sc*yc) / (ss*cc - sc*sc) + b1 := (ss*yc - sc*ys) / (ss*cc - sc*sc) + return math.Hypot(a1, b1), math.Atan2(b1, a1) + } + sine := make([]float64, n) + for i := range sine { + sine[i] = math.Sin(wTone * float64(i)) + } + out, err := Filtfilt(b, a, mustFloats(t, sine)) + if err != nil { + t.Fatalf("Filtfilt: %v", err) + } + amp, phase := fitPhase(out, clean, n-clean) + if math.Abs(phase) > 1e-9 { + t.Fatalf("passband phase %g rad, want 0", phase) + } + wantAmp := cmag(freqResponse(b, a, wTone)) + wantAmp *= wantAmp + if math.Abs(amp-wantAmp) > 1e-9 { + t.Fatalf("passband amplitude %.10f, want |H|^2 = %.10f", amp, wantAmp) + } + // The one-pass contrast: same fit shows the causal group delay. + one, err := FilterApply(b, a, mustFloats(t, sine)) + if err != nil { + t.Fatalf("FilterApply: %v", err) + } + _, onePhase := fitPhase(one, clean, n-clean) + if math.Abs(onePhase) < 0.1 { + t.Fatalf("one-pass phase %g rad should be plainly nonzero", onePhase) + } +} + +// TestFiltfiltMatchesManualComposition pins the architecture: the +// documented pad, two FilterApply passes with a time-reversal between +// them, and the crop reproduce Filtfilt sample for sample, bit for +// bit. +func TestFiltfiltMatchesManualComposition(t *testing.T) { + b, a := filtfiltTestFilter(t) + g := core.NewGenerator(11) + const n = 300 + vals := make([]float64, n) + for i := range vals { + vals[i] = g.NormalUnit() + } + x := mustFloats(t, vals) + got, err := Filtfilt(b, a, x) + if err != nil { + t.Fatalf("Filtfilt: %v", err) + } + // The pad of the documentation, built by hand. + nfilt := max(len(b), len(a)) + pad := 3 * (nfilt - 1) + p := n + 2*pad + padded := make([]float64, p) + copy(padded[pad:pad+n], vals) + for k := range pad { + padded[k] = vals[pad-k] + padded[pad+n+k] = vals[n-2-k] + } + y1, err := FilterApply(b, a, mustFloats(t, padded)) + if err != nil { + t.Fatalf("FilterApply forward: %v", err) + } + rev := make([]float64, p) + for i := range p { + rev[i] = y1.FloatAt(p - 1 - i) + } + y2, err := FilterApply(b, a, mustFloats(t, rev)) + if err != nil { + t.Fatalf("FilterApply backward: %v", err) + } + for ts := range n { + want := y2.FloatAt(p - 1 - pad - ts) + if got.FloatAt(ts) != want { + t.Fatalf("sample %d: %v, want %v", ts, got.FloatAt(ts), want) + } + } +} + +// TestFiltfiltSuppressesMoreThanOnePass pins the doubled magnitude +// response: in the stopband the filtfilt attenuation is the square of +// the one-pass attenuation at the same frequency. +func TestFiltfiltSuppressesMoreThanOnePass(t *testing.T) { + b, a := filtfiltTestFilter(t) + const fs = 100.0 + const fStop = 40.0 + wStop := 2 * math.Pi * fStop / fs + const n = 800 + vals := make([]float64, n) + for i := range vals { + vals[i] = math.Sin(wStop * float64(i)) + } + x := mustFloats(t, vals) + two, err := Filtfilt(b, a, x) + if err != nil { + t.Fatalf("Filtfilt: %v", err) + } + one, err := FilterApply(b, a, x) + if err != nil { + t.Fatalf("FilterApply: %v", err) + } + rms := func(y *core.Array) float64 { + var s float64 + for i := 200; i < n-200; i++ { + s += y.FloatAt(i) * y.FloatAt(i) + } + // Whole periods sit in the window, so the RMS of a unit + // sinusoid of amplitude A is A/√2 and the factor below + // returns the amplitude. + return math.Sqrt2 * math.Sqrt(s/(n-400)) + } + gotTwo := rms(two) + gotOne := rms(one) + wantTwo := cmag(freqResponse(b, a, wStop)) + wantTwo *= wantTwo + if math.Abs(gotTwo-wantTwo) > 0.05*wantTwo+1e-9 { + t.Fatalf("filtfilt stopband amplitude %.3g, want |H|^2 = %.3g", gotTwo, wantTwo) + } + wantOne := cmag(freqResponse(b, a, wStop)) + if math.Abs(gotOne-wantOne) > 0.05*wantOne+1e-9 { + t.Fatalf("one-pass stopband amplitude %.3g, want |H| = %.3g", gotOne, wantOne) + } + if gotTwo >= gotOne { + t.Fatalf("filtfilt %.3g not quieter than one pass %.3g", gotTwo, gotOne) + } +} + +// TestFiltfiltConstantTracks pins the DC behaviour: away from the +// cropped transient regions a constant comes back at the constant, +// squared DC gain being exactly one for the Butterworth design. +func TestFiltfiltConstantTracks(t *testing.T) { + b, a := filtfiltTestFilter(t) + const n = 400 + vals := make([]float64, n) + for i := range vals { + vals[i] = 2.5 + } + got, err := Filtfilt(b, a, mustFloats(t, vals)) + if err != nil { + t.Fatalf("Filtfilt: %v", err) + } + for i := 100; i < n-100; i++ { + if math.Abs(got.FloatAt(i)-2.5) > 1e-9 { + t.Fatalf("sample %d: %g, want 2.5", i, got.FloatAt(i)) + } + } +} + +// TestFiltfiltDtype pins the dtype contract: float32 keeps its dtype, +// int widens to float. +func TestFiltfiltDtype(t *testing.T) { + b, a := filtfiltTestFilter(t) + f32, err := core.FromFloat32s([]float32{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13}, 13) + if err != nil { + t.Fatal(err) + } + out, err := Filtfilt(b, a, f32) + if err != nil { + t.Fatalf("Filtfilt float32: %v", err) + } + if out.Dtype() != core.Float32 { + t.Fatalf("float32 input produced dtype %s", out.Dtype()) + } + ints, err := core.FromInts([]int64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13}, 13) + if err != nil { + t.Fatal(err) + } + out, err = Filtfilt(b, a, ints) + if err != nil { + t.Fatalf("Filtfilt int: %v", err) + } + if out.Dtype() != core.Float { + t.Fatalf("int input produced dtype %s", out.Dtype()) + } +} + +// TestFiltfiltErrors pins the input gates: the shape and coefficient +// checks FilterApply applies, and the pad-length refusal for short +// signals. +func TestFiltfiltErrors(t *testing.T) { + b, a := filtfiltTestFilter(t) + nfilt := max(len(b), len(a)) + short := mustFloats(t, make([]float64, 3*(nfilt-1))) + if _, err := Filtfilt(b, a, short); err == nil { + t.Error("signal at exactly the pad length accepted") + } + rank2 := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2) + if _, err := Filtfilt(b, a, rank2); err == nil { + t.Error("rank-2 accepted") + } + cx := mustComplexes(t, []complex128{1, 2, 3, 4, 5}, 5) + if _, err := Filtfilt(b, a, cx); err == nil { + t.Error("complex accepted") + } + exact := mustFloats(t, make([]float64, 3*(nfilt-1)+1)) + if _, err := Filtfilt([]float64{}, a, exact); err == nil || !strings.Contains(err.Error(), "coefficient lists") { + t.Errorf("empty b: want the coefficient gate, got %v", err) + } + if _, err := Filtfilt([]float64{1}, []float64{0, 1}, exact); err == nil { + t.Error("a[0] = 0 accepted") + } + if _, err := Filtfilt(b, a, mustFloats(t, nil)); err == nil { + t.Error("empty signal accepted") + } +} diff --git a/signal/float32_test.go b/signal/float32_test.go new file mode 100644 index 0000000..074e462 --- /dev/null +++ b/signal/float32_test.go @@ -0,0 +1,26 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestFloat32FFT checks that the transform computes in float64 from +// float32 inputs. +func TestFloat32FFT(t *testing.T) { + a, err := core.FromFloat32s([]float32{1, 1, 1, 1}, 4) + if err != nil { + t.Fatal(err) + } + spec, err := FFT(a) + if err != nil { + t.Fatalf("FFT float32: %v", err) + } + if v := spec.ComplexAt(0); v != 4 { + t.Fatalf("FFT DC: %v", v) + } +} diff --git a/signal/helpers_test.go b/signal/helpers_test.go new file mode 100644 index 0000000..17ed6ae --- /dev/null +++ b/signal/helpers_test.go @@ -0,0 +1,50 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// mustFromFloats builds a float array, failing the test on a bad shape. +func mustFromFloats(t *testing.T, vals []float64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, shape...) + if err != nil { + t.Fatalf("FromFloats(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustComplexes builds a complex array, failing the test on a bad shape. +func mustComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array { + t.Helper() + a, err := core.FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustFromInts builds an int array, failing the test on a bad shape. +func mustFromInts(t *testing.T, vals []int64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromInts(vals, shape...) + if err != nil { + t.Fatalf("FromInts(%v, %v): %v", vals, shape, err) + } + return a +} + +// 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)} + } + return mustFromFloats(t, vals, shape...) +} diff --git a/signal/hilbert.go b/signal/hilbert.go new file mode 100644 index 0000000..ae4bda6 --- /dev/null +++ b/signal/hilbert.go @@ -0,0 +1,72 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// AnalyticSignal builds the analytic signal z = x + i·H(x) of a real +// series, where H is the Hilbert transform: the negative frequencies +// are removed, the positive ones doubled and the DC and Nyquist bins +// left alone, so z carries the instantaneous amplitude and phase of +// the series. The Fourier definition treats the series as one period, +// so it is exact for an integer number of tones and produces the +// familiar edge swings on an aperiodic series, the same behaviour any +// Fourier-domain filter shows there. +func AnalyticSignal(data *core.Array) (*core.Array, error) { + const name = "AnalyticSignal" + if data.NDim() != 1 { + return nil, base.Errf("%s: the series must be a vector, got shape %s", name, base.ShapeText(data.Shape())) + } + if data.Dtype() == core.Complex { + return nil, base.Errf("%s: complex series are not supported", name) + } + n := data.Len() + if n == 0 { + return nil, base.Errf("%s: the series must not be empty", name) + } + spec, err := FFT(data) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + bins := spec.RawComplexes()[:spec.Len()] + half := (n + 1) / 2 + for k := range bins { + switch { + case k == 0 || (n%2 == 0 && k == n/2): + // DC and, for even lengths, Nyquist: single-sided bins. + case k < half: + bins[k] *= 2 + default: + bins[k] = 0 + } + } + out, err := IFFT(spec) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + return out, nil +} + +// Envelope returns the instantaneous amplitude of a series, the +// modulus of its analytic signal. For a narrowband series this traces +// the curve a peak detector would find, without the smoothing lag. +func Envelope(data *core.Array) (*core.Array, error) { + const name = "Envelope" + z, err := AnalyticSignal(data) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + bins := z.RawComplexes()[:z.Len()] + out := core.New(core.Float, z.Len()) + vals := out.RawFloats() + for i, c := range bins { + vals[i] = math.Hypot(real(c), imag(c)) + } + return out, nil +} diff --git a/signal/hilbert_test.go b/signal/hilbert_test.go new file mode 100644 index 0000000..4051524 --- /dev/null +++ b/signal/hilbert_test.go @@ -0,0 +1,91 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// hilbertTone builds an integer number of periods of a cosine, where +// the periodic Fourier definition of the analytic signal is exact. +func hilbertTone(t *testing.T, cycles, n int) *core.Array { + t.Helper() + vals := make([]float64, n) + for i := range n { + vals[i] = math.Cos(2 * math.Pi * float64(cycles) * float64(i) / float64(n)) + } + a, err := core.FromFloats(vals, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestAnalyticSignalTone checks the defining property on a pure tone: +// the analytic signal of cos is exp(iωt), so the real part is the +// input, the imaginary part the quadrature sine and the modulus one. +func TestAnalyticSignalTone(t *testing.T) { + const cycles, n = 13, 256 + z, err := AnalyticSignal(hilbertTone(t, cycles, n)) + if err != nil { + t.Fatalf("AnalyticSignal: %v", err) + } + bins := z.RawComplexes()[:z.Len()] + for i, c := range bins { + want := 2 * math.Pi * float64(cycles) * float64(i) / float64(n) + if math.Abs(real(c)-math.Cos(want)) > 1e-9 { + t.Fatalf("bin %d real part %.12g, want the input cosine", i, real(c)) + } + if math.Abs(imag(c)-math.Sin(want)) > 1e-9 { + t.Fatalf("bin %d imaginary part %.12g, want the quadrature sine", i, imag(c)) + } + if e := math.Hypot(real(c), imag(c)); math.Abs(e-1) > 1e-9 { + t.Fatalf("bin %d modulus %.12g, want 1", i, e) + } + } +} + +// TestEnvelopeAmplitudeModulation checks the envelope on an AM +// carrier: the modulus of the analytic signal must recover the +// modulating wave, the quantity an AM receiver is after. +func TestEnvelopeAmplitudeModulation(t *testing.T) { + const n = 512 + vals := make([]float64, n) + for i := range n { + mod := 1 + 0.5*math.Cos(2*math.Pi*4*float64(i)/float64(n)) + carrier := math.Cos(2 * math.Pi * 40 * float64(i) / float64(n)) + vals[i] = mod * carrier + } + a, err := core.FromFloats(vals, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + env, err := Envelope(a) + if err != nil { + t.Fatalf("Envelope: %v", err) + } + vals = env.RawFloats()[:env.Len()] + for i, got := range vals { + want := 1 + 0.5*math.Cos(2*math.Pi*4*float64(i)/float64(n)) + if math.Abs(got-want) > 1e-9 { + t.Fatalf("sample %d envelope %.12g, want %.12g", i, got, want) + } + } +} + +// TestAnalyticSignalRefusals checks the shape and emptiness guards. +func TestAnalyticSignalRefusals(t *testing.T) { + if _, err := AnalyticSignal(core.New(core.Float, 3, 2)); err == nil { + t.Fatal("matrix accepted") + } + if _, err := AnalyticSignal(core.New(core.Float, 0)); err == nil { + t.Fatal("zero-length series accepted") + } + if _, err := Envelope(core.New(core.Float, 3, 3, 3)); err == nil { + t.Fatal("3-D array accepted") + } +} diff --git a/signal/kalman.go b/signal/kalman.go new file mode 100644 index 0000000..cf6ce85 --- /dev/null +++ b/signal/kalman.go @@ -0,0 +1,941 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// State-space filtering: the Kalman family over the model +// x_{t+1} = f(x_t) + w_t with w ~ N(0, Q), z_t = h(x_t) + v_t with +// v ~ N(0, R). The linear filter runs the standard Riccati recursion +// with a Joseph-form correction, the extended filter linearises f and +// h at the current estimate (analytic Jacobians when supplied, the +// house central-difference helper otherwise), and the unscented filter +// carries the state distribution through the deterministic sigma-point +// set. All three accumulate the exact Gaussian log-likelihood of the +// innovation sequence, and every covariance they publish is mirrored +// into its symmetric average, with positive definiteness enforced +// where the mathematics demands it: a Cholesky factorisation that +// cannot be taken names the step it failed at. + +// StateFunc maps one state vector to the image vector the model's +// transition or observation applies: the f of x_{t+1} = f(x_t) + w_t, +// the h of z_t = h(x_t) + v_t. The input array is private to the call +// and may be kept until the call returns; the output must be a real +// rank-1 array of one fixed length. +type StateFunc func(x *core.Array) (*core.Array, error) + +// JacobianFunc returns the Jacobian of a StateFunc at x as an +// (m × n) float array, row i holding the partials of output i. An +// analytic Jacobian is always the better instrument; when the option +// is left unset the extended filter builds one by central differences +// at every step. +type JacobianFunc func(x *core.Array) (*core.Array, error) + +// KalmanOptions carries the initial condition, the noise levels and +// the filter-specific knobs. An unset (nil) array field takes its +// documented default: a zero initial state, an identity initial +// covariance, zero process noise, an identity measurement noise. +// SigmaAlpha, SigmaBeta and SigmaKappa follow the same rule with the +// defaults 0.001, 2 and 0, the standard scaled unscented choice for a +// Gaussian prior. The extended and unscented filters cannot infer the +// state dimension from their callbacks, so at least one of +// InitialState and InitialCovariance must be set for them. +// +// TransitionJacobian and ObservationJacobian give the extended filter +// the analytic partials it linearises with; a nil one is replaced by +// central differences on the underlying map, two evaluations per +// partial per step. +type KalmanOptions struct { + InitialState *core.Array + InitialCovariance *core.Array + ProcessNoise *core.Array + MeasurementNoise *core.Array + TransitionJacobian JacobianFunc + ObservationJacobian JacobianFunc + SigmaAlpha float64 + SigmaBeta float64 + SigmaKappa float64 +} + +// KalmanResult holds one filtering pass over the measurement stack. +type KalmanResult struct { + // States is the (n × d) stack of filtered means x̂_{t|t}, one row + // per measurement. + States *core.Array + // Covariances is the (n × d × d) stack of filtered covariances + // P_{t|t}, one symmetric positive-definite block per measurement. + Covariances *core.Array + // Innovations is the (n × m) stack of one-step prediction errors + // z_t − h(x̂_{t|t−1}). + Innovations *core.Array + // InnovationCovariances is the (n × m × m) stack of the innovation + // covariances S_t the likelihood reads. + InnovationCovariances *core.Array + // LogLikelihood is Σ_t log N(z_t; h(x̂_{t|t−1}), S_t), the exact + // Gaussian likelihood of the measurement sequence under the model + // and the quantity noise and parameter estimation maximises. + LogLikelihood float64 +} + +// kalmanStepOut carries one step's outputs from a filter's step +// closure into the result stack. +type kalmanStepOut struct { + mean []float64 + covariance []float64 + innovation []float64 + innovationCov []float64 + logLikelihood float64 +} + +// kfSymmetryEps is the relative tolerance a covariance's mirror check +// allows, the same reading the multivariate normal takes: matrices +// assembled from products differ from their mirror by an ulp of +// rounding, a genuinely asymmetric pair by far more. +const kfSymmetryEps = 1e-12 + +// kfFinite refuses the non-finite values a filter would otherwise +// carry silently through every recursion. +func kfFinite(name, what string, vals []float64) error { + for i, v := range vals { + if math.IsNaN(v) || math.IsInf(v, 0) { + return base.Errf("%s: %s holds the non-finite value %g at %d", name, what, v, i) + } + } + return nil +} + +// kfVector reads a finite real rank-1 array into a float slice, the +// payload itself for a contiguous float64 array, +// with copies made only where the payload cannot serve, of exactly want entries when want is non-negative. +func kfVector(name, what string, a *core.Array, want int) ([]float64, error) { + if a.NDim() != 1 { + return nil, base.Errf("%s: %s must be a vector, got shape %s", name, what, base.ShapeText(a.Shape())) + } + if a.Dtype() == core.Complex { + return nil, base.Errf("%s: %s must be real, got complex", name, what) + } + if want >= 0 && a.Len() != want { + return nil, base.Errf("%s: %s holds %d entries, want %d", name, what, a.Len(), want) + } + vals := widenFloats(a) + if err := kfFinite(name, what, vals); err != nil { + return nil, err + } + return vals, nil +} + +// kfMatrix reads a finite real rank-2 array into a fresh row-major +// float slice of exactly rows×cols, either side wildcarded at −1. +func kfMatrix(name, what string, a *core.Array, rows, cols int) ([]float64, error) { + if a.NDim() != 2 { + return nil, base.Errf("%s: %s must be rank 2, got shape %s", name, what, base.ShapeText(a.Shape())) + } + if a.Dtype() == core.Complex { + return nil, base.Errf("%s: %s must be real, got complex", name, what) + } + shape := a.Shape() + if (rows >= 0 && shape[0] != rows) || (cols >= 0 && shape[1] != cols) { + return nil, base.Errf("%s: %s has shape %s, want %d×%d", name, what, base.ShapeText(shape), rows, cols) + } + vals := widenFloats(a) + if err := kfFinite(name, what, vals); err != nil { + return nil, err + } + return vals, nil +} + +// kfSymmetric demands every mirror pair agree within a relative +// tolerance, because only one triangle is ever read. +func kfSymmetric(name, what string, a []float64, n int) error { + for i := range n { + for j := range i { + lo, hi := a[i*n+j], a[j*n+i] + if math.Abs(lo-hi) > kfSymmetryEps*math.Max(math.Abs(lo), math.Abs(hi)) { + return base.Errf("%s: %s is not symmetric at (%d, %d): %g against %g", + name, what, i+1, j+1, lo, hi) + } + } + } + return nil +} + +// kfNoise reads one noise covariance: an unset array takes its +// documented default (zeros for the process noise, whose only job is +// to enter sums, an identity for the measurement noise, which is +// factored every step and so must be positive definite), a set one +// must be a finite symmetric matrix of the right size. +func kfNoise(name, what string, a *core.Array, dim int, identityDefault bool) ([]float64, error) { + if a == nil { + if identityDefault { + return kfIdentity(dim), nil + } + return make([]float64, dim*dim), nil + } + vals, err := kfMatrix(name, what, a, dim, dim) + if err != nil { + return nil, err + } + if err := kfSymmetric(name, what, vals, dim); err != nil { + return nil, err + } + if identityDefault { + if _, err := kfCholesky(what, vals, dim); err != nil { + return nil, base.Errf("%s: %w", name, err) + } + } + return vals, nil +} + +// kfInitialCondition reads the initial mean and covariance, zeros and +// identity where the options leave them unset. The initial covariance +// must be symmetric positive definite: the first predict would +// otherwise hand the recursion a structure it cannot factor. +func kfInitialCondition(name string, opts KalmanOptions, d int) (x0, p0 []float64, err error) { + if opts.InitialState == nil { + x0 = make([]float64, d) + } else if x0, err = kfVector(name, "the initial state", opts.InitialState, d); err != nil { + return nil, nil, err + } + if opts.InitialCovariance == nil { + p0 = kfIdentity(d) + } else { + if p0, err = kfMatrix(name, "the initial covariance", opts.InitialCovariance, d, d); err != nil { + return nil, nil, err + } + if err := kfSymmetric(name, "the initial covariance", p0, d); err != nil { + return nil, nil, err + } + if _, err := kfCholesky("the initial covariance", p0, d); err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + } + return x0, p0, nil +} + +// kfNonlinearInputs resolves the shared inputs of the extended and +// unscented filters: the state dimension, which the callbacks cannot +// carry, comes from the initial state or the initial covariance, at +// least one of which must be set. +func kfNonlinearInputs(name string, opts KalmanOptions, m int) (d int, x0, p0, q, r []float64, err error) { + switch { + case opts.InitialState != nil: + if x0, err = kfVector(name, "the initial state", opts.InitialState, -1); err != nil { + return 0, nil, nil, nil, nil, err + } + d = len(x0) + case opts.InitialCovariance != nil: + if opts.InitialCovariance.NDim() != 2 { + return 0, nil, nil, nil, nil, base.Errf("%s: the initial covariance must be rank 2, got shape %s", + name, base.ShapeText(opts.InitialCovariance.Shape())) + } + d = opts.InitialCovariance.Shape()[0] + default: + return 0, nil, nil, nil, nil, base.Errf("%s: the state dimension must come from an initial state or an initial covariance; neither is set", name) + } + if d < 1 { + return 0, nil, nil, nil, nil, base.Errf("%s: the state dimension must be at least 1, got %d", name, d) + } + if x0, p0, err = kfInitialCondition(name, opts, d); err != nil { + return 0, nil, nil, nil, nil, err + } + if q, err = kfNoise(name, "the process noise", opts.ProcessNoise, d, false); err != nil { + return 0, nil, nil, nil, nil, err + } + if r, err = kfNoise(name, "the measurement noise", opts.MeasurementNoise, m, true); err != nil { + return 0, nil, nil, nil, nil, err + } + return d, x0, p0, q, r, nil +} + +// kfMeasurements reads the measurement stack: a rank-1 array holds n +// scalar observations, a rank-2 array holds n rows of m channels. At +// least one measurement is needed, every value must be finite. +func kfMeasurements(name string, z *core.Array) (n, m int, rows []float64, err error) { + if z.NDim() != 1 && z.NDim() != 2 { + return 0, 0, nil, base.Errf("%s: the measurements must be rank 1 or rank 2, got shape %s", + name, base.ShapeText(z.Shape())) + } + if z.Dtype() == core.Complex { + return 0, 0, nil, base.Errf("%s: complex measurements are not supported", name) + } + if z.NDim() == 1 { + m = 1 + } else { + m = z.Shape()[1] + } + if m < 1 { + return 0, 0, nil, base.Errf("%s: the measurement width must be at least 1, got %d", name, m) + } + n = z.Len() / m + if n < 1 { + return 0, 0, nil, base.Errf("%s: at least one measurement is needed", name) + } + rows = widenFloats(z) + if err := kfFinite(name, "the measurements", rows); err != nil { + return 0, 0, nil, err + } + return n, m, rows, nil +} + +// kfIdentity returns the n×n identity, row-major. +func kfIdentity(n int) []float64 { + a := make([]float64, n*n) + for i := range n { + a[i*n+i] = 1 + } + return a +} + +// kfTranspose returns the transpose of the row-major rows×cols +// matrix. +func kfTranspose(a []float64, rows, cols int) []float64 { + out := make([]float64, rows*cols) + for i := range rows { + for j := range cols { + out[j*rows+i] = a[i*cols+j] + } + } + return out +} + +// kfMatVec returns A·x for the row-major rows×cols matrix A. +func kfMatVec(a []float64, rows, cols int, x []float64) []float64 { + out := make([]float64, rows) + for i := range rows { + total := 0.0 + row := a[i*cols : (i+1)*cols] + for j, v := range row { + total += v * x[j] + } + out[i] = total + } + return out +} + +// kfMatMul returns A·B for the row-major A of size ra×ca and B of +// size ca×cb. +func kfMatMul(a []float64, ra, ca int, b []float64, cb int) []float64 { + out := make([]float64, ra*cb) + for i := range ra { + row := out[i*cb : (i+1)*cb] + arow := a[i*ca : (i+1)*ca] + for k, aik := range arow { + brow := b[k*cb : (k+1)*cb] + for j := range cb { + row[j] += aik * brow[j] + } + } + } + return out +} + +// kfCholesky factors the symmetric positive-definite row-major n×n +// matrix into the lower triangular L with A = L·Lᵀ, only the lower +// mirror read. A non-positive pivot names its row. +func kfCholesky(what string, a []float64, n int) ([]float64, error) { + l := make([]float64, n*n) + for i := range n { + for j := range i + 1 { + total := a[i*n+j] + for k := range j { + total -= l[i*n+k] * l[j*n+k] + } + if i == j { + if !(total > 0) { + return nil, base.Errf("%s is not positive definite at row %d", what, i+1) + } + l[i*n+i] = math.Sqrt(total) + } else { + l[i*n+j] = total / l[j*n+j] + } + } + } + return l, nil +} + +// kfCholSolve solves L·Lᵀ·x = b through the forward and backward +// substitutions. +func kfCholSolve(l []float64, n int, b []float64) []float64 { + x := make([]float64, n) + for i := range n { + total := b[i] + for k := range i { + total -= l[i*n+k] * x[k] + } + x[i] = total / l[i*n+i] + } + for i := n - 1; i >= 0; i-- { + total := x[i] + for k := i + 1; k < n; k++ { + total -= l[k*n+i] * x[k] + } + x[i] = total / l[i*n+i] + } + return x +} + +// kfCholSolveMatrix solves L·Lᵀ·X = B for the row-major B of size +// n×cols, row by row. +func kfCholSolveMatrix(l []float64, n int, b []float64, cols int) []float64 { + x := make([]float64, len(b)) + copy(x, b) + for i := range n { + row := x[i*cols : (i+1)*cols] + for k := range i { + lk := l[i*n+k] + krow := x[k*cols : (k+1)*cols] + for j := range cols { + row[j] -= lk * krow[j] + } + } + di := l[i*n+i] + for j := range cols { + row[j] /= di + } + } + for i := n - 1; i >= 0; i-- { + row := x[i*cols : (i+1)*cols] + for k := i + 1; k < n; k++ { + lk := l[k*n+i] + krow := x[k*cols : (k+1)*cols] + for j := range cols { + row[j] -= lk * krow[j] + } + } + di := l[i*n+i] + for j := range cols { + row[j] /= di + } + } + return x +} + +// kfLogDet returns the log determinant of the Cholesky factor: twice +// the sum of the log diagonal. +func kfLogDet(l []float64, n int) float64 { + total := 0.0 + for i := range n { + total += math.Log(l[i*n+i]) + } + return 2 * total +} + +// kfSymmetrise replaces A by (A + Aᵀ)/2 in place: the rounding drift +// of a recursion that touches a covariance only through symmetric +// expressions cannot survive the mirror average. +func kfSymmetrise(a []float64, n int) { + for i := range n { + for j := range i { + avg := (a[i*n+j] + a[j*n+i]) / 2 + a[i*n+j] = avg + a[j*n+i] = avg + } + } +} + +// kalmanRun walks the measurement stack through the step closure, +// which owns one predict-and-correct cycle: it receives the current +// mean and covariance (read-only; it must return fresh slices) and +// the step's measurement row, and returns everything the result stack +// records, the Gaussian log-likelihood contribution included. +func kalmanRun(nMeas, m, d int, meas []float64, x0, p0 []float64, + step func(t int, zRow, mean, cov []float64) (kalmanStepOut, error), +) (*KalmanResult, error) { + out := &KalmanResult{ + States: core.New(core.Float, nMeas, d), + Covariances: core.New(core.Float, nMeas, d, d), + Innovations: core.New(core.Float, nMeas, m), + InnovationCovariances: core.New(core.Float, nMeas, m, m), + } + mean := slices.Clone(x0) + cov := slices.Clone(p0) + total := 0.0 + for t := range nMeas { + res, err := step(t, meas[t*m:(t+1)*m], mean, cov) + if err != nil { + return nil, err + } + copy(out.States.RawFloats()[t*d:(t+1)*d], res.mean) + copy(out.Covariances.RawFloats()[t*d*d:(t+1)*d*d], res.covariance) + copy(out.Innovations.RawFloats()[t*m:(t+1)*m], res.innovation) + copy(out.InnovationCovariances.RawFloats()[t*m*m:(t+1)*m*m], res.innovationCov) + total += res.logLikelihood + mean, cov = res.mean, res.covariance + } + out.LogLikelihood = total + return out, nil +} + +// kfCorrect runs the shared correction of the linear and extended +// filters: the innovation against the predicted observation, its +// covariance S = H·P⁻·Hᵀ + R, the gain K = P⁻·Hᵀ·S⁻¹ through the +// Cholesky solve, the state update and the Joseph-form covariance +// (I−K·H)·P⁻·(I−K·H)ᵀ + K·R·Kᵀ, mirrored into its symmetric average. +// The Joseph form keeps the filtered covariance symmetric positive +// definite by construction, where the plain (I−K·H)·P⁻ recursion +// preserves it only in exact arithmetic. +func kfCorrect(name string, t int, zRow, predicted, xp, pp []float64, d, m int, + h, ht, r []float64, +) (kalmanStepOut, error) { + innovation := make([]float64, m) + for i := range m { + innovation[i] = zRow[i] - predicted[i] + } + hp := kfMatMul(h, m, d, pp, d) + s := kfMatMul(hp, m, d, ht, m) + for i := range s { + s[i] += r[i] + } + ls, err := kfCholesky("the innovation covariance", s, m) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t, err) + } + solved := kfCholSolve(ls, m, innovation) + quad := 0.0 + for i := range m { + quad += innovation[i] * solved[i] + } + // K = P⁻·Hᵀ·S⁻¹ arrives through the transposed solve: S·X = + // H·P⁻ gives X = S⁻¹·H·P⁻ = Kᵀ and K = Xᵀ. + gain := kfTranspose(kfCholSolveMatrix(ls, m, hp, d), m, d) + mean := make([]float64, d) + copy(mean, xp) + for i := range d { + total := 0.0 + grow := gain[i*m : (i+1)*m] + for j, y := range innovation { + total += grow[j] * y + } + mean[i] += total + } + kh := kfMatMul(gain, d, m, h, d) + a := make([]float64, d*d) + for i := range d { + for j := range d { + a[i*d+j] = -kh[i*d+j] + } + a[i*d+i]++ + } + at := kfTranspose(a, d, d) + cov := kfMatMul(kfMatMul(a, d, d, pp, d), d, d, at, d) + krkt := kfMatMul(kfMatMul(gain, d, m, r, m), d, m, kfTranspose(gain, d, m), d) + for i := range cov { + cov[i] += krkt[i] + } + kfSymmetrise(cov, d) + logLik := -0.5 * (float64(m)*math.Log(2*math.Pi) + kfLogDet(ls, m) + quad) + return kalmanStepOut{mean: mean, covariance: cov, innovation: innovation, + innovationCov: s, logLikelihood: logLik}, nil +} + +// KalmanFilter runs the linear Kalman filter over the measurement +// stack z under the model x_{t+1} = F·x_t + w_t, z_t = H·x_t + v_t: +// transition is the (d × d) state matrix F, observation the (m × d) +// matrix H. Every step predicts with F and corrects with a +// Joseph-form update, and the exact Gaussian log-likelihood of the +// innovation sequence accumulates into the result. +// +// A rank-1 z holds n scalar observations, a rank-2 z holds n rows of +// m channels. Nil option fields take their defaults (see +// KalmanOptions). A non-square transition, a shape mismatch, an +// asymmetric noise covariance, a singular measurement noise, a +// non-finite value anywhere, and an innovation covariance that loses +// positive definiteness (naming the step) are errors. +func KalmanFilter(z *core.Array, transition, observation *core.Array, opts KalmanOptions) (*KalmanResult, error) { + const name = "KalmanFilter" + if transition == nil || observation == nil { + return nil, base.Errf("%s: the transition and observation matrices are required", name) + } + if transition.NDim() != 2 { + return nil, base.Errf("%s: the transition matrix must be rank 2, got shape %s", + name, base.ShapeText(transition.Shape())) + } + shape := transition.Shape() + if shape[0] != shape[1] { + return nil, base.Errf("%s: the transition matrix must be square, got %d×%d", name, shape[0], shape[1]) + } + d := shape[0] + if d < 1 { + return nil, base.Errf("%s: the state dimension must be at least 1, got %d", name, d) + } + nMeas, m, meas, err := kfMeasurements(name, z) + if err != nil { + return nil, err + } + f, err := kfMatrix(name, "the transition matrix", transition, d, d) + if err != nil { + return nil, err + } + h, err := kfMatrix(name, "the observation matrix", observation, m, d) + if err != nil { + return nil, err + } + q, err := kfNoise(name, "the process noise", opts.ProcessNoise, d, false) + if err != nil { + return nil, err + } + r, err := kfNoise(name, "the measurement noise", opts.MeasurementNoise, m, true) + if err != nil { + return nil, err + } + x0, p0, err := kfInitialCondition(name, opts, d) + if err != nil { + return nil, err + } + ft := kfTranspose(f, d, d) + ht := kfTranspose(h, m, d) + return kalmanRun(nMeas, m, d, meas, x0, p0, func(t int, zRow, x, p []float64) (kalmanStepOut, error) { + // Predict: x⁻ = F·x, P⁻ = F·P·Fᵀ + Q. + xp := kfMatVec(f, d, d, x) + pp := kfMatMul(kfMatMul(f, d, d, p, d), d, d, ft, d) + for i := range pp { + pp[i] += q[i] + } + predicted := kfMatVec(h, m, d, xp) + return kfCorrect(name, t+1, zRow, predicted, xp, pp, d, m, h, ht, r) + }) +} + +// ExtendedKalmanFilter runs the extended Kalman filter over the +// measurement stack z under the nonlinear model x_{t+1} = f(x_t) + +// w_t, z_t = h(x_t) + v_t: every step predicts by propagating the +// mean through f and the covariance through the linearisation F = +// ∂f/∂x at x̂_{t|t}, then corrects through the Jacobian H = ∂h/∂x at +// the predicted mean, with the same Joseph-form update and the same +// exact log-likelihood the linear filter carries. The Jacobians come +// from the options when supplied analytically and from central +// differences otherwise (see JacobianFunc). +// +// The extended filter is the linear one applied to local linear +// models: it inherits the Kalman recursions and with them the first +// order's blindness to the curvature of f and h, so a strongly bent +// observation map wants the unscented filter instead. The state +// dimension must be fixed by an initial state or covariance (see +// KalmanOptions); everything else follows the linear filter's +// contract, with a failing callback or Jacobian reported with the +// step it failed at. +func ExtendedKalmanFilter(z *core.Array, transition, observation StateFunc, opts KalmanOptions) (*KalmanResult, error) { + const name = "ExtendedKalmanFilter" + if transition == nil || observation == nil { + return nil, base.Errf("%s: the transition and observation maps are required", name) + } + nMeas, m, meas, err := kfMeasurements(name, z) + if err != nil { + return nil, err + } + d, x0, p0, q, r, err := kfNonlinearInputs(name, opts, m) + if err != nil { + return nil, err + } + fjac := opts.TransitionJacobian + if fjac == nil { + fjac = func(x *core.Array) (*core.Array, error) { + return core.Jacobian(transition, x, core.JacobianOptions{}) + } + } + hjac := opts.ObservationJacobian + if hjac == nil { + hjac = func(x *core.Array) (*core.Array, error) { + return core.Jacobian(observation, x, core.JacobianOptions{}) + } + } + return kalmanRun(nMeas, m, d, meas, x0, p0, func(t int, zRow, x, p []float64) (kalmanStepOut, error) { + xArr, err := core.FromFloats(x, d) + if err != nil { + return kalmanStepOut{}, err + } + fOut, err := transition(xArr) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: the transition: %w", name, t+1, err) + } + xp, err := kfVector(name, "the transition output", fOut, d) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + fJacArr, err := fjac(xArr) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: the transition Jacobian: %w", name, t+1, err) + } + fMat, err := kfMatrix(name, "the transition Jacobian", fJacArr, d, d) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + // Predict through the linearisation at the filtered mean. + pp := kfMatMul(kfMatMul(fMat, d, d, p, d), d, d, kfTranspose(fMat, d, d), d) + for i := range pp { + pp[i] += q[i] + } + // Correct through the linearisation at the predicted mean, but + // innovate against the true observation map. + xpArr, err := core.FromFloats(xp, d) + if err != nil { + return kalmanStepOut{}, err + } + hOut, err := observation(xpArr) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: the observation: %w", name, t+1, err) + } + predicted, err := kfVector(name, "the observation output", hOut, m) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + hJacArr, err := hjac(xpArr) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: the observation Jacobian: %w", name, t+1, err) + } + hMat, err := kfMatrix(name, "the observation Jacobian", hJacArr, m, d) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + return kfCorrect(name, t+1, zRow, predicted, xp, pp, d, m, hMat, kfTranspose(hMat, m, d), r) + }) +} + +// UnscentedKalmanFilter runs the unscented Kalman filter over the +// measurement stack z under the nonlinear model x_{t+1} = f(x_t) + +// w_t, z_t = h(x_t) + v_t: the state distribution N(x̂, P) is carried +// through f and h exactly to second order by the deterministic +// sigma-point set x̂ ± sqrt(d+λ)·L[:, i], L the Cholesky factor of P, +// whose weighted moments reconstruct the predicted mean and +// covariance. The update re-draws the set from the predicted +// distribution, so the state-observation cross-covariance carries the +// process noise the innovation covariance does. λ = α²(d+κ) − d, with +// d the state dimension the sigma-point set spans, comes from +// KalmanOptions' SigmaAlpha, SigmaBeta and SigmaKappa; the +// weights are the standard scaled set, with the covariance weight of +// the centre point carrying the 1 − α² + β prior correction. +// +// The correction carries no Joseph form: the gain comes from the +// explicit cross-covariance of state and observation rather than an +// observation matrix, so the covariance leaves the update as +// P⁻ − K·S·Kᵀ mirrored into its symmetric average, and positive +// definiteness is enforced by the next predict's Cholesky +// factorisation, which names the step when it fails. The log-likelihood +// accumulates exactly as in the linear filter. On a linear model the +// sigma transforms are exact and the filter degenerates to the +// Kalman answer to rounding; the state dimension must be fixed by an +// initial state or covariance (see KalmanOptions). +func UnscentedKalmanFilter(z *core.Array, transition, observation StateFunc, opts KalmanOptions) (*KalmanResult, error) { + const name = "UnscentedKalmanFilter" + if transition == nil || observation == nil { + return nil, base.Errf("%s: the transition and observation maps are required", name) + } + nMeas, m, meas, err := kfMeasurements(name, z) + if err != nil { + return nil, err + } + d, x0, p0, q, r, err := kfNonlinearInputs(name, opts, m) + if err != nil { + return nil, err + } + alpha := opts.SigmaAlpha + if alpha == 0 { + alpha = 0.001 + } + beta := opts.SigmaBeta + if beta == 0 { + beta = 2 + } + kappa := opts.SigmaKappa + if math.IsNaN(alpha) || alpha < 0 { + return nil, base.Errf("%s: the sigma alpha must be unset (the default 0.001) or positive, got %g", + name, opts.SigmaAlpha) + } + if math.IsNaN(beta) || beta < 0 { + return nil, base.Errf("%s: the sigma beta must be unset (the default 2) or positive, got %g", + name, opts.SigmaBeta) + } + if math.IsNaN(kappa) { + return nil, base.Errf("%s: the sigma kappa must not be NaN, got %g", name, kappa) + } + scale := alpha * alpha * (float64(d) + kappa) + if scale <= 0 { + return nil, base.Errf("%s: the sigma spread vanishes: alpha %g and kappa %g leave no positive scale for %d states", + name, alpha, kappa, d) + } + points := 2*d + 1 + lambda := scale - float64(d) + wm := make([]float64, points) + wc := make([]float64, points) + wm[0] = lambda / scale + wc[0] = wm[0] + 1 - alpha*alpha + beta + for i := 1; i < points; i++ { + wm[i] = 1 / (2 * scale) + wc[i] = wm[i] + } + spread := math.Sqrt(scale) + sig := make([]float64, points*d) + prop := make([]float64, points*d) + obs := make([]float64, points*m) + return kalmanRun(nMeas, m, d, meas, x0, p0, func(t int, zRow, x, p []float64) (kalmanStepOut, error) { + // The sigma set: the mean plus each Cholesky column of the + // covariance, pushed out by sqrt(d+λ) both ways. + l, err := kfCholesky("the state covariance", p, d) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + copy(sig[:d], x) + for i := range d { + plus := sig[(1+2*i)*d : (2+2*i)*d] + minus := sig[(2+2*i)*d : (3+2*i)*d] + copy(plus, x) + copy(minus, x) + for j := range d { + v := spread * l[j*d+i] + plus[j] += v + minus[j] -= v + } + } + // Predict: every sigma through f, then the weighted moments. + for s := range points { + sa, err := core.FromFloats(sig[s*d:(s+1)*d], d) + if err != nil { + return kalmanStepOut{}, err + } + fa, err := transition(sa) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: the transition: %w", name, t+1, err) + } + fv, err := kfVector(name, "the transition output", fa, d) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + copy(prop[s*d:(s+1)*d], fv) + } + xp := make([]float64, d) + for s := range points { + for j := range d { + xp[j] += wm[s] * prop[s*d+j] + } + } + pp := make([]float64, d*d) + for s := range points { + w := wc[s] + prows := prop[s*d : (s+1)*d] + for a := range d { + da := prows[a] - xp[a] + for b := range d { + pp[a*d+b] += w * da * (prows[b] - xp[b]) + } + } + } + for i := range pp { + pp[i] += q[i] + } + kfSymmetrise(pp, d) + // The update re-draws the sigma set from the predicted + // distribution N(x⁻, P⁻), the process noise included: the + // state-observation cross-covariance must carry the same + // spread the innovation covariance does, or the gain loses the + // Q·Hᵀ term and a linear model stops degenerating to the + // Kalman answer. + lp2, err := kfCholesky("the predicted state covariance", pp, d) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + copy(sig[:d], xp) + for i := range d { + plus := sig[(1+2*i)*d : (2+2*i)*d] + minus := sig[(2+2*i)*d : (3+2*i)*d] + copy(plus, xp) + copy(minus, xp) + for j := range d { + v := spread * lp2[j*d+i] + plus[j] += v + minus[j] -= v + } + } + // Observation: the predicted sigmas through h, the weighted + // moments, and the state-observation cross-covariance. + for s := range points { + pa, err := core.FromFloats(sig[s*d:(s+1)*d], d) + if err != nil { + return kalmanStepOut{}, err + } + ha, err := observation(pa) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: the observation: %w", name, t+1, err) + } + hv, err := kfVector(name, "the observation output", ha, m) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + copy(obs[s*m:(s+1)*m], hv) + } + zp := make([]float64, m) + for s := range points { + for j := range m { + zp[j] += wm[s] * obs[s*m+j] + } + } + s := make([]float64, m*m) + for k := range points { + w := wc[k] + orow := obs[k*m : (k+1)*m] + for a := range m { + da := orow[a] - zp[a] + for b := range m { + s[a*m+b] += w * da * (orow[b] - zp[b]) + } + } + } + for i := range s { + s[i] += r[i] + } + kfSymmetrise(s, m) + cross := make([]float64, d*m) + for k := range points { + w := wc[k] + prow := sig[k*d : (k+1)*d] + orow := obs[k*m : (k+1)*m] + for a := range d { + da := prow[a] - xp[a] + for b := range m { + cross[a*m+b] += w * da * (orow[b] - zp[b]) + } + } + } + ls, err := kfCholesky("the innovation covariance", s, m) + if err != nil { + return kalmanStepOut{}, base.Errf("%s: at step %d: %w", name, t+1, err) + } + innovation := make([]float64, m) + for i := range m { + innovation[i] = zRow[i] - zp[i] + } + solved := kfCholSolve(ls, m, innovation) + quad := 0.0 + for i := range m { + quad += innovation[i] * solved[i] + } + // K = P_xz·S⁻¹ through the transposed solve: S·X = P_xzᵀ + // gives X = S⁻¹·P_xzᵀ and K = Xᵀ. + gain := kfTranspose(kfCholSolveMatrix(ls, m, kfTranspose(cross, d, m), d), m, d) + mean := make([]float64, d) + copy(mean, xp) + for i := range d { + total := 0.0 + grow := gain[i*m : (i+1)*m] + for j, y := range innovation { + total += grow[j] * y + } + mean[i] += total + } + kskt := kfMatMul(kfMatMul(gain, d, m, s, m), d, m, kfTranspose(gain, d, m), d) + covariance := make([]float64, d*d) + for i := range covariance { + covariance[i] = pp[i] - kskt[i] + } + kfSymmetrise(covariance, d) + logLik := -0.5 * (float64(m)*math.Log(2*math.Pi) + kfLogDet(ls, m) + quad) + return kalmanStepOut{mean: mean, covariance: covariance, innovation: innovation, + innovationCov: s, logLikelihood: logLik}, nil + }) +} diff --git a/signal/kalman_test.go b/signal/kalman_test.go new file mode 100644 index 0000000..f3dd591 --- /dev/null +++ b/signal/kalman_test.go @@ -0,0 +1,739 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "errors" + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// scalarKFRun is the scalar Kalman recursion written out by hand, the +// closed-form reference the library's matrix machinery is pinned +// against: the same predict-and-correct cycle in plain arithmetic. +func scalarKFRun(z []float64, f, q, h, r, x0, p0 float64) (xs, ps, inns, ss []float64, logLik float64) { + x, p := x0, p0 + xs = make([]float64, len(z)) + ps = make([]float64, len(z)) + inns = make([]float64, len(z)) + ss = make([]float64, len(z)) + for t, zt := range z { + xp := f * x + pp := f*p*f + q + s := h*pp*h + r + k := pp * h / s + inn := zt - h*xp + x = xp + k*inn + p = pp - k*s*k // the standard form: (1−k·h)·pp + xs[t], ps[t], inns[t], ss[t] = x, p, inn, s + logLik += -0.5 * (math.Log(2*math.Pi*s) + inn*inn/s) + } + return xs, ps, inns, ss, logLik +} + +// cholPD reports whether a symmetric row-major block factors as +// positive definite. +func cholPD(a []float64, n int) bool { + l := make([]float64, n*n) + for i := range n { + for j := range i + 1 { + total := a[i*n+j] + for k := range j { + total -= l[i*n+k] * l[j*n+k] + } + if i == j { + if !(total > 0) { + return false + } + l[i*n+i] = math.Sqrt(total) + } else { + l[i*n+j] = total / l[j*n+j] + } + } + } + return true +} + +// kalmanResultBlock reads one (rows × cols) block out of a result +// stack at step t. +func kalmanResultBlock(t *testing.T, a *core.Array, step, rows, cols int) []float64 { + t.Helper() + block := make([]float64, rows*cols) + for i := range rows * cols { + block[i] = a.FloatAt(step*rows*cols + i) + } + return block +} + +// TestKalmanScalarSteadyState pins the Riccati recursion against the +// scalar steady state solved by hand: for x_{t+1} = x_t + w, +// z_t = x_t + v with unit variances the fixed point of p⁻ = p + 1, +// p = p⁻·r/(p⁻ + r) is the golden ratio, and the filtered covariance +// settles on its reciprocal. +func TestKalmanScalarSteadyState(t *testing.T) { + const n = 400 + g := core.NewGenerator(3) + z := make([]float64, n) + for i := range z { + z[i] = g.NormalUnit() + } + f, err := KalmanFilter(mustFloats(t, z), mustFloats(t, []float64{1}, 1, 1), mustFloats(t, []float64{1}, 1, 1), + KalmanOptions{ProcessNoise: mustFloats(t, []float64{1}, 1, 1), MeasurementNoise: mustFloats(t, []float64{1}, 1, 1)}) + if err != nil { + t.Fatalf("KalmanFilter: %v", err) + } + wantP := (math.Sqrt(5) - 1) / 2 // 1/φ + if got := f.Covariances.FloatAt(n - 1); math.Abs(got-wantP) > 1e-9 { + t.Fatalf("steady filtered covariance = %.12f, want %.12f", got, wantP) + } + // The state path must match the hand-run scalar recursion. + xs, _, _, _, _ := scalarKFRun(z, 1, 1, 1, 1, 0, 1) + if got, want := f.States.FloatAt(n-1), xs[n-1]; math.Abs(got-want) > 1e-9 { + t.Fatalf("steady filtered state = %.12f, hand recursion %.12f", got, want) + } +} + +// TestKalmanAR1Reference pins the filtered mean of a scalar AR(1) plus +// noise against the closed-form recursion, and the accumulated +// log-likelihood against the direct sum of the Gaussian log densities +// of the innovations and innovation variances the result itself +// publishes. +func TestKalmanAR1Reference(t *testing.T) { + const ( + phi = 0.7 + q = 0.04 + r = 0.25 + n = 80 + ) + g := core.NewGenerator(5) + z := make([]float64, n) + for i := range z { + z[i] = g.NormalUnit() + } + f, err := KalmanFilter(mustFloats(t, z), mustFloats(t, []float64{phi}, 1, 1), + mustFloats(t, []float64{1}, 1, 1), + KalmanOptions{ + ProcessNoise: mustFloats(t, []float64{q}, 1, 1), + MeasurementNoise: mustFloats(t, []float64{r}, 1, 1), + }) + if err != nil { + t.Fatalf("KalmanFilter: %v", err) + } + xs, ps, inns, ss, wantLL := scalarKFRun(z, phi, q, 1, r, 0, 1) + for step := range n { + if got := f.States.FloatAt(step); math.Abs(got-xs[step]) > 1e-9 { + t.Fatalf("filtered mean at %d = %.12f, want %.12f", step, got, xs[step]) + } + if got := f.Covariances.FloatAt(step); math.Abs(got-ps[step]) > 1e-9 { + t.Fatalf("filtered variance at %d = %.12f, want %.12f", step, got, ps[step]) + } + if got := f.Innovations.FloatAt(step); math.Abs(got-inns[step]) > 1e-9 { + t.Fatalf("innovation at %d = %.12f, want %.12f", step, got, inns[step]) + } + if got := f.InnovationCovariances.FloatAt(step); math.Abs(got-ss[step]) > 1e-9 { + t.Fatalf("innovation variance at %d = %.12f, want %.12f", step, got, ss[step]) + } + } + if math.Abs(f.LogLikelihood-wantLL) > 1e-7 { + t.Fatalf("log-likelihood = %.10f, closed form %.10f", f.LogLikelihood, wantLL) + } + // The published likelihood must be the sum of the published + // innovations' own Gaussian log densities. + total := 0.0 + for t := range n { + inn := f.Innovations.FloatAt(t) + s := f.InnovationCovariances.FloatAt(t) + total += -0.5 * (math.Log(2*math.Pi*s) + inn*inn/s) + } + if math.Abs(total-f.LogLikelihood) > 1e-9 { + t.Fatalf("log-likelihood %.12f is not the innovations' own sum %.12f", f.LogLikelihood, total) + } +} + +// TestKalmanMatchesUKFLinear runs the unscented filter on an exactly +// linear model: the sigma-point transforms are exact there, so the +// UKF must degenerate to the Kalman answer to rounding. +func TestKalmanMatchesUKFLinear(t *testing.T) { + const n = 120 + fMat := mustFloats(t, []float64{1, 1, 0, 1}, 2, 2) + hMat := mustFloats(t, []float64{1, 0}, 1, 2) + opts := KalmanOptions{ + InitialState: mustFloats(t, []float64{1, 0}), + InitialCovariance: mustFloats(t, []float64{4, 0, 0, 1}, 2, 2), + ProcessNoise: mustFloats(t, []float64{1e-4, 0, 0, 0.01}, 2, 2), + MeasurementNoise: mustFloats(t, []float64{0.25}, 1, 1), + } + g := core.NewGenerator(9) + z := make([]float64, n) + for i := range z { + z[i] = g.NormalUnit() + } + zArr := mustFloats(t, z) + kf, err := KalmanFilter(zArr, fMat, hMat, opts) + if err != nil { + t.Fatalf("KalmanFilter: %v", err) + } + linear := func(m *core.Array) StateFunc { + return func(x *core.Array) (*core.Array, error) { + rows := m.Len() / x.Len() + out := make([]float64, rows) + for i := range out { + total := 0.0 + for j := range x.Len() { + total += m.FloatAt(i*x.Len()+j) * x.FloatAt(j) + } + out[i] = total + } + return core.FromFloats(out, len(out)) + } + } + ukf, err := UnscentedKalmanFilter(zArr, linear(fMat), linear(hMat), opts) + if err != nil { + t.Fatalf("UnscentedKalmanFilter: %v", err) + } + for step := range n { + for i := range 2 { + got, want := ukf.States.FloatAt(step*2+i), kf.States.FloatAt(step*2+i) + if math.Abs(got-want) > 1e-8 { + t.Fatalf("UKF state (%d, %d) = %.10f against KF %.10f", step, i, got, want) + } + } + for i := range 4 { + got, want := ukf.Covariances.FloatAt(step*4+i), kf.Covariances.FloatAt(step*4+i) + if math.Abs(got-want) > 1e-8 { + t.Fatalf("UKF covariance (%d, %d) = %.10f against KF %.10f", step, i, got, want) + } + } + } + if math.Abs(ukf.LogLikelihood-kf.LogLikelihood) > 1e-6 { + t.Fatalf("UKF log-likelihood %.10f against KF %.10f", ukf.LogLikelihood, kf.LogLikelihood) + } + // The extended filter with central-difference Jacobians sees the + // same linear maps exactly, so it must land on the KF too. + ekf, err := ExtendedKalmanFilter(zArr, linear(fMat), linear(hMat), opts) + if err != nil { + t.Fatalf("ExtendedKalmanFilter: %v", err) + } + for step := range n { + for i := range 2 { + got, want := ekf.States.FloatAt(step*2+i), kf.States.FloatAt(step*2+i) + if math.Abs(got-want) > 1e-8 { + t.Fatalf("EKF state (%d, %d) = %.10f against KF %.10f", step, i, got, want) + } + } + } +} + +// TestUnscentedQuadraticMoment pins the sigma-point transform on a +// genuinely nonlinear map: with x ~ N(0, 2) and f(x) = x², the exact +// predicted second moment is 2, and with α = 1, κ = 0 the symmetric +// sigma pair reproduces it exactly. The inert measurement noise keeps +// the correction from moving the answer. +func TestUnscentedQuadraticMoment(t *testing.T) { + square := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{x.FloatAt(0) * x.FloatAt(0)}, 1) + } + identity := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{x.FloatAt(0)}, 1) + } + f, err := UnscentedKalmanFilter(mustFloats(t, []float64{0}), square, identity, KalmanOptions{ + InitialState: mustFloats(t, []float64{0}), + InitialCovariance: mustFloats(t, []float64{2}, 1, 1), + ProcessNoise: mustFloats(t, []float64{0}, 1, 1), + MeasurementNoise: mustFloats(t, []float64{1e18}, 1, 1), + SigmaAlpha: 1, + SigmaBeta: 2, + SigmaKappa: 0, + }) + if err != nil { + t.Fatalf("UnscentedKalmanFilter: %v", err) + } + if got, want := f.States.FloatAt(0), 2.0; math.Abs(got-want) > 1e-9 { + t.Fatalf("predicted second moment = %.12f, want %.12f", got, want) + } +} + +// TestExtendedKalmanBearingOnly tracks a constant-velocity target from +// bearing-only measurements, the classical case where the observation +// map is genuinely nonlinear and the posterior is not Gaussian. The +// honest pin is Monte Carlo: a bootstrap particle filter over the same +// model and data approximates the true posterior mean, and both the +// extended and the unscented filter must land within a few posterior +// standard deviations of it, and of each other. +func TestExtendedKalmanBearingOnly(t *testing.T) { + const ( + n = 60 + sigmaB = 0.005 + truthX0 = 2000.0 + truthY0 = 100.0 + truthVX = -15.0 + truthVY = 3.0 + ) + // Truth and bearings, through the house generator. + g := core.NewGenerator(101) + z := make([]float64, n) + for t := range n { + px := truthX0 + truthVX*float64(t+1) + py := truthY0 + truthVY*float64(t+1) + z[t] = math.Atan2(py, px) + sigmaB*g.NormalUnit() + } + cv := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{ + x.FloatAt(0) + x.FloatAt(2), + x.FloatAt(1) + x.FloatAt(3), + x.FloatAt(2), + x.FloatAt(3), + }, 4) + } + bearing := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{math.Atan2(x.FloatAt(1), x.FloatAt(0))}, 1) + } + bearingJac := func(x *core.Array) (*core.Array, error) { + px, py := x.FloatAt(0), x.FloatAt(1) + r2 := px*px + py*py + return core.FromFloats([]float64{-py / r2, px / r2, 0, 0}, 1, 4) + } + opts := KalmanOptions{ + InitialState: mustFloats(t, []float64{1900, 80, 0, 0}), + InitialCovariance: mustFloats(t, []float64{1e4, 0, 0, 0, 0, 1e4, 0, 0, 0, 0, 25, 0, 0, 0, 0, 25}, 4, 4), + ProcessNoise: mustFloats(t, []float64{1e-4, 0, 0, 0, 0, 1e-4, 0, 0, 0, 0, 0.01, 0, 0, 0, 0, 0.01}, 4, 4), + MeasurementNoise: mustFloats(t, []float64{sigmaB * sigmaB}, 1, 1), + ObservationJacobian: bearingJac, + } + zArr := mustFloats(t, z) + ekf, err := ExtendedKalmanFilter(zArr, cv, bearing, opts) + if err != nil { + t.Fatalf("ExtendedKalmanFilter: %v", err) + } + ukf, err := UnscentedKalmanFilter(zArr, cv, bearing, opts) + if err != nil { + t.Fatalf("UnscentedKalmanFilter: %v", err) + } + // The bootstrap particle filter: 20000 particles propagated + // through the same dynamics, weighted by the bearing likelihood, + // systematically resampled every step. + const particles = 20000 + pg := core.NewGenerator(202) + px := make([]float64, particles) + py := make([]float64, particles) + vx := make([]float64, particles) + vy := make([]float64, particles) + weights := make([]float64, particles) + for i := range particles { + px[i] = 1900 + 100*pg.NormalUnit() + py[i] = 80 + 100*pg.NormalUnit() + vx[i] = 0 + 5*pg.NormalUnit() + vy[i] = 0 + 5*pg.NormalUnit() + } + var meanX, meanY, varX, varY float64 + for t := range n { + for i := range particles { + px[i] += vx[i] + 0.01*pg.NormalUnit() + py[i] += vy[i] + 0.01*pg.NormalUnit() + vx[i] += 0.1 * pg.NormalUnit() + vy[i] += 0.1 * pg.NormalUnit() + } + // Log weights against this step's bearing. + maxLog := math.Inf(-1) + for i := range particles { + res := (z[t] - math.Atan2(py[i], px[i])) / sigmaB + weights[i] = -0.5 * res * res + maxLog = max(maxLog, weights[i]) + } + total := 0.0 + for i := range particles { + weights[i] = math.Exp(weights[i] - maxLog) + total += weights[i] + } + // The weighted posterior moments, recorded at the last step. + if t == n-1 { + for i := range particles { + w := weights[i] / total + meanX += w * px[i] + meanY += w * py[i] + } + for i := range particles { + w := weights[i] / total + varX += w * (px[i] - meanX) * (px[i] - meanX) + varY += w * (py[i] - meanY) * (py[i] - meanY) + } + } + // Systematic resampling: one uniform start, N evenly spaced + // pointers walked through the cumulative weights. + ancestorsX := make([]float64, particles) + ancestorsY := make([]float64, particles) + ancestorsVX := make([]float64, particles) + ancestorsVY := make([]float64, particles) + copy(ancestorsX, px) + copy(ancestorsY, py) + copy(ancestorsVX, vx) + copy(ancestorsVY, vy) + u := pg.Unit() / particles + cum := 0.0 + j := 0 + cum += weights[0] / total + for i := range particles { + pointer := u + float64(i)/particles + for cum < pointer && j < particles-1 { + j++ + cum += weights[j] / total + } + px[i] = ancestorsX[j] + py[i] = ancestorsY[j] + vx[i] = ancestorsVX[j] + vy[i] = ancestorsVY[j] + } + } + // The EKF and UKF positions against the Monte Carlo posterior, + // with a tolerance of five posterior standard deviations. + for name, state := range map[string]float64{ + "ekf x": ekf.States.FloatAt((n-1)*4 + 0), + "ekf y": ekf.States.FloatAt((n-1)*4 + 1), + "ukf x": ukf.States.FloatAt((n-1)*4 + 0), + "ukf y": ukf.States.FloatAt((n-1)*4 + 1), + } { + want, spread := meanX, 5*math.Sqrt(varX) + if strings.HasSuffix(name, "y") { + want, spread = meanY, 5*math.Sqrt(varY) + } + if math.Abs(state-want) > spread { + t.Fatalf("%s = %.3f, Monte Carlo posterior mean %.3f with sd %.3f", + name, state, want, math.Sqrt(spread*spread/25)) + } + } + // The two filters agree with each other to well under a posterior + // standard deviation. + for i, name := range []string{"x", "y"} { + d := math.Abs(ekf.States.FloatAt((n-1)*4+i) - ukf.States.FloatAt((n-1)*4+i)) + spread := math.Sqrt(varX) + if i == 1 { + spread = math.Sqrt(varY) + } + if d > spread { + t.Fatalf("EKF and UKF %s differ by %.3f, a posterior sd of %.3f", name, d, spread) + } + } + // Plausible posterior: the filtered position tracks the truth + // within four of its own standard deviations, and every published + // covariance stayed symmetric and positive definite. + cov := kalmanResultBlock(t, ekf.Covariances, n-1, 4, 4) + errX := math.Abs(ekf.States.FloatAt((n-1)*4+0) - (truthX0 + truthVX*float64(n))) + errY := math.Abs(ekf.States.FloatAt((n-1)*4+1) - (truthY0 + truthVY*float64(n))) + if errX > 4*math.Sqrt(cov[0]) || errY > 4*math.Sqrt(cov[5]) { + t.Fatalf("final position error (%.3f, %.3f) exceeds four filtered sds (%.3f, %.3f)", + errX, errY, 4*math.Sqrt(cov[0]), 4*math.Sqrt(cov[5])) + } + if !cholPD(cov, 4) { + t.Fatal("the final EKF covariance is not positive definite") + } +} + +// TestKalmanMultiChannel pins the multi-channel path: two observed +// states with cross-coupled noise, against an independent two-by-two +// reference recursion written in the test. +func TestKalmanMultiChannel(t *testing.T) { + const n = 100 + g := core.NewGenerator(13) + z := make([]float64, 2*n) + for i := range z { + z[i] = g.NormalUnit() + } + f, err := KalmanFilter(mustFloats(t, z, n, 2), mustFloats(t, []float64{1, 1, 0, 1}, 2, 2), + mustFloats(t, []float64{1, 0, 0, 1}, 2, 2), + KalmanOptions{ + ProcessNoise: mustFloats(t, []float64{1e-4, 0, 0, 1e-2}, 2, 2), + MeasurementNoise: mustFloats(t, []float64{0.1, 0, 0, 0.4}, 2, 2), + }) + if err != nil { + t.Fatalf("KalmanFilter: %v", err) + } + // The independent reference: the same recursion with a + // standard-form update and an explicit two-by-two solve. + x := []float64{0, 0} + p := []float64{1, 0, 0, 1} + logLik := 0.0 + solve2 := func(s []float64, b []float64) []float64 { + det := s[0]*s[3] - s[1]*s[2] + return []float64{(s[3]*b[0] - s[1]*b[1]) / det, (s[0]*b[1] - s[2]*b[0]) / det} + } + // mat2 multiplies two row-major two-by-two matrices. + mat2 := func(x, y []float64) []float64 { + return []float64{ + x[0]*y[0] + x[1]*y[2], x[0]*y[1] + x[1]*y[3], + x[2]*y[0] + x[3]*y[2], x[2]*y[1] + x[3]*y[3], + } + } + for step := range n { + xp := []float64{x[0] + x[1], x[1]} + pp := []float64{ + p[0] + p[1] + p[2] + p[3] + 1e-4, p[1] + p[3], + p[2] + p[3], p[3] + 1e-2, + } + inn := []float64{z[2*step] - xp[0], z[2*step+1] - xp[1]} + s := []float64{pp[0] + 0.1, pp[1], pp[2], pp[3] + 0.4} + k := [][]float64{ + solve2(s, []float64{pp[0], pp[2]}), + solve2(s, []float64{pp[1], pp[3]}), + } + solved := solve2(s, inn) + quad := inn[0]*solved[0] + inn[1]*solved[1] + det := s[0]*s[3] - s[1]*s[2] + logLik += -0.5 * (2*math.Log(2*math.Pi) + math.Log(det) + quad) + x = []float64{xp[0] + k[0][0]*inn[0] + k[0][1]*inn[1], xp[1] + k[1][0]*inn[0] + k[1][1]*inn[1]} + a := []float64{1 - k[0][0], -k[0][1], -k[1][0], 1 - k[1][1]} + at := []float64{a[0], a[2], a[1], a[3]} + kr := []float64{k[0][0] * 0.1, k[0][1] * 0.4, k[1][0] * 0.1, k[1][1] * 0.4} + kt := []float64{k[0][0], k[1][0], k[0][1], k[1][1]} + aph := mat2(a, pp) + joseph := mat2(aph, at) + krkt := mat2(kr, kt) + p = []float64{joseph[0] + krkt[0], joseph[1] + krkt[1], joseph[2] + krkt[2], joseph[3] + krkt[3]} + if got := f.States.FloatAt(2 * step); math.Abs(got-x[0]) > 1e-9 { + t.Fatalf("filtered mean 0 at %d = %.10f, want %.10f", step, got, x[0]) + } + if got := f.States.FloatAt(2*step + 1); math.Abs(got-x[1]) > 1e-9 { + t.Fatalf("filtered mean 1 at %d = %.10f, want %.10f", step, got, x[1]) + } + } + if math.Abs(f.LogLikelihood-logLik) > 1e-6 { + t.Fatalf("log-likelihood %.8f, reference %.8f", f.LogLikelihood, logLik) + } +} + +// TestKalmanCovarianceGuards walks a long filter and demands every +// published covariance come back symmetric and positive definite, +// then pins the input gates and the step-named failure reports. +func TestKalmanCovarianceGuards(t *testing.T) { + const n = 200 + g := core.NewGenerator(17) + z := make([]float64, n) + for i := range z { + z[i] = g.NormalUnit() + } + f, err := KalmanFilter(mustFloats(t, z), mustFloats(t, []float64{1, 1, 0, 1}, 2, 2), + mustFloats(t, []float64{1, 0}, 1, 2), + KalmanOptions{ + ProcessNoise: mustFloats(t, []float64{1e-4, 0, 0, 1e-2}, 2, 2), + MeasurementNoise: mustFloats(t, []float64{0.25}, 1, 1), + }) + if err != nil { + t.Fatalf("KalmanFilter: %v", err) + } + for step := range n { + cov := kalmanResultBlock(t, f.Covariances, step, 2, 2) + if math.Abs(cov[1]-cov[2]) > 1e-12*math.Max(math.Abs(cov[1]), math.Abs(cov[2])) { + t.Fatalf("covariance at %d is not symmetric: %g against %g", step, cov[1], cov[2]) + } + if !cholPD(cov, 2) { + t.Fatalf("covariance at %d is not positive definite", step) + } + if !cholPD(kalmanResultBlock(t, f.InnovationCovariances, step, 1, 1), 1) { + t.Fatalf("innovation covariance at %d is not positive definite", step) + } + } + // The input gates. + good := mustFloats(t, []float64{1, 1, 0, 1}, 2, 2) + hGood := mustFloats(t, []float64{1, 0}, 1, 2) + zOne := mustFloats(t, []float64{0.5}) + bad := []struct { + name string + run func() (*KalmanResult, error) + }{ + {"nil transition", func() (*KalmanResult, error) { return KalmanFilter(zOne, nil, hGood, KalmanOptions{}) }}, + {"non-square transition", func() (*KalmanResult, error) { + return KalmanFilter(zOne, mustFloats(t, []float64{1, 1}, 1, 2), hGood, KalmanOptions{}) + }}, + {"observation width mismatch", func() (*KalmanResult, error) { + return KalmanFilter(zOne, good, mustFloats(t, []float64{1, 0, 0}, 1, 3), KalmanOptions{}) + }}, + {"asymmetric process noise", func() (*KalmanResult, error) { + return KalmanFilter(zOne, good, hGood, KalmanOptions{ProcessNoise: mustFloats(t, []float64{1, 1, 0, 1}, 2, 2)}) + }}, + {"singular measurement noise", func() (*KalmanResult, error) { + return KalmanFilter(zOne, good, hGood, KalmanOptions{MeasurementNoise: mustFloats(t, []float64{0}, 1, 1)}) + }}, + {"non-positive-definite initial covariance", func() (*KalmanResult, error) { + return KalmanFilter(zOne, good, hGood, KalmanOptions{ + InitialCovariance: mustFloats(t, []float64{1, 0, 0, -1}, 2, 2)}) + }}, + {"empty measurements", func() (*KalmanResult, error) { + return KalmanFilter(mustFloats(t, nil), good, hGood, KalmanOptions{}) + }}, + {"non-finite measurement", func() (*KalmanResult, error) { + return KalmanFilter(mustFloats(t, []float64{math.NaN()}), good, hGood, KalmanOptions{}) + }}, + {"rank-3 measurements", func() (*KalmanResult, error) { + return KalmanFilter(mustFloats(t, []float64{1, 0, 1, 0, 1, 0, 1, 0}, 2, 2, 2), good, hGood, KalmanOptions{}) + }}, + } + for _, b := range bad { + if _, err := b.run(); err == nil { + t.Errorf("%s accepted", b.name) + } + } + // The nonlinear filters need the state dimension spelled out. + cv := func(x *core.Array) (*core.Array, error) { return x, nil } + id := func(x *core.Array) (*core.Array, error) { return x, nil } + if _, err := ExtendedKalmanFilter(zOne, cv, id, KalmanOptions{}); err == nil { + t.Error("the extended filter accepted no initial state") + } + if _, err := UnscentedKalmanFilter(zOne, cv, id, KalmanOptions{}); err == nil { + t.Error("the unscented filter accepted no initial state") + } + if _, err := UnscentedKalmanFilter(zOne, cv, id, KalmanOptions{ + InitialState: mustFloats(t, []float64{0}), + SigmaAlpha: -1, + }); err == nil { + t.Error("the unscented filter accepted a negative alpha") + } + if _, err := UnscentedKalmanFilter(zOne, cv, id, KalmanOptions{ + InitialState: mustFloats(t, []float64{0}), + SigmaAlpha: 1e-3, + SigmaKappa: -2, + }); err == nil { + t.Error("the unscented filter accepted a kappa that kills the sigma scale") + } + // A singular state covariance mid-run is refused with the step + // named: a transition that annihilates every state with zero + // process noise collapses the predicted spread onto a point, and + // the next predict cannot factor it. + annihilate := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{0}, 1) + } + _, err = UnscentedKalmanFilter(mustFloats(t, []float64{0, 0}), annihilate, id, KalmanOptions{ + InitialState: mustFloats(t, []float64{0}), + InitialCovariance: mustFloats(t, []float64{2}, 1, 1), + ProcessNoise: mustFloats(t, []float64{0}, 1, 1), + MeasurementNoise: mustFloats(t, []float64{1e18}, 1, 1), + SigmaAlpha: 1, + }) + if err == nil || !strings.Contains(err.Error(), "at step 1") { + t.Fatalf("the collapsed sigma spread was not reported at step 1, got %v", err) + } +} + +// TestKalmanNonlinearGates pins the validation branches the nonlinear +// filters own: callback presence, the state dimension's sources, the +// complex and shape gates on every array the filters read, and the +// step-named reports for failing or misbehaving callbacks. +func TestKalmanNonlinearGates(t *testing.T) { + identity := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{x.FloatAt(0)}, 1) + } + doubling := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{2 * x.FloatAt(0)}, 1) + } + base := KalmanOptions{ + InitialState: mustFloats(t, []float64{1}), + InitialCovariance: mustFloats(t, []float64{1}, 1, 1), + MeasurementNoise: mustFloats(t, []float64{1}, 1, 1), + } + zOne := mustFloats(t, []float64{1}) + if _, err := ExtendedKalmanFilter(zOne, nil, identity, base); err == nil { + t.Error("the extended filter accepted a nil transition") + } + if _, err := UnscentedKalmanFilter(zOne, identity, nil, base); err == nil { + t.Error("the unscented filter accepted a nil observation") + } + // The state dimension may come from the initial covariance alone. + withoutState := base + withoutState.InitialState = nil + run, err := UnscentedKalmanFilter(zOne, doubling, identity, withoutState) + if err != nil { + t.Fatalf("the unscented filter refused a dimension from the initial covariance: %v", err) + } + // The gain from a unit prior under the doubling map: pp = 4, + // s = 5, K = 4/5, and the measurement 1 lands the state at 4/5. + if math.Abs(run.States.FloatAt(0)-0.8) > 1e-12 { + t.Fatalf("state %.12g from a covariance-only start, want 0.8", run.States.FloatAt(0)) + } + badRank, _ := core.FromFloats([]float64{1}, 1) + covBadRank := withoutState + covBadRank.InitialCovariance = badRank + if _, err := ExtendedKalmanFilter(zOne, doubling, identity, covBadRank); err == nil { + t.Error("a rank-1 initial covariance accepted") + } + // The complex gates on the arrays the filters read. + complexState, _ := core.ComplexFromArray([]complex128{1}, 1) + withComplex := base + withComplex.InitialState = complexState + if _, err := ExtendedKalmanFilter(zOne, doubling, identity, withComplex); err == nil { + t.Error("a complex initial state accepted") + } + complexCov, _ := core.ComplexFromArray([]complex128{1}, 1, 1) + withComplexCov := withoutState + withComplexCov.InitialCovariance = complexCov + if _, err := ExtendedKalmanFilter(zOne, doubling, identity, withComplexCov); err == nil { + t.Error("a complex initial covariance accepted") + } + complexZ, _ := core.ComplexFromArray([]complex128{1}, 1) + if _, err := ExtendedKalmanFilter(complexZ, doubling, identity, base); err == nil { + t.Error("complex measurements accepted") + } + if _, err := UnscentedKalmanFilter(complexZ, doubling, identity, base); err == nil { + t.Error("complex measurements accepted by the unscented filter") + } + // A measurement stack with zero width is refused before the walk. + empty, _ := core.FromFloats([]float64{}, 2, 0) + if _, err := ExtendedKalmanFilter(empty, doubling, identity, base); err == nil { + t.Error("zero-width measurements accepted") + } + // A misbehaving transition is reported with the step. + exploding := func(x *core.Array) (*core.Array, error) { + return nil, errors.New("no") + } + _, err = ExtendedKalmanFilter(mustFloats(t, []float64{1, 1}), exploding, identity, base) + if err == nil || !strings.Contains(err.Error(), "the transition") || !strings.Contains(err.Error(), "at step 1") { + t.Fatalf("a failing transition was not reported at its step, got %v", err) + } + // A transition with the wrong output length likewise. + wrong := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 2}, 2) + } + if _, err := ExtendedKalmanFilter(mustFloats(t, []float64{1, 1}), wrong, identity, base); err == nil { + t.Error("a transition with the wrong output length accepted") + } + if _, err := UnscentedKalmanFilter(mustFloats(t, []float64{1, 1}), identity, wrong, base); err == nil { + t.Error("an observation with the wrong output length accepted") + } + // A non-finite transition output is refused where it appears. + nonFinite := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{math.NaN()}, 1) + } + if _, err := UnscentedKalmanFilter(mustFloats(t, []float64{1, 1}), nonFinite, identity, base); err == nil { + t.Error("a non-finite transition output accepted") + } + // A Jacobian of the wrong shape or a failing one is reported. + badJac := func(x *core.Array) (*core.Array, error) { + return core.FromFloats([]float64{1, 1}, 2, 1) + } + withBadJac := base + withBadJac.ObservationJacobian = badJac + if _, err := ExtendedKalmanFilter(mustFloats(t, []float64{1, 1}), doubling, identity, withBadJac); err == nil { + t.Error("an observation Jacobian of the wrong shape accepted") + } + errJac := func(x *core.Array) (*core.Array, error) { + return nil, errors.New("no") + } + withErrJac := base + withErrJac.TransitionJacobian = errJac + _, err = ExtendedKalmanFilter(mustFloats(t, []float64{1, 1}), doubling, identity, withErrJac) + if err == nil || !strings.Contains(err.Error(), "the transition Jacobian") { + t.Fatalf("a failing transition Jacobian was not reported, got %v", err) + } + withErrObsJac := base + withErrObsJac.ObservationJacobian = errJac + _, err = ExtendedKalmanFilter(mustFloats(t, []float64{1, 1}), doubling, identity, withErrObsJac) + if err == nil || !strings.Contains(err.Error(), "the observation Jacobian") { + t.Fatalf("a failing observation Jacobian was not reported, got %v", err) + } + // The linear filter's rank gate on the transition matrix. + if _, err := KalmanFilter(mustFloats(t, []float64{1}), mustFloats(t, []float64{1, 1}, 1, 2), + mustFloats(t, []float64{1, 0}, 1, 2), KalmanOptions{}); err == nil { + t.Error("a rank-1 transition matrix accepted") + } +} diff --git a/signal/misc_bench_test.go b/signal/misc_bench_test.go new file mode 100644 index 0000000..376733c --- /dev/null +++ b/signal/misc_bench_test.go @@ -0,0 +1,159 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "testing" + + core "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Signal bulk-path benchmarks: the wavelet filter banks, the spectral +// estimators and the finite-difference stencils at analysis-like sizes. + +func benchWave(b *testing.B, seed, n int, shape ...int) *core.Array { + b.Helper() + v := make([]float64, n) + for i := range v { + v[i] = float64(i%31)*float64(seed%5)*0.25 + float64(i%13) - 6 + } + a, err := core.FromFloats(v, shape...) + if err != nil { + b.Fatal(err) + } + return a +} + +func BenchmarkDWT4096(b *testing.B) { + x := benchWave(b, 1, 4096, 4096) + b.ReportAllocs() + for b.Loop() { + if _, err := DWT(x, 8); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkIDWT4096(b *testing.B) { + x := benchWave(b, 2, 4096, 4096) + coef, err := DWT(x, 8) + if err != nil { + b.Fatal(err) + } + b.ResetTimer() + b.ReportAllocs() + for b.Loop() { + if _, err := IDWT(coef, 8); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCWTMorlet(b *testing.B) { + x := benchWave(b, 3, 4096, 4096) + scales := make([]float64, 32) + for i := range scales { + scales[i] = float64(int(1) << (i / 4)) + } + b.ReportAllocs() + for b.Loop() { + if _, err := CWT(x, Morlet, scales, 1.0/512); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkWelchPSD(b *testing.B) { + x := benchWave(b, 4, 1<<16, 1<<16) + b.ReportAllocs() + for b.Loop() { + if _, _, err := WelchPSD(x, 1024, 1024, 512, "hann"); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLombScargle(b *testing.B) { + t := benchWave(b, 5, 2048, 2048) + y := benchWave(b, 6, 2048, 2048) + b.ReportAllocs() + for b.Loop() { + if _, _, err := LombScargle(t, y, 0.01, 1, 512); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSavitzkyGolay(b *testing.B) { + x := benchWave(b, 7, 1<<18, 1<<18) + b.ReportAllocs() + for b.Loop() { + if _, err := SavitzkyGolay(x, 11, 3); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkFilterApply(b *testing.B) { + x := benchWave(b, 8, 1<<18, 1<<18) + coefB, coefA, err := ButterworthLowPass(4, 1024, 64) + if err != nil { + b.Fatal(err) + } + b.ResetTimer() + b.ReportAllocs() + for b.Loop() { + if _, err := FilterApply(coefB, coefA, x); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLaplacian2D(b *testing.B) { + g := benchWave(b, 9, 512*512, 512, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := Laplacian(g, 1, 1); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkARMASpectrum(b *testing.B) { + res := &ARMAResult{ + AR: []float64{0.75, -0.5, 0.2, -0.1}, + MA: []float64{0.4, -0.25, 0.1}, + InnovationVariance: 1.5, + } + b.ReportAllocs() + for b.Loop() { + if _, _, err := ARMASpectrum(res, 4096); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNUFFTType1(b *testing.B) { + const n = 4096 + xv := make([]float64, n) + cv := make([]complex128, n) + for i := range n { + xv[i] = float64(i%1024)/1024 - 0.5 + cv[i] = complex(float64(i%7)*0.5-1, float64(i%5)*0.25) + } + xs, err := core.FromFloats(xv, n) + if err != nil { + b.Fatal(err) + } + cs, err := core.FromComplexes(cv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := NUFFTType1(xs, cs, 2048); err != nil { + b.Fatal(err) + } + } +} diff --git a/signal/nan_propagation_pins_test.go b/signal/nan_propagation_pins_test.go new file mode 100644 index 0000000..ae7f301 --- /dev/null +++ b/signal/nan_propagation_pins_test.go @@ -0,0 +1,177 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// NaN-propagation pins: max pools that silently dropped NaN, empty +// windows behind oversized padding, estimators that published NaN +// without an error, and the odd-length inverse real FFT with no test +// at all. + +// TestMaxPoolNaNPropagates: v > best reads false against NaN, so the +// max pools answered the largest finite neighbour and hid the NaN; a +// NaN in a window now wins the comparison and propagates. +func TestMaxPoolNaNPropagates(t *testing.T) { + lane := mustFromFloats(t, []float64{1, 3, math.NaN(), 2}, 1, 1, 4) + got, err := MaxPool1D(lane, 2, 2, 0) + if err != nil { + t.Fatal(err) + } + if v0, _ := core.FloatAt(got, 0, 0, 0); v0 != 3 { + t.Errorf("MaxPool1D finite window: got %v, want 3", v0) + } + if v1, _ := core.FloatAt(got, 0, 0, 1); !math.IsNaN(v1) { + t.Errorf("MaxPool1D window with NaN: got %v, want NaN", v1) + } + sq := mustFromFloats(t, []float64{ + math.NaN(), 2, 3, + 4, 5, 6, + 7, 8, 9, + }, 1, 1, 3, 3) + got2, err := MaxPool2D(sq, 2, 1, 0) + if err != nil { + t.Fatal(err) + } + // The top-left window holds the NaN; the bottom-right one does not. + if v, _ := core.FloatAt(got2, 0, 0, 0, 0); !math.IsNaN(v) { + t.Errorf("MaxPool2D window with NaN: got %v, want NaN", v) + } + if v, _ := core.FloatAt(got2, 0, 0, 1, 1); v != 9 { + t.Errorf("MaxPool2D finite window: got %v, want 9", v) + } + // The adaptive and global variants share the walk. + got3, err := AdaptiveMaxPool2D(sq, 2, 2) + if err != nil { + t.Fatal(err) + } + if v, _ := core.FloatAt(got3, 0, 0, 0, 0); !math.IsNaN(v) { + t.Errorf("AdaptiveMaxPool2D window with NaN: got %v, want NaN", v) + } + got4, err := GlobalMaxPool1D(mustFromFloats(t, []float64{1, 2, 3, math.NaN()}, 1, 1, 4)) + if err != nil { + t.Fatal(err) + } + if v, _ := core.FloatAt(got4, 0, 0, 0); !math.IsNaN(v) { + t.Errorf("GlobalMaxPool1D with NaN: got %v, want NaN", v) + } +} + +// TestPoolPaddingBelowKernelRefused: padding at or above the kernel +// leaves output windows entirely inside the padding, where a max +// answered −Inf and an average 0; such configurations are refused. +func TestPoolPaddingBelowKernelRefused(t *testing.T) { + lane := mustFromFloats(t, []float64{1, 2, 3, 4}, 1, 1, 4) + if _, err := MaxPool1D(lane, 1, 1, 2); err == nil || !strings.Contains(err.Error(), "below the kernel") { + t.Errorf("MaxPool1D padding 2 kernel 1: err = %v", err) + } + if _, err := AvgPool1D(lane, 1, 1, 2, true); err == nil || !strings.Contains(err.Error(), "below the kernel") { + t.Errorf("AvgPool1D padding 2 kernel 1: err = %v", err) + } + // padding == kernel empties the first window too, so it is + // refused all the same. + sq := mustFromFloats(t, make([]float64, 9), 1, 1, 3, 3) + if _, err := MaxPool2D(sq, 2, 1, 2); err == nil || !strings.Contains(err.Error(), "below the kernel") { + t.Errorf("MaxPool2D padding equal to the kernel: err = %v", err) + } + cube := mustFromFloats(t, make([]float64, 8), 1, 1, 2, 2, 2) + if _, err := MaxPool3D(cube, [3]int{2, 1, 1}, [3]int{1, 1, 1}, [3]int{2, 0, 0}); err == nil || !strings.Contains(err.Error(), "below the kernel") { + t.Errorf("MaxPool3D oversized padding: err = %v", err) + } + // Valid configurations are untouched: padding below the kernel + // still pools, and every window keeps at least one real sample. + if _, err := MaxPool1D(lane, 2, 2, 1); err != nil { + t.Errorf("MaxPool1D padding 1 kernel 2 refused: %v", err) + } + if _, err := MaxPool2D(sq, 3, 1, 1); err != nil { + t.Errorf("MaxPool2D padding 1 kernel 3 refused: %v", err) + } +} + +// TestCorrelatePeriodogramNonFiniteRefused: a non-finite sample drove +// the estimators' normalisers NaN and published NaN estimates with no +// error; the inputs are refused up front now. +func TestCorrelatePeriodogramNonFiniteRefused(t *testing.T) { + clean := mustFromFloats(t, []float64{1, 2, 3, 4, 5}, 5) + nan := mustFromFloats(t, []float64{1, 2, math.NaN(), 4, 5}, 5) + inf := mustFromFloats(t, []float64{1, 2, math.Inf(1), 4, 5}, 5) + if _, err := Autocorrelate(nan, 2); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Errorf("Autocorrelate with NaN: err = %v", err) + } + if _, err := PartialAutocorrelate(nan, 2); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Errorf("PartialAutocorrelate with NaN: err = %v", err) + } + if _, err := CrossCorrelate(clean, inf); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Errorf("CrossCorrelate with Inf: err = %v", err) + } + times := mustFromFloats(t, []float64{0, 1, 2, 3, math.NaN()}, 5) + vals := mustFromFloats(t, []float64{1, -1, 1, -1, 1}, 5) + if _, _, err := LombScargle(times, vals, 0.1, 0.4, 8); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Errorf("LombScargle with NaN times: err = %v", err) + } + timesClean := mustFromFloats(t, []float64{0, 1, 2, 3, 4}, 5) + valsBad := mustFromFloats(t, []float64{1, -1, math.NaN(), -1, 1}, 5) + if _, _, err := LombScargle(timesClean, valsBad, 0.1, 0.4, 8); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Errorf("LombScargle with NaN values: err = %v", err) + } + // Clean inputs still estimate. + if _, err := Autocorrelate(clean, 2); err != nil { + t.Errorf("Autocorrelate clean: %v", err) + } +} + +// TestIRFFTOddLengthRoundTrip: the odd-length inverse real FFT had no +// coverage; the spectrum must match the naive DFT and the round trip +// must restore the series. +func TestIRFFTOddLengthRoundTrip(t *testing.T) { + for _, n := range []int{5, 7, 9} { + vals := make([]float64, n) + g := core.NewGenerator(int64(n)) + for i := range vals { + f, _ := core.Floats(g, 1) + v, _ := core.FloatAt(f, 0) + vals[i] = v*10 - 5 + } + src := mustFromFloats(t, vals, n) + spec, err := RFFT(src) + if err != nil { + t.Fatalf("RFFT(%d): %v", n, err) + } + complexVals := make([]complex128, n) + for i := range vals { + complexVals[i] = complex(vals[i], 0) + } + naive := naiveDFT(complexVals, -1) + half := n/2 + 1 + for k := range half { + got, _ := core.ComplexAt(spec, k) + if math.Abs(real(got)-real(naive[k])) > 1e-9 || math.Abs(imag(got)-imag(naive[k])) > 1e-9 { + t.Fatalf("RFFT(%d)[%d]: got %v, want %v", n, k, got, naive[k]) + } + } + back, err := IRFFT(spec, n) + if err != nil { + t.Fatalf("IRFFT(%d): %v", n, err) + } + if back.Len() != n { + t.Fatalf("IRFFT(%d) length %d, want %d", n, back.Len(), n) + } + for i := range n { + got, _ := core.FloatAt(back, i) + if math.Abs(got-vals[i]) > 1e-9 { + t.Fatalf("IRFFT(%d)[%d]: got %.14g, want %.14g", n, i, got, vals[i]) + } + } + } + // The bin count of an odd spectrum is (n+1)/2, so the documented + // default n = 2·(half−1) cannot recover an odd length: odd + // round trips must pass n explicitly, which is exactly the path + // checked above. +} diff --git a/signal/nufft.go b/signal/nufft.go new file mode 100644 index 0000000..12c3a16 --- /dev/null +++ b/signal/nufft.go @@ -0,0 +1,124 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// The type-1 nonuniform fast Fourier transform. Interferometric +// imaging and irregularly sampled spectroscopy ask for the spectrum +// of data whose samples sit at arbitrary coordinates: f_k = +// Σ_j c_j·e^{2πi·k·x_j} over a uniform output grid, with the x_j +// spread anywhere in [−1/2, 1/2). The direct sum costs O(n·m); the +// gridding route costs O(n + m·log m) by scattering the samples onto +// an oversampled grid through a localised Gaussian kernel, running +// one FFT, and undoing the kernel's own Fourier footprint by pointwise +// division. The Gaussian is the kernel of choice because its Fourier +// image is Gaussian too, so the deconvolution is as local as the +// spreading. + +// NUFFTType1 computes f_k = Σ_j c_j·e^{2πi·k·x_j} for k = 0 … n−1, +// where x holds the nonuniform sample coordinates in [−1/2, 1/2) and +// c the complex values. The answer carries the gridding error of the +// Gaussian kernel, a few digits short of the direct sum at the +// default kernel width, in exchange for the FFT's speed. Coordinates +// outside [−1/2, 1/2), mismatched lengths, or a non-positive output +// size is an error. +func NUFFTType1(x, c *core.Array, n int) (*core.Array, error) { + const name = "NUFFTType1" + if x.NDim() != 1 || c.NDim() != 1 { + return nil, base.Errf("%s: coordinates and values must be vectors", name) + } + if x.Len() != c.Len() { + return nil, base.Errf("%s: %d coordinates for %d values", name, x.Len(), c.Len()) + } + if x.Dtype() == core.Complex { + return nil, base.Errf("%s: the coordinates must be real, got %s", name, x.Dtype()) + } + if n <= 0 { + return nil, base.Errf("%s: the output size must be positive, got %d", name, n) + } + count := x.Len() + // Oversampling factor and kernel geometry: the wider the kernel + // relative to the grid, the smaller the aliasing and truncation + // error, at linear cost in the spreading step. + upsampled := 8 + for upsampled < 4*n { + upsampled *= 2 + } + const sigma = 2.0 + const radius = 10 + grid := make([]complex128, upsampled) + coords := make([]float64, count) + for j := range count { + coords[j] = x.FloatAt(j) + // NaN defeats both range comparisons below, so it is refused + // by name: it would poison the whole grid and the answer + // would come back all-NaN with no error. + if math.IsNaN(coords[j]) { + return nil, base.Errf("%s: coordinate %d is NaN", name, j) + } + if coords[j] < -0.5 || coords[j] >= 0.5 { + return nil, base.Errf("%s: coordinate %d = %g lies outside [−1/2, 1/2)", name, j, coords[j]) + } + } + // Spread: each sample lands on the nearest grid point and bleeds + // into its radius neighbours through the kernel. The sample's value is + // read once, not once per tap. + for j := range count { + g := coords[j] * float64(upsampled) + near := math.Floor(g + 0.5) + frac := g - near + origin := int(near) + cj := c.ComplexAt(j) + for o := -radius; o <= radius; o++ { + idx := ((origin+o)%upsampled + upsampled) % upsampled + t := float64(o) - frac + grid[idx] += cj * complex(math.Exp(-t*t/(2*sigma*sigma)), 0) + } + } + // One inverse FFT of the oversampled grid, scaled by 1/m exactly as + // the IFFT entry point applies it: a division, not a multiplication + // by the reciprocal, so the rounding matches. + transform(grid, +1) + m := complex(float64(upsampled), 0) + for k := range upsampled { + grid[k] /= m + } + // Undo the kernel: the discrete Fourier image of the sampled + // Gaussian at frequency k, computed by the same short sum. The + // Gaussian envelope does not depend on k, so it is built once. + gaussian := make([]float64, 2*radius+1) + for o := -radius; o <= radius; o++ { + gaussian[o+radius] = math.Exp(-float64(o*o) / (2 * sigma * sigma)) + } + out, oerr := core.Zeros(core.Complex, []int{n}...) + if oerr != nil { + return nil, oerr + } + spectrum := out.RawComplexes() + for k := range n { + ku := k + if ku > upsampled/2 { + ku = upsampled - ku + } + image := complex(0, 0) + // The same taps in the same order; the row index is carried by + // the range so the kernel weight needs no index add. The + // rotation comes straight from math.Sincos, whose pair is the + // one CmplxPolar multiplies through (verified bit-for-bit), so + // each tap halves its trig work. + for oi := range gaussian { + o := oi - radius + s, c := math.Sincos(-2 * math.Pi * float64(o) * float64(ku) / float64(upsampled)) + image += complex(gaussian[oi]*c, gaussian[oi]*s) + } + spectrum[k] = grid[k] / image * complex(float64(upsampled), 0) + } + return out, nil +} diff --git a/signal/nufft_test.go b/signal/nufft_test.go new file mode 100644 index 0000000..28df9bd --- /dev/null +++ b/signal/nufft_test.go @@ -0,0 +1,119 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import ( + "math" + "math/cmplx" + "testing" +) + +// nufftCase builds a deterministic nonuniform problem: coordinates +// spread over the band, complex values with structure at several +// scales. +func nufftCase(count int) ([]float64, []complex128) { + x := make([]float64, count) + c := make([]complex128, count) + for j := range count { + x[j] = -0.5 + float64(j)*0.97/float64(count) + c[j] = cmplx.Exp(complex(0.3*float64(j%7), 0.9*float64(j%5))) + } + return x, c +} + +// directDFT evaluates the type-1 sum the slow honest way. +func directDFT(x []float64, c []complex128, n int) []complex128 { + out := make([]complex128, n) + for k := range n { + s := complex(0, 0) + for j := range x { + s += c[j] * cmplx.Exp(complex(0, 2*math.Pi*float64(k)*x[j])) + } + out[k] = s + } + return out +} + +// TestNUFFTType1MatchesDirectDFT checks the gridded transform against +// the direct sum over the whole output band: the Gaussian kernel's +// guarantee is a handful of relative digits at every frequency, not +// just the low ones. +func TestNUFFTType1MatchesDirectDFT(t *testing.T) { + const count, n = 64, 64 + x, c := nufftCase(count) + xa, cerr := core.FromFloats(x, count) + if cerr != nil { + t.Fatal(cerr) + } + ca, cerr := core.FromComplexes(c, count) + if cerr != nil { + t.Fatal(cerr) + } + got, err := NUFFTType1(xa, ca, n) + if err != nil { + t.Fatalf("NUFFTType1: %v", err) + } + want := directDFT(x, c, n) + worst, scale := 0.0, 0.0 + for k := range n { + d := base.AbsComplex(got.ComplexAt(k) - want[k]) + worst = math.Max(worst, d) + scale = math.Max(scale, base.AbsComplex(want[k])) + } + if worst/scale > 1e-4 { + t.Fatalf("worst relative error %.3g, want under 1e-4", worst/scale) + } +} + +// TestNUFFTType1LargeGrid covers an output grid wider than the sample +// count, the regime interferometric imaging actually runs. +func TestNUFFTType1LargeGrid(t *testing.T) { + const count, n = 24, 96 + x := make([]float64, count) + c := make([]complex128, count) + for j := range count { + x[j] = -0.5 + (float64(j)*0.77+0.1)/float64(count) + c[j] = complex(math.Cos(float64(j)), math.Sin(float64(2*j))) + } + xa, _ := core.FromFloats(x, count) + ca, _ := core.FromComplexes(c, count) + got, err := NUFFTType1(xa, ca, n) + if err != nil { + t.Fatalf("NUFFTType1: %v", err) + } + want := directDFT(x, c, n) + worst := 0.0 + for k := range n { + worst = math.Max(worst, base.AbsComplex(got.ComplexAt(k)-want[k])) + } + if worst > 1e-4 { + t.Fatalf("worst absolute error %.3g, want under 1e-4", worst) + } +} + +// TestNUFFTType1Errors pins the validation contract. +func TestNUFFTType1Errors(t *testing.T) { + xa := mustFloats(t, []float64{-0.4, 0.1, 0.3}, 3) + ca, _ := core.FromComplexes([]complex128{1, 1i, 2}, 3) + if _, err := NUFFTType1(xa, ca, 0); err == nil { + t.Fatal("expected an error for a zero output size") + } + ca2, _ := core.FromComplexes([]complex128{1, 1i}, 2) + if _, err := NUFFTType1(xa, ca2, 8); err == nil { + t.Fatal("expected an error for mismatched lengths") + } + outside := mustFloats(t, []float64{-0.6, 0.1, 0.3}, 3) + if _, err := NUFFTType1(outside, ca, 8); err == nil { + t.Fatal("expected an error for a coordinate below −1/2") + } + atEdge := mustFloats(t, []float64{-0.4, 0.1, 0.5}, 3) + if _, err := NUFFTType1(atEdge, ca, 8); err == nil { + t.Fatal("expected an error for a coordinate at +1/2") + } +} diff --git a/signal/periodogram.go b/signal/periodogram.go new file mode 100644 index 0000000..c3f985c --- /dev/null +++ b/signal/periodogram.go @@ -0,0 +1,202 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +import ( + "math" + "sync" +) + +// Periodograms. The Lomb-Scargle periodogram answers "at which +// frequency does unevenly sampled data oscillate" without the +// interpolation a resampled FFT would need: each trial frequency gets +// its own least-squares fit of a sine and cosine through the actual +// observation times, with the phase reference τ chosen so the two +// fitted components are exactly orthogonal at that frequency. + +// lombParallelMinN is the observation count above which a single +// frequency's two sine/cosine passes are worth a worker's spawn cost. +// Below it the frequency grid walk stays on the calling goroutine. +const lombParallelMinN = 1 << 9 + +// LombScargle computes the normalised Lomb-Scargle periodogram of the +// observations values taken at times, over the nFreq frequencies +// evenly spaced from minFreq to maxFreq inclusive (a single frequency +// when nFreq is 1) and returns the frequency grid and the power at +// each frequency. The power carries the classical +// normalisation: a pure sinusoid of amplitude A at a frequency on the +// grid peaks near A²·n/(4·var(values)), so the scale is comparable +// across data sets. An empty or two-point time base, a length +// mismatch, an all-equal time base, a non-positive variance, or a +// frequency range that does not satisfy 0 < minFreq ≤ maxFreq is an +// error. +func LombScargle(times, values *core.Array, minFreq, maxFreq float64, nFreq int) (freqs, power *core.Array, err error) { + const name = "LombScargle" + if times.NDim() != 1 || values.NDim() != 1 { + return nil, nil, base.Errf("%s: times and values must be vectors", name) + } + if times.Dtype() == core.Complex || values.Dtype() == core.Complex { + return nil, nil, base.Errf("%s: complex arrays are not supported", name) + } + n := values.Len() + if n < 3 { + return nil, nil, base.Errf("%s: at least three observations are needed, got %d", name, n) + } + if times.Len() != n { + return nil, nil, base.Errf("%s: times has %d entries for %d values", name, times.Len(), n) + } + if nFreq < 1 { + return nil, nil, base.Errf("%s: nFreq must be at least 1, got %d", name, nFreq) + } + // The gate is NaN-rejecting and Inf-rejecting at once: +Inf passes + // a bare > 0, and Inf endpoints turn every interpolated frequency + // into NaN with no error. + if !(minFreq > 0) || math.IsInf(minFreq, 0) || math.IsInf(maxFreq, 0) || maxFreq < minFreq { + return nil, nil, base.Errf("%s: the frequency range must satisfy 0 < minFreq ≤ maxFreq over finite frequencies, got [%g, %g]", + name, minFreq, maxFreq) + } + t := make([]float64, n) + x := make([]float64, n) + mean := 0.0 + // A non-finite time or value would drive the variance NaN, slip + // past its gate and publish NaN powers with no error, so both + // arrays are refused up front (the guard SolvePoissonPeriodic + // applies to its source). + for i := range n { + t[i] = times.FloatAt(i) + if math.IsNaN(t[i]) || math.IsInf(t[i], 0) { + return nil, nil, base.Errf("%s: times holds the non-finite value %g at %d", name, t[i], i) + } + x[i] = values.FloatAt(i) + if math.IsNaN(x[i]) || math.IsInf(x[i], 0) { + return nil, nil, base.Errf("%s: values holds the non-finite value %g at %d", name, x[i], i) + } + mean += x[i] + } + mean /= float64(n) + // A constant time base carries no phase information: every trial + // frequency drives the sine fit to 0/0. + if allEqual(t) { + return nil, nil, base.Errf("%s: the times must not all be equal", name) + } + variance := 0.0 + for i := range n { + x[i] -= mean + variance += x[i] * x[i] + } + variance /= float64(n - 1) + if variance <= 0 { + return nil, nil, base.Errf("%s: the values have zero variance", name) + } + tCenter := t[n/2] + // The offsets from the phase centre feed every trig argument of + // every frequency: (t[i]−tCenter) is recomputed twice per + // observation per frequency, so it is evaluated once here and + // reused. The stored value is the subtraction result itself, so + // every argument keeps the exact bits it had. + dt := make([]float64, n) + for i := range n { + dt[i] = t[i] - tCenter + } + freqsArr := core.New(core.Float, nFreq) + powerArr := core.New(core.Float, nFreq) + freqRow := freqsArr.RawFloats() + powerRow := powerArr.RawFloats() + // A frequency whose sine or cosine sum vanishes cannot be fitted + // on this time base. The serial walk reported the lowest such + // frequency; the split keeps that contract by remembering the + // smallest offending index and erroring after the join. + var ( + badMu sync.Mutex + bad = -1 + ) + unresolvable := func(f int) { + badMu.Lock() + defer badMu.Unlock() + if bad < 0 || f < bad { + bad = f + } + } + // fitAt runs the whole per-frequency pipeline: the grid frequency, + // the orthogonalising phase reference τ and both least-squares + // fits. Every read is from the shared time and value slices, every + // write lands in this frequency's own slot of the two outputs, and + // the per-frequency arithmetic sequence is the serial one + // unchanged, so the split cannot move an addend. + fitAt := func(f int) { + freq := minFreq + if nFreq > 1 { + freq = minFreq + (maxFreq-minFreq)*float64(f)/float64(nFreq-1) + } + freqRow[f] = freq + omega := 2 * math.Pi * freq + // The phase reference τ keeps the sine and cosine fits + // orthogonal at this frequency. Both passes need the sine and + // the cosine of the same argument; math.Sincos shares the range + // reduction between the two and returns exactly the pair + // math.Sin and math.Cos produce (verified bit-for-bit), so the + // sums are unchanged while the trig work halves. + sumSin2, sumCos2 := 0.0, 0.0 + for i := range n { + arg := omega * dt[i] + s, c := math.Sincos(2 * arg) + sumSin2 += s + sumCos2 += c + } + tau := 0.5 * math.Atan2(sumSin2, sumCos2) / omega + sumCos, sumSin, sumCosSq, sumSinSq := 0.0, 0.0, 0.0, 0.0 + for i := range n { + arg := omega * (dt[i] - tau) + c, s := math.Sincos(arg) + sumCos += x[i] * c + sumSin += x[i] * s + sumCosSq += c * c + sumSinSq += s * s + } + if sumSinSq == 0 || sumCosSq == 0 { + unresolvable(f) + return + } + power := (sumCos*sumCos)/sumCosSq + (sumSin*sumSin)/sumSinSq + powerRow[f] = power / (2 * variance) + } + if n >= lombParallelMinN { + // The frequencies split across workers: disjoint output slots, + // per-frequency normalisations computed inside the worker that + // owns the frequency. + engine.Parallel(nFreq, func(fs, fe int) { + for f := fs; f < fe; f++ { + fitAt(f) + } + }) + } else { + for f := range nFreq { + fitAt(f) + } + } + if bad >= 0 { + freq := minFreq + if nFreq > 1 { + freq = minFreq + (maxFreq-minFreq)*float64(bad)/float64(nFreq-1) + } + return nil, nil, base.Errf("%s: the time base cannot resolve the frequency %g", name, freq) + } + return freqsArr, powerArr, nil +} + +// allEqual reports whether every slice entry matches the first. +func allEqual(v []float64) bool { + for _, x := range v[1:] { + if x != v[0] { + return false + } + } + return true +} diff --git a/signal/periodogram_test.go b/signal/periodogram_test.go new file mode 100644 index 0000000..29b0bfe --- /dev/null +++ b/signal/periodogram_test.go @@ -0,0 +1,312 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "testing" +) + +// TestLombScargleFindsPeak recovers the known 0.1-cycle-per-sample +// period from unevenly sampled data: the grid point nearest the true +// frequency must carry the largest power by a clear margin. +func TestLombScargleFindsPeak(t *testing.T) { + const trueFreq = 0.1 + n := 200 + times := make([]float64, n) + values := make([]float64, n) + for i := range n { + // Deterministic uneven sampling: every third tick skipped. + times[i] = float64(i) * 1.3 + values[i] = math.Sin(2*math.Pi*trueFreq*times[i]) + 0.3*math.Cos(2*math.Pi*0.31*times[i]) + } + freqs, power, err := LombScargle(mustFloats(t, times, n), mustFloats(t, values, n), + 0.02, 0.45, 400) + if err != nil { + t.Fatalf("LombScargle: %v", err) + } + best, bestF := 0, 0.0 + for i := range power.Len() { + if power.FloatAt(i) > float64(best) { + best = i + bestF = freqs.FloatAt(i) + } + } + if math.Abs(bestF-trueFreq) > 0.45/400*3 { + t.Fatalf("peak at %.5f, want within three bins of %.5f", bestF, trueFreq) + } + // The second tone must also stand out on the grid. + power2 := 0.0 + for i := range power.Len() { + if math.Abs(freqs.FloatAt(i)-0.31) < 0.01 && power.FloatAt(i) > power2 { + power2 = power.FloatAt(i) + } + } + if power2 <= 0 { + t.Fatal("the secondary tone left no peak") + } +} + +// TestLombScargleScale pins the classical normalisation: a pure +// unit-amplitude sinusoid sampled evenly at exactly its period peaks +// at n/(4·var) ≈ n/2 for unit-amplitude data of unit variance… checked +// against the directly evaluated defining sum instead of a hand rule. +func TestLombScargleScale(t *testing.T) { + n := 60 + times := make([]float64, n) + values := make([]float64, n) + for i := range n { + times[i] = float64(i) + values[i] = math.Sin(2 * math.Pi * float64(i) / 12) + } + _, power, err := LombScargle(mustFloats(t, times, n), mustFloats(t, values, n), + 1.0/12.0, 1.0/12.0, 1) + if err != nil { + t.Fatalf("LombScargle: %v", err) + } + // Direct reference: the two orthogonal sums at the exact tone. + mean := 0.0 + for i := range n { + mean += values[i] + } + mean /= float64(n) + variance := 0.0 + for i := range n { + d := values[i] - mean + variance += d * d + } + variance /= float64(n - 1) + omega := 2 * math.Pi / 12 + sc, ss := 0.0, 0.0 + for i := range n { + sc += math.Cos(omega * float64(i)) + ss += math.Sin(omega * float64(i)) + } + tau := 0.5 * math.Atan2(ss, sc) / omega + sumCos, sumSin, sumCosSq, sumSinSq := 0.0, 0.0, 0.0, 0.0 + for i := range n { + arg := omega * (float64(i) - tau) + c, s := math.Cos(arg), math.Sin(arg) + sumCos += (values[i] - mean) * c + sumSin += (values[i] - mean) * s + sumCosSq += c * c + sumSinSq += s * s + } + want := (sumCos*sumCos/sumCosSq + sumSin*sumSin/sumSinSq) / (2 * variance) + if got := power.FloatAt(0); math.IsNaN(got) || math.IsInf(got, 0) { + t.Fatalf("power = %v, want a finite value comparable to %.12g", got, want) + } + if math.Abs(power.FloatAt(0)-want) > 1e-9 { + t.Fatalf("power = %.12g, want the direct sum %.12g", power.FloatAt(0), want) + } +} + +// TestLombScargleErrors pins the validation contract. +func TestLombScargleErrors(t *testing.T) { + times := mustFloats(t, []float64{0, 1, 2, 3}, 4) + values := mustFloats(t, []float64{1, 2, 1, 2}, 4) + pair := mustFloats(t, []float64{0, 1}, 2) + if _, _, err := LombScargle(pair, mustFloats(t, []float64{1, 2}, 2), 0.1, 1, 5); err == nil { + t.Fatal("expected an error for a two-point time base") + } + short := mustFloats(t, []float64{0, 1, 2}, 3) + flat := mustFloats(t, []float64{1, 1, 1}, 3) + if _, _, err := LombScargle(short, flat, 0.1, 1, 5); err == nil { + t.Fatal("expected an error for zero variance") + } + if _, _, err := LombScargle(times, mustFloats(t, []float64{1, 2}, 2), 0.1, 1, 5); err == nil { + t.Fatal("expected an error for a length mismatch") + } + if _, _, err := LombScargle(times, values, 0, 1, 5); err == nil { + t.Fatal("expected an error for a non-positive minFreq") + } + if _, _, err := LombScargle(times, values, 1, 0.5, 5); err == nil { + t.Fatal("expected an error for an inverted range") + } + if _, _, err := LombScargle(times, values, 0.1, 1, 0); err == nil { + t.Fatal("expected an error for zero frequency points") + } +} + +// TestWelchPSDWhiteNoise checks the level: filtered deterministic +// samples with unit variance must estimate a flat band near one. +func TestWelchPSDWhiteNoise(t *testing.T) { + n := 4096 + x := make([]float64, n) + g := core.NewGenerator(7) + draws, err := core.Normal(g, n, 0, 1) + if err != nil { + t.Fatalf("Normal: %v", err) + } + for i := range n { + x[i] = draws.FloatAt(i) + } + _, psd, err := WelchPSD(mustFloats(t, x, n), 1000, 256, 128, "hann") + if err != nil { + t.Fatalf("WelchPSD: %v", err) + } + mean := 0.0 + for i := range psd.Len() { + mean += psd.FloatAt(i) / float64(psd.Len()) + } + // A flat unit-variance band integrates to σ² over fs/2, so the + // level sits at 2σ²/fs in PSD units of x²/Hz. + if r := mean / (2 / 1000.0); math.Abs(r-1) > 0.15 { + t.Fatalf("mean PSD = %.6f, want 2σ²/fs = %.6f (ratio %.3f)", mean, 2/1000.0, r) + } +} + +// TestWelchPSDSinusoid pins the peak location and the Parseval +// balance: the PSD integrated over frequency returns the signal +// variance. +func TestWelchPSDSinusoid(t *testing.T) { + const fs = 128.0 + const tone = 16.0 + n := 2048 + x := make([]float64, n) + for i := range n { + x[i] = math.Sin(2 * math.Pi * tone * float64(i) / fs) + } + freqs, psd, err := WelchPSD(mustFloats(t, x, n), fs, 256, 128, "hann") + if err != nil { + t.Fatalf("WelchPSD: %v", err) + } + peak, best := 0, 0.0 + for i := range psd.Len() { + if psd.FloatAt(i) > best { + best = psd.FloatAt(i) + peak = i + } + } + if math.Abs(freqs.FloatAt(peak)-tone) > fs/256 { + t.Fatalf("PSD peak at %.4f Hz, want %.1f", freqs.FloatAt(peak), tone) + } + // Parseval: sum(psd)·df ≈ variance for a windowed estimate on a + // signal with negligible edge leakage. + total, df := 0.0, fs/256 + for i := range psd.Len() { + total += psd.FloatAt(i) + } + total *= df + sq := 0.0 + for i := range n { + sq += x[i] * x[i] + } + variance := sq / float64(n) + if math.Abs(total-variance)/variance > 0.1 { + t.Fatalf("Parseval: ∫PSD = %.4f, variance = %.4f (rel %.3f)", total, variance, + math.Abs(total-variance)/variance) + } +} + +// TestWelchPSDErrors pins the validation contract. +func TestWelchPSDErrors(t *testing.T) { + x := mustFloats(t, make([]float64, 64), 64) + if _, _, err := WelchPSD(x, 100, 128, 0, "hann"); err == nil { + t.Fatal("expected an error for a segment longer than the signal") + } + if _, _, err := WelchPSD(x, 100, 32, 32, "hann"); err == nil { + t.Fatal("expected an error for a full overlap") + } + if _, _, err := WelchPSD(x, 100, 32, -1, "hann"); err == nil { + t.Fatal("expected an error for a negative overlap") + } + if _, _, err := WelchPSD(x, 0, 32, 16, "hann"); err == nil { + t.Fatal("expected an error for a zero sampling rate") + } + if _, _, err := WelchPSD(x, 100, 32, 16, "kaiser"); err == nil { + t.Fatal("expected an error for an unknown window") + } + rank2, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, _, err := WelchPSD(rank2, 100, 2, 1, "hann"); err == nil { + t.Fatal("expected an error for a rank-2 signal") + } +} + +// TestWelchPSDSegmentCount pins the segment arithmetic: a signal of +// n samples with non-overlapping segments fills exactly +// (n−overlap)/(segment−overlap) segments, the last one included. The +// third of the three segments here carries the only energy, so a +// dropped final segment would leave the Nyquist bin empty, and the +// average over three segments puts |X|²=64 at exactly 64/(3·fs·wPower). +func TestWelchPSDSegmentCount(t *testing.T) { + n := 24 + x := make([]float64, n) // segments 1 and 2 are silent + for i := 16; i < n; i++ { + x[i] = 1 // the alternating ±1 pattern shifted to +1: 16+i even + if (i-16)%2 == 1 { + x[i] = -1 + } + } + freqs, psd, err := WelchPSD(mustFloats(t, x, n), 1, 8, 0, "box") + if err != nil { + t.Fatalf("WelchPSD: %v", err) + } + if psd.Len() != 5 { + t.Fatalf("bins = %d, want 5", psd.Len()) + } + for k := range 5 { + want := 0.0 + if k == 4 { + want = 64.0 / (3 * 1 * 8) // one segment of three carries |X[4]|² = 64 + } + if math.Abs(psd.FloatAt(k)-want) > 1e-9 { + t.Fatalf("psd[%d] = %.9g, want %.9g", k, psd.FloatAt(k), want) + } + if k == 4 && math.Abs(freqs.FloatAt(k)-0.5) > 1e-9 { + t.Fatalf("bin 4 sits at %.6f Hz, want the 0.5 Nyquist", freqs.FloatAt(k)) + } + } + // A signal exactly one segment long is a legal single-segment + // periodogram, not an error. + solo, psd1, err := WelchPSD(mustFloats(t, x[16:], 8), 1, 8, 0, "box") + if err != nil { + t.Fatalf("WelchPSD single segment: %v", err) + } + if solo.Len() != 5 || psd1.Len() != 5 { + t.Fatalf("single-segment shapes = %d/%d, want 5/5", solo.Len(), psd1.Len()) + } + if math.Abs(psd1.FloatAt(4)-8) > 1e-9 { + t.Fatalf("single-segment psd[4] = %.9g, want 64/(1·1·8) = 8", psd1.FloatAt(4)) + } +} + +// TestLombScargleSingleFrequency pins the nFreq == 1 contract: the +// grid holds the single frequency minFreq and a finite power (the +// old grid formula divided 0/0 and produced NaN). +func TestLombScargleSingleFrequency(t *testing.T) { + n := 40 + times := make([]float64, n) + values := make([]float64, n) + for i := range n { + times[i] = 1.3 * float64(i) + values[i] = math.Sin(2 * math.Pi * 0.1 * times[i]) + } + freqs, power, err := LombScargle(mustFloats(t, times, n), mustFloats(t, values, n), 0.1, 0.7, 1) + if err != nil { + t.Fatalf("LombScargle: %v", err) + } + if freqs.Len() != 1 || power.Len() != 1 { + t.Fatalf("shapes %d/%d, want 1/1", freqs.Len(), power.Len()) + } + if freqs.FloatAt(0) != 0.1 { + t.Fatalf("single frequency = %g, want minFreq 0.1", freqs.FloatAt(0)) + } + if p := power.FloatAt(0); math.IsNaN(p) || math.IsInf(p, 0) || p <= 0 { + t.Fatalf("single-frequency power = %g, want a finite positive value", p) + } +} + +// TestLombScargleConstantTimesErrors pins the degenerate time base: a +// constant time base carries no phase information and used to drive +// the sine fit to 0/0 NaN. +func TestLombScargleConstantTimesErrors(t *testing.T) { + times := mustFloats(t, []float64{2, 2, 2, 2, 2}, 5) + values := mustFloats(t, []float64{1, -1, 1, -1, 1}, 5) + if _, _, err := LombScargle(times, values, 0.1, 1, 5); err == nil { + t.Fatal("expected an error for an all-equal time base") + } +} diff --git a/signal/poisson.go b/signal/poisson.go new file mode 100644 index 0000000..69983ad --- /dev/null +++ b/signal/poisson.go @@ -0,0 +1,112 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "math" + +// Spectral solution of the Poisson equation on a periodic domain. The +// Laplacian is diagonal in the Fourier basis: each plane wave +// e^{i(kx·x + ky·y)} is an eigenfunction of −Δ with eigenvalue +// (2πkx/Lx)² + (2πky/Ly)², so solving −Δu = f on the torus is one +// forward FFT, one division per mode and one inverse FFT. The +// convergence is spectral: for a smooth right-hand side the error +// falls faster than any power of the grid spacing, which no finite +// difference stencil matches at comparable cost. + +// SolvePoissonPeriodic solves −Δu = f on the periodic square +// [0, Lx] × [0, Ly] sampled on the (rows, cols) grid carried by the +// shape of f, and returns u as a float array of the same shape. Row r +// of f samples y = r·Ly/rows, column c samples x = c·Lx/cols. f must +// be a float64 array (the computation runs in float64 end to end and +// the solution is exact in no narrower dtype), and a float32, int or +// complex f is an error. +// +// A periodic solution exists only when f has zero mean: the zero +// Fourier mode of −Δu vanishes, so a nonzero mean is an error rather +// than a quietly shifted problem (the round-off band around zero is +// accepted and the constant mode of u is set to zero). A rank other +// than 2, a grid dimension below 2 or a non-positive side length is +// an error. +func SolvePoissonPeriodic(f *core.Array, lx, ly float64) (*core.Array, error) { + const name = "SolvePoissonPeriodic" + if f.NDim() != 2 { + return nil, base.Errf("%s: f must be rank 2, got shape %s", name, base.ShapeText(f.Shape())) + } + if f.Dtype() != core.Float { + return nil, base.Errf("%s: f must be a float64 array, got %s", name, f.Dtype()) + } + rows, cols := f.Shape()[0], f.Shape()[1] + if rows < 2 || cols < 2 { + return nil, base.Errf("%s: the grid must be at least 2×2, got %d×%d", name, rows, cols) + } + if lx <= 0 || ly <= 0 { + return nil, base.Errf("%s: the side lengths must be positive, got %g and %g", name, lx, ly) + } + spectrum, err := FFT2(f) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + wave := func(n int, i int, l float64) float64 { + k := i + if k > n/2 { + k -= n + } + s := 2 * math.Pi * float64(k) / l + return s * s + } + // A non-finite source would drive the compatibility tolerance NaN + // and slip through its gate, so it is refused up front. + for i := range rows * cols { + if v := f.FloatAt(i); math.IsInf(v, 0) || math.IsNaN(v) { + return nil, base.Errf("%s: f holds the non-finite value %g", name, v) + } + } + largest := 0.0 + for _, v := range poissonFloats(f) { + largest = max(largest, math.Abs(v)) + } + // The DFT of the constant mode is the plain sum of the samples; + // anything above the accumulated round-off breaks solvability. The + // tolerance is relative to the source: the sum scales with the + // sample magnitude and the element count, so an absolute floor would + // accept a tiny source with a constant mode far above its own size. + meanTolerance := 64 * base.EpsF * float64(rows*cols) * largest + mean := spectrum.ComplexAt(0) + if math.Abs(real(mean)) > meanTolerance || math.Abs(imag(mean)) > meanTolerance { + return nil, base.Errf("%s: f has a nonzero mean %g, the periodic problem has no solution", + name, real(mean)/float64(rows*cols)) + } + adjusted := make([]complex128, rows*cols) + for r := range rows { + // Rows sample y = r·ly/rows, columns sample x = c·lx/cols, so + // the row wavenumber carries ly and the column wavenumber lx. + ly2 := wave(rows, r, ly) + for c := range cols { + lambda := ly2 + wave(cols, c, lx) + if lambda == 0 { + continue // the constant mode stays zero + } + adjusted[r*cols+c] = spectrum.ComplexAt(r*cols+c) / complex(lambda, 0) + } + } + spectrumArray, err := core.FromComplexes(adjusted, rows, cols) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + inverted, err := IFFT2(spectrumArray) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + out := core.New(f.Dtype(), rows, cols) + vals := out.RawFloats() + for i := range vals { + vals[i] = real(inverted.ComplexAt(i)) + } + return out, nil +} diff --git a/signal/poisson_complex_pins_test.go b/signal/poisson_complex_pins_test.go new file mode 100644 index 0000000..5813678 --- /dev/null +++ b/signal/poisson_complex_pins_test.go @@ -0,0 +1,198 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// anisoGrid samples f(x, y) = sin(2x)·sin(y) on a rows×cols grid over +// [0, Lx]×[0, Ly], x along columns, y along rows. +func anisoGrid(t *testing.T, rows, cols int, lx, ly float64) *core.Array { + t.Helper() + vals := make([]float64, rows*cols) + for r := range rows { + y := float64(r) * ly / float64(rows) + for c := range cols { + x := float64(c) * lx / float64(cols) + vals[r*cols+c] = math.Sin(2*x) * math.Sin(y) + } + } + a, err := core.FromFloats(vals, rows, cols) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestPoissonAnisotropicDomain pins the axis-length pairing: on a +// domain with Lx ≠ Ly the eigenvalues must pair (2πk/Lx)² for columns +// with (2πm/Ly)² for rows. The swapped pairing the code used to have +// inverts the aspect ratio. +func TestPoissonAnisotropicDomain(t *testing.T) { + const rows, cols = 24, 32 + const lx, ly = 4 * math.Pi, 2 * math.Pi + f := anisoGrid(t, rows, cols, lx, ly) + u, err := SolvePoissonPeriodic(f, lx, ly) + if err != nil { + t.Fatalf("SolvePoissonPeriodic: %v", err) + } + // The analytic solution of −Δu = f on the periodic domain is + // u = f/λ: f = sin(2x)·sin(y) carries the modes 2πk/Lx = 2 (k = 4) + // and 2πm/Ly = 1 (m = 1), so λ = 2² + 1² = 5. + const lam = 5 + maxErr := 0.0 + for r := range rows { + for c := range cols { + want := f.FloatAt(r*cols+c) / lam + maxErr = max(maxErr, math.Abs(u.FloatAt(r*cols+c)-want)) + } + } + if maxErr > 1e-10 { + t.Fatalf("max error %g, want below 1e-10 (aspect ratio mixed up)", maxErr) + } +} + +// TestPoissonComplexRejects pins the dtype contract. +func TestPoissonComplexRejects(t *testing.T) { + a, err := core.FromComplexes([]complex128{1i, 2, 3, 4}, 2, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + if _, err := SolvePoissonPeriodic(a, 1, 1); err == nil { + t.Fatal("SolvePoissonPeriodic accepted complex input") + } +} + +// TestPoolingComplexRejects pins the dtype guard on the pooling +// surface: complex input must be an error, not a FloatAt panic. +func TestPoolingComplexRejects(t *testing.T) { + mk := func(shape ...int) *core.Array { + n := 1 + for _, d := range shape { + n *= d + } + vals := make([]complex128, n) + for i := range vals { + vals[i] = complex(float64(i), 1) + } + a, _ := core.FromComplexes(vals, shape...) + return a + } + in2 := mk(1, 1, 4, 4) + if _, err := MaxPool2D(in2, 2, 2, 0); err == nil { + t.Error("MaxPool2D accepted complex input") + } + if _, err := AvgPool2D(in2, 2, 2, 0, true); err == nil { + t.Error("AvgPool2D accepted complex input") + } + if _, err := AdaptiveMaxPool2D(in2, 2, 2); err == nil { + t.Error("AdaptiveMaxPool2D accepted complex input") + } + if _, err := AdaptiveAvgPool2D(in2, 2, 2); err == nil { + t.Error("AdaptiveAvgPool2D accepted complex input") + } + if _, err := GlobalAvgPool2D(in2); err == nil { + t.Error("GlobalAvgPool2D accepted complex input") + } + in1 := mk(1, 1, 8) + if _, err := MaxPool1D(in1, 2, 2, 0); err == nil { + t.Error("MaxPool1D accepted complex input") + } + if _, err := AdaptiveAvgPool1D(in1, 2); err == nil { + t.Error("AdaptiveAvgPool1D accepted complex input") + } + if _, err := GlobalMaxPool1D(in1); err == nil { + t.Error("GlobalMaxPool1D accepted complex input") + } + in3 := mk(1, 1, 4, 4, 4) + if _, err := MaxPool3D(in3, [3]int{2, 2, 2}, [3]int{2, 2, 2}, [3]int{0, 0, 0}); err == nil { + t.Error("MaxPool3D accepted complex input") + } + if _, err := AdaptiveAvgPool3D(in3, 2, 2, 2); err == nil { + t.Error("AdaptiveAvgPool3D accepted complex input") + } + if _, err := GlobalMaxPool3D(in3); err == nil { + t.Error("GlobalMaxPool3D accepted complex input") + } +} + +// TestTransformsComplexRejects pins the dtype guard on Welch, Lomb- +// Scargle, DCT/DST and the stencils. +func TestTransformsComplexRejects(t *testing.T) { + v, _ := core.FromComplexes([]complex128{1 + 1i, 2, 3 + 2i, 4, 5, 6, 7, 8}, 8) + if _, _, err := WelchPSD(v, 1.0, 4, 1, "hann"); err == nil { + t.Error("WelchPSD accepted complex input") + } + ts, _ := core.FromFloats([]float64{0, 1, 2, 3}, 4) + if _, _, err := LombScargle(ts, v, 0.1, 1, 8); err == nil { + t.Error("LombScargle accepted complex values") + } + if _, err := DCT(v, 2); err == nil { + t.Error("DCT accepted complex input") + } + if _, err := DST(v, 1); err == nil { + t.Error("DST accepted complex input") + } + if _, err := Gradient1D(v, 0.5); err == nil { + t.Error("Gradient1D accepted complex input") + } + v2, _ := core.FromComplexes([]complex128{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3) + if _, err := Laplacian(v2, 1, 1); err == nil { + t.Error("Laplacian accepted complex input") + } +} + +// TestFFTDoesNotMutateComplexInput pins the immutability contract: the +// 1-D transforms must not overwrite the caller's complex payload, which +// ComplexValues aliasing used to let happen. +func TestFFTDoesNotMutateComplexInput(t *testing.T) { + vals := []complex128{1 + 2i, 3 - 1i, -0.5 + 0.5i, 2} + a, err := core.FromComplexes(vals, 4) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + if _, err := FFT(a); err != nil { + t.Fatalf("FFT: %v", err) + } + for i, want := range vals { + if a.ComplexAt(i) != want { + t.Fatalf("FFT mutated input[%d] = %v, want %v", i, a.ComplexAt(i), want) + } + } + if _, err := IFFT(a); err != nil { + t.Fatalf("IFFT: %v", err) + } + for i, want := range vals { + if a.ComplexAt(i) != want { + t.Fatalf("IFFT mutated input[%d] = %v, want %v", i, a.ComplexAt(i), want) + } + } + // Double transform stays correct: FFT twice then IFFT twice is a + // round trip, only meaningful when neither call clobbers the input. + f1, err := FFT(a) + if err != nil { + t.Fatalf("FFT: %v", err) + } + f2, err := FFT(f1) + if err != nil { + t.Fatalf("FFT: %v", err) + } + i1, err := IFFT(f2) + if err != nil { + t.Fatalf("IFFT: %v", err) + } + i2, err := IFFT(i1) + if err != nil { + t.Fatalf("IFFT: %v", err) + } + for i, want := range vals { + got := i2.ComplexAt(i) + if math.Abs(real(got)-real(want)) > 1e-12 || math.Abs(imag(got)-imag(want)) > 1e-12 { + t.Fatalf("round trip[%d] = %v, want %v", i, got, want) + } + } +} diff --git a/signal/poisson_test.go b/signal/poisson_test.go new file mode 100644 index 0000000..84237c8 --- /dev/null +++ b/signal/poisson_test.go @@ -0,0 +1,143 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "testing" +) + +// poissonGrid samples a function of (x, y) on the periodic square +// [0, lx] × [0, ly] as a (rows × cols) array, matching the sampling +// SolvePoissonPeriodic documents. +func poissonGrid(t *testing.T, fn func(x, y float64) float64, rows, cols int, lx, ly float64) *core.Array { + t.Helper() + vals := make([]float64, rows*cols) + for r := range rows { + for c := range cols { + vals[r*cols+c] = fn(float64(c)*lx/float64(cols), float64(r)*ly/float64(rows)) + } + } + return mustFloats(t, vals, rows, cols) +} + +// poissonError returns the largest absolute difference between the +// solution and an exact function of (x, y) on the same grid. +func poissonError(u *core.Array, want func(x, y float64) float64, rows, cols int, lx, ly float64) float64 { + worst := 0.0 + for r := range rows { + for c := range cols { + d := math.Abs(u.FloatAt(r*cols+c) - want(float64(c)*lx/float64(cols), float64(r)*ly/float64(rows))) + worst = math.Max(worst, d) + } + } + return worst +} + +// TestSolvePoissonPeriodicModes checks the plane-wave answers the +// diagonalisation gets exactly: for u = sin x·sin y the Laplacian +// gives −2u, and for u = sin 3x·cos 2y it gives −13u. Both land on +// round-off. +func TestSolvePoissonPeriodicModes(t *testing.T) { + const lx, ly = 2 * math.Pi, 2 * math.Pi + cases := []struct { + f, u func(x, y float64) float64 + }{ + { + f: func(x, y float64) float64 { return 2 * math.Sin(x) * math.Sin(y) }, + u: func(x, y float64) float64 { return math.Sin(x) * math.Sin(y) }, + }, + { + f: func(x, y float64) float64 { return 13 * math.Sin(3*x) * math.Cos(2*y) }, + u: func(x, y float64) float64 { return math.Sin(3*x) * math.Cos(2*y) }, + }, + } + for k, tc := range cases { + grid := poissonGrid(t, tc.f, 16, 16, lx, ly) + u, err := SolvePoissonPeriodic(grid, lx, ly) + if err != nil { + t.Fatalf("case %d: SolvePoissonPeriodic: %v", k, err) + } + if u.Dtype() == core.Complex { + t.Fatalf("case %d: the solution must be real", k) + } + if got := poissonError(u, tc.u, 16, 16, lx, ly); got > 1e-10 { + t.Fatalf("case %d: max error %.3g, want round-off", k, got) + } + } +} + +// TestSolvePoissonPeriodicSpectral shows the convergence is spectral: +// for u = sin x·e^{sin y}, whose Fourier coefficients decay faster +// than any power, refining the grid drives the error from visible to +// round-off in one refinement. +func TestSolvePoissonPeriodicSpectral(t *testing.T) { + const lx, ly = 2 * math.Pi, 2 * math.Pi + f := func(x, y float64) float64 { + return math.Sin(x) * math.Exp(math.Sin(y)) * (math.Sin(y)*math.Sin(y) + math.Sin(y)) + } + u := func(x, y float64) float64 { return math.Sin(x) * math.Exp(math.Sin(y)) } + errAt := func(n int) float64 { + sol, err := SolvePoissonPeriodic(poissonGrid(t, f, n, n, lx, ly), lx, ly) + if err != nil { + t.Fatalf("SolvePoissonPeriodic(%d): %v", n, err) + } + return poissonError(sol, u, n, n, lx, ly) + } + coarse, fine, finer := errAt(16), errAt(32), errAt(64) + if coarse < 1e-9 { + t.Skipf("coarse grid already at round-off (%v)", coarse) + } + if fine >= coarse { + t.Fatalf("refining the grid did not help: %.3g then %.3g", coarse, fine) + } + if fine > 1e-11 || finer > 1e-11 { + t.Fatalf("spectral accuracy not reached: %.3g, %.3g", fine, finer) + } +} + +// TestSolvePoissonPeriodicErrors pins the validation contract: the +// zero-mean compatibility condition, the shape and size of the grid, +// and positive side lengths. +func TestSolvePoissonPeriodicErrors(t *testing.T) { + const lx, ly = 2 * math.Pi, 2 * math.Pi + constant := poissonGrid(t, func(x, y float64) float64 { return 1 }, 8, 8, lx, ly) + if _, err := SolvePoissonPeriodic(constant, lx, ly); err == nil { + t.Fatal("expected an error for a right-hand side with nonzero mean") + } + rank1 := mustFloats(t, []float64{0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8}, 8) + if _, err := SolvePoissonPeriodic(rank1, lx, ly); err == nil { + t.Fatal("expected an error for a rank-1 right-hand side") + } + small := poissonGrid(t, func(x, y float64) float64 { return x }, 1, 8, lx, ly) + if _, err := SolvePoissonPeriodic(small, lx, ly); err == nil { + t.Fatal("expected an error for a grid dimension below 2") + } + grid := poissonGrid(t, func(x, y float64) float64 { return math.Sin(x) }, 8, 8, lx, ly) + if _, err := SolvePoissonPeriodic(grid, 0, ly); err == nil { + t.Fatal("expected an error for a zero side length") + } +} + +// TestSolvePoissonPeriodicDtypeContract pins the dtype contract: the +// solver computes in float64 end to end, so a float32, int or complex +// right-hand side is an error rather than a silently all-zero +// solution (RawFloats is nil for those dtypes). +func TestSolvePoissonPeriodicDtypeContract(t *testing.T) { + const lx, ly = 2 * math.Pi, 2 * math.Pi + f32, _ := core.FromFloat32s([]float32{1, 2, 3, 4, 5, 6, 7, 8, 8, 7, 6, 5, 4, 3, 2, 1}, 4, 4) + if _, err := SolvePoissonPeriodic(f32, lx, ly); err == nil { + t.Fatal("expected an error for a float32 right-hand side") + } + ints, _ := core.FromInts([]int64{1, 2, 3, 4, 5, 6, 7, 8, 8, 7, 6, 5, 4, 3, 2, 1}, 4, 4) + if _, err := SolvePoissonPeriodic(ints, lx, ly); err == nil { + t.Fatal("expected an error for an int right-hand side") + } + cplx, _ := core.FromComplexes(make([]complex128, 16), 4, 4) + if _, err := SolvePoissonPeriodic(cplx, lx, ly); err == nil { + t.Fatal("expected an error for a complex right-hand side") + } +} diff --git a/signal/poisson_view_pin_test.go b/signal/poisson_view_pin_test.go new file mode 100644 index 0000000..753ce01 --- /dev/null +++ b/signal/poisson_view_pin_test.go @@ -0,0 +1,99 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The spectral Poisson solves read their source through poissonFloats, +// which must materialise a rebased view in element order. A view's +// payload is longer than its element count and starts past offset zero, +// so a raw payload read would solve a different problem: this pin +// compares a 2-D view answer against the same values in a fresh array, +// bit for bit, on both solvers. +func TestSolvePoissonViewsMatchFresh(t *testing.T) { + const rows, cols = 9, 9 + win := make([]float64, rows*cols) + for i := range win { + win[i] = math.Sin(float64(i)*0.37)*0.5 + float64(i%7)*0.1 - 0.3 + } + // Neumann's compatibility gate measures the source with the + // trapezoidal rule: interior weight 1, edge weight 1/2, corner + // weight 1/4. Subtract the trapezoidal mean so the window's measure + // is exactly zero. + weights := make([]float64, rows*cols) + wSum := 0.0 + for r := range rows { + for c := range cols { + w := 1.0 + if r == 0 || r == rows-1 { + w /= 2 + } + if c == 0 || c == cols-1 { + w /= 2 + } + weights[r*cols+c] = w + wSum += w + } + } + mean := 0.0 + for i := range win { + mean += weights[i] * win[i] + } + mean /= wSum + for i := range win { + win[i] -= mean + } + padded := make([]float64, (rows+2)*(cols+2)) + for r := range rows { + copy(padded[(r+1)*(cols+2)+1:(r+1)*(cols+2)+1+cols], win[r*cols:(r+1)*cols]) + } + back, err := core.FromFloats(padded, rows+2, cols+2) + if err != nil { + t.Fatal(err) + } + v1, err := core.Slice(back, 0, 1, rows+1) + if err != nil { + t.Fatal(err) + } + view, err := core.Slice(v1, 1, 1, cols+1) + if err != nil { + t.Fatal(err) + } + fresh, err := core.FromFloats(win, rows, cols) + if err != nil { + t.Fatal(err) + } + solvers := map[string]func(*core.Array) (*core.Array, error){ + "SolvePoissonDirichlet": func(a *core.Array) (*core.Array, error) { + return SolvePoissonDirichlet(a, 1, 1) + }, + "SolvePoissonNeumann": func(a *core.Array) (*core.Array, error) { + return SolvePoissonNeumann(a, 1, 1) + }, + } + for name, solve := range solvers { + gv, err := solve(view) + if err != nil { + t.Fatalf("%s(view): %v", name, err) + } + gf, err := solve(fresh) + if err != nil { + t.Fatalf("%s(fresh): %v", name, err) + } + if gv.Len() != gf.Len() { + t.Fatalf("%s: view length %d, fresh %d", name, gv.Len(), gf.Len()) + } + for i := range gv.Len() { + if math.Float64bits(gv.FloatAt(i)) != math.Float64bits(gf.FloatAt(i)) { + t.Fatalf("%s element %d: view %#x, fresh %#x", + name, i, math.Float64bits(gv.FloatAt(i)), math.Float64bits(gf.FloatAt(i))) + } + } + } +} diff --git a/signal/poissondirichlet.go b/signal/poissondirichlet.go new file mode 100644 index 0000000..e224653 --- /dev/null +++ b/signal/poissondirichlet.go @@ -0,0 +1,468 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Spectral solution of the Poisson equation on a rectangle with +// homogeneous boundary data. The sine basis vanishes on the boundary +// and diagonalises −Δ for Dirichlet data; the cosine basis has zero +// slope there and does the same for Neumann. Both solves run one +// separable transform pair per axis and one division per mode, and +// both refuse the inputs the mathematics refuses: nonzero-mean +// sources for Neumann, exactly as the periodic solve does. + +// poissonRect validates the shared rectangle arguments and returns +// the grid shape and spacings. +func poissonRect(name string, f *core.Array, lx, ly float64) (rows, cols int, hx, hy float64, err error) { + if f.NDim() != 2 { + return 0, 0, 0, 0, base.Errf("%s: f must be rank 2, got shape %s", name, base.ShapeText(f.Shape())) + } + if f.Dtype() != core.Float { + return 0, 0, 0, 0, base.Errf("%s: f must be a float64 array, got %s", name, f.Dtype()) + } + rows, cols = f.Shape()[0], f.Shape()[1] + if rows < 3 || cols < 3 { + return 0, 0, 0, 0, base.Errf("%s: the grid must be at least 3×3 to hold interior points, got %d×%d", name, rows, cols) + } + if lx <= 0 || ly <= 0 { + return 0, 0, 0, 0, base.Errf("%s: the side lengths must be positive, got %g and %g", name, lx, ly) + } + // A non-finite source would drive the compatibility tolerances NaN + // and slip through their gates, so it is refused up front. A dense + // float64 payload is scanned slice-wise: the elements are the ones + // FloatAt returns. + vals := poissonFloats(f) + for _, v := range vals { + if math.IsInf(v, 0) || math.IsNaN(v) { + return 0, 0, 0, 0, base.Errf("%s: f holds the non-finite value %g", name, v) + } + } + hx = lx / float64(cols-1) + hy = ly / float64(rows-1) + return rows, cols, hx, hy, nil +} + +// poissonFloats streams the grid's elements in element order: the raw +// payload when the array is dense, a materialised copy through FloatAt +// when it is a strided view, whose payload order is not the element +// order. Treat an aliased result as read-only. +func poissonFloats(f *core.Array) []float64 { + if !f.Strided() { + return f.RawFloats() + } + out := make([]float64, f.Len()) + for i := range out { + out[i] = f.FloatAt(i) + } + return out +} + +// poissonTrapzSum returns the trapezoidal-weighted sum of f over the +// whole grid: the boundary rows and columns count half, the weight the +// mirrored stencil's own left null vector carries. It is the constant +// mode of the cosine transform pair, the measure the Neumann solve is +// compatible in. +func poissonTrapzSum(f *core.Array) float64 { + rows, cols := f.Shape()[0], f.Shape()[1] + vals := poissonFloats(f) + sum := 0.0 + for r := range rows { + wr := 1.0 + if r == 0 || r == rows-1 { + wr = 0.5 + } + for c := range cols { + wc := 1.0 + if c == 0 || c == cols-1 { + wc = 0.5 + } + sum += wr * wc * vals[r*cols+c] + } + } + return sum +} + +// SolvePoissonDirichlet solves −Δu = f on the rectangle with u held +// at zero on the whole boundary. f is sampled on the (rows × cols) +// grid including the boundary rows and columns, whose entries play no +// role in the solve; u comes back on the same grid with a zero +// boundary. Row r of f samples y = r·hy, column c samples x = c·hx, +// with hx = lx/(cols−1) and hy = ly/(rows−1). The sine transform +// pair along each axis diagonalises the interior second-difference +// operator, so the error is the stencil's O(h²), decaying with the +// grid. Non-float dtypes, grids under 3×3 and non-positive side +// lengths are errors. +func SolvePoissonDirichlet(f *core.Array, lx, ly float64) (*core.Array, error) { + const name = "SolvePoissonDirichlet" + rows, cols, hx, hy, err := poissonRect(name, f, lx, ly) + if err != nil { + return nil, err + } + interiorR, interiorC := rows-2, cols-2 + // The interior right-hand side, transformed along both axes with + // the orthonormal DST-I, whose basis vanishes on the boundary. + spectrum, err := transformBlock(f, 1, 1, interiorR, interiorC, true) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + // Divide by the stencil eigenvalues: two per axis from the cosine + // identity on the second difference. Each eigenvalue depends on one + // index alone, so the per-axis values are built once instead of once + // per mode; the expressions are the ones the inner loop evaluated. + eigX := make([]float64, interiorC+1) + for kx := 1; kx <= interiorC; kx++ { + eigX[kx] = 2 / (hx * hx) * (1 - math.Cos(math.Pi*float64(kx)/float64(interiorC+1))) + } + eigY := make([]float64, interiorR+1) + for ky := 1; ky <= interiorR; ky++ { + eigY[ky] = 2 / (hy * hy) * (1 - math.Cos(math.Pi*float64(ky)/float64(interiorR+1))) + } + for ky := 1; ky <= interiorR; ky++ { + ly := eigY[ky] + row := (ky - 1) * interiorC + for kx := 1; kx <= interiorC; kx++ { + spectrum[row+kx-1] /= eigX[kx] + ly + } + } + interior, err := inverseBlock(spectrum, interiorR, interiorC, true) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + out := core.New(core.Float, rows, cols) + into := out.RawFloats() + for r := range interiorR { + for c := range interiorC { + into[(r+1)*cols+(c+1)] = interior[r*interiorC+c] + } + } + return out, nil +} + +// SolvePoissonNeumann solves −Δu = f with zero normal derivative on +// the whole boundary: u comes back on the full grid, boundary points +// included, every sample an unknown. Conventions mirror +// SolvePoissonDirichlet, with the cosine transform pair along each +// axis. The stencil is the ghost-mirror one, the neighbour outside the +// boundary standing in for zero slope, so a boundary row couples its +// mirrored neighbour with 2/h² where an interior row uses 1/h² (the +// boundary control volume is half an interior one). Its cosine modes +// cos(πkx x/lx)·cos(πky y/ly), kx = 0…cols−1 and ky = 0…rows−1, carry +// the eigenvalues 2/hx²·(1−cos(πkx/(cols−1))) + +// 2/hy²·(1−cos(πky/(rows−1))) per axis, and the error against the +// closed form is the stencil's O(h²). +// +// The operator annihilates the constants, so a solution exists only +// for sources whose trapezoidal-weighted sum vanishes; that sum is the +// operator's own compatibility condition (the trapezoidal measure is +// its left null vector), not the plain mean, and a nonzero one is an +// error rather than a quietly shifted problem. The free constant of u +// is fixed by setting the constant mode to zero, so the result has +// zero trapezoidal mean. Non-float dtypes, grids under 3×3 and +// non-positive side lengths are errors. +func SolvePoissonNeumann(f *core.Array, lx, ly float64) (*core.Array, error) { + const name = "SolvePoissonNeumann" + rows, cols, hx, hy, err := poissonRect(name, f, lx, ly) + if err != nil { + return nil, err + } + // The compatibility gate: the constant mode of the transformed + // source is the trapezoidal-weighted sum, so anything above the + // accumulated round-off of that sum breaks solvability. The + // tolerance follows the periodic solve's constant-mode gate, and it + // is relative to the source: the sum of a trapezoidal grid scales + // with the sample magnitude and the element count, so an absolute + // floor would let a tiny source through with a residual far above + // its own size (a 1e-9 source used to be accepted with a 1.8e-5 + // relative residual). + sum := poissonTrapzSum(f) + tol := 64 * base.EpsF * float64(rows*cols) * maxAbsF(f) + if math.Abs(sum) > tol { + return nil, base.Errf("%s: f has the nonzero trapezoidal mean %g, the Neumann problem has no solution", + name, sum/float64((rows-1)*(cols-1))) + } + // The pipeline. The mirrored operator M (rows [2,−2] on the + // boundary, [−1,2,−1] inside, over h²) is diagonalised by the + // UNNORMALISED DCT-I, whose boundary input weight of 1/2 is the + // trapezoidal measure; the library carries the orthonormal DCT-I + // (boundary weight 1/sqrt2), and W = diag(1/2,1,…,1,1/2) turns one + // into the other. So the forward trip scales by W^(1/2) before the + // DCT-I, divides by the eigenvalues, and the return trip scales by + // W^(-1/2) after the second DCT-I, which is its own inverse. + weighted := core.New(core.Float, rows, cols) + src := poissonFloats(f) + dst := weighted.RawFloats() + for r := range rows { + sr := 1.0 + if r == 0 || r == rows-1 { + sr = math.Sqrt2 / 2 + } + for c := range cols { + sc := 1.0 + if c == 0 || c == cols-1 { + sc = math.Sqrt2 / 2 + } + dst[r*cols+c] = src[r*cols+c] * sr * sc + } + } + spectrum, err := transformBlock(weighted, 0, 0, rows, cols, false) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + // The x eigenvalues depend on kx alone, so the row is built once and + // read by every ky; the expression is the one the inner loop + // evaluated. The y eigenvalue was already hoisted to the outer loop. + eigX := make([]float64, cols) + for kx := range cols { + eigX[kx] = 2 / (hx * hx) * (1 - math.Cos(math.Pi*float64(kx)/float64(cols-1))) + } + for ky := range rows { + ly := 2 / (hy * hy) * (1 - math.Cos(math.Pi*float64(ky)/float64(rows-1))) + row := ky * cols + for kx := range cols { + i := row + kx + if kx == 0 && ky == 0 { + // The constant mode is the operator's null space: it + // stays zero, which is what fixes the solution's free + // constant, instead of dividing by its zero eigenvalue. + spectrum[i] = 0 + continue + } + spectrum[i] /= eigX[kx] + ly + } + } + interior, err := inverseBlock(spectrum, rows, cols, false) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + // The return trip: undo the W^(1/2) scaling, and strip the + // trapezoidal mean, the measure the solve works in and the one the + // operator leaves free. + out := core.New(core.Float, rows, cols) + into := out.RawFloats() + sum, weight := 0.0, 0.0 + for r := range rows { + sr, wr := 1.0, 1.0 + if r == 0 || r == rows-1 { + sr, wr = math.Sqrt2, 0.5 + } + for c := range cols { + sc, wc := 1.0, 1.0 + if c == 0 || c == cols-1 { + sc, wc = math.Sqrt2, 0.5 + } + v := interior[r*cols+c] * sr * sc + into[r*cols+c] = v + sum += wr * wc * v + weight += wr * wc + } + } + mean := sum / weight + for i := range into { + into[i] -= mean + } + return out, nil +} + +// maxAbsF returns the largest absolute sample of a float array. +func maxAbsF(f *core.Array) float64 { + largest := 0.0 + for _, v := range poissonFloats(f) { + largest = max(largest, math.Abs(v)) + } + return largest +} + +// lineTransformPlan carries the chirp-z pieces one transform length +// needs for the orthonormal DST-I (sine) or DCT-I (cosine) the Poisson +// solves apply along a line, plus the buffers the per-line work reuses. +// +// Both transforms are the real or imaginary part of a sum of complex +// exponentials over the half-frequency grid the trigonometric identity +// (j+1)(k+1) = ((j+1)² + (k+1)² − (j−k)²)/2 exposes: a per-sample +// phase, a per-bin rotation and an even chirp kernel, so the sum over +// the bins is a circular convolution of length the smallest power of +// two ≥ 2n−1. The padded full-length route each line used to take ran +// a convolution of roughly twice that length to reach the same bins. +type lineTransformPlan struct { + n int + m int + sine bool + // pre weights the samples, post rotates the bins, kern is the + // kernel's forward transform on the convolution grid. + pre []complex128 + post []complex128 + kern []complex128 + norm float64 + // buf carries one line's spectrum through the convolution; the plan + // is used from one goroutine and the solves are single-threaded, so + // the buffer is fully rewritten on every apply. + buf []complex128 +} + +// newLineTransformPlan builds the plan for n points. The chirp angles +// are reduced modulo the kernel's period in exact integer arithmetic +// before the transcendental call, so no angle grows with n. +func newLineTransformPlan(n int, sine bool) *lineTransformPlan { + p := &lineTransformPlan{n: n, sine: sine} + m := 1 + for m < 2*n-1 { + m <<= 1 + } + p.m = m + p.pre = make([]complex128, n) + p.post = make([]complex128, n) + p.kern = make([]complex128, m) + p.buf = make([]complex128, m) + if sine { + // DST-I: out[k] = norm·Im[ post_k · Σ_j x_j·pre_j·e^{−iθ(j−k)²/2} ] + // with θ = π/(n+1), pre_j = e^{iθ(j + j²/2)}, + // post_k = e^{iθ(k+1 + k²/2)}. + period := 4 * (n + 1) + theta := math.Pi / float64(n+1) + for j := range n { + t := (2*j + j*j) % period + p.pre[j] = polar(1, theta*float64(t)/2) + } + for k := range n { + t := (2*(k+1) + k*k) % period + p.post[k] = polar(1, theta*float64(t)/2) + } + for l := range n { + t := (l * l) % period + w := polar(1, -theta*float64(t)/2) + p.kern[l] = w + if l != 0 { + p.kern[m-l] = w + } + } + p.norm = math.Sqrt(2 / float64(n+1)) + } else { + // DCT-I: out[k] = norm·Re[ w_k·e^{iθk²/2} · + // Σ_j x_j·w_j·e^{iθj²/2}·e^{−iθ(j−k)²/2} ] with θ = π/(n−1) + // and w the half weight on the endpoints. + period := 4 * (n - 1) + theta := math.Pi / float64(n-1) + for j := range n { + t := (j * j) % period + w := 1.0 + if j == 0 || j == n-1 { + w = math.Sqrt2 / 2 + } + p.pre[j] = complex(w, 0) * polar(1, theta*float64(t)/2) + } + for k := range n { + t := (k * k) % period + w := 1.0 + if k == 0 || k == n-1 { + w = math.Sqrt2 / 2 + } + p.post[k] = complex(w, 0) * polar(1, theta*float64(t)/2) + } + for l := range n { + t := (l * l) % period + w := polar(1, -theta*float64(t)/2) + p.kern[l] = w + if l != 0 { + p.kern[m-l] = w + } + } + p.norm = math.Sqrt(2 / float64(n-1)) + } + // The kernel's spectrum depends only on the plan; it is the same + // forward transform every line multiplies against. + transform(p.kern, -1) + return p +} + +// polar returns the unit complex number of the given angle. +// math.Sincos returns exactly the pair math.Cos and math.Sin produce +// (verified bit-for-bit), so the value is unchanged. +func polar(magnitude, angle float64) complex128 { + s, c := math.Sincos(angle) + return complex(magnitude*c, magnitude*s) +} + +// apply transforms src into dst; both have length p.n and must not +// alias. The convolution runs on the plan's reused spectrum buffer, +// which every call overwrites completely before it is read. +func (p *lineTransformPlan) apply(dst, src []float64) { + n, m := p.n, p.m + buf := p.buf + if p.sine && n == 1 { + dst[0] = src[0] + return + } + for j := range n { + buf[j] = complex(src[j], 0) * p.pre[j] + } + clear(buf[n:m]) + transform(buf, -1) + for i := range m { + buf[i] *= p.kern[i] + } + transform(buf, +1) + scale := complex(1/float64(m), 0) + if p.sine { + for k := range n { + v := buf[k] * scale * p.post[k] + dst[k] = p.norm * imag(v) + } + return + } + for k := range n { + v := buf[k] * scale * p.post[k] + dst[k] = p.norm * real(v) + } +} + +// transformBlock applies the named self-inverse orthonormal transform +// (DST-I when sine, DCT-I when not) to every row and then every column +// of the rows×cols block of f whose top-left corner sits at (row0, +// col0), returning the flat spectrum. One plan per axis length serves +// every line of that axis. +func transformBlock(f *core.Array, row0, col0, rows, cols int, sine bool) ([]float64, error) { + srcVals := poissonFloats(f) + srcCols := f.Shape()[1] + work := make([]float64, rows*cols) + longest := max(rows, cols) + line := make([]float64, longest) + out := make([]float64, longest) + // poissonFloats materialises a strided source in element order; a + // raw payload read would keep the physical order instead. + rowPlan := newLineTransformPlan(cols, sine) + for r := range rows { + for c := range cols { + line[c] = srcVals[(r+row0)*srcCols+(c+col0)] + } + rowPlan.apply(out[:cols], line[:cols]) + copy(work[r*cols:(r+1)*cols], out[:cols]) + } + colPlan := newLineTransformPlan(rows, sine) + for c := range cols { + for r := range rows { + line[r] = work[r*cols+c] + } + colPlan.apply(out[:rows], line[:rows]) + for r := range rows { + work[r*cols+c] = out[r] + } + } + return work, nil +} + +// inverseBlock applies the same self-inverse transform to a flat +// spectrum, the return trip of transformBlock. +func inverseBlock(spectrum []float64, rows, cols int, sine bool) ([]float64, error) { + arr, err := core.FromFloats(spectrum, rows, cols) + if err != nil { + return nil, err + } + return transformBlock(arr, 0, 0, rows, cols, sine) +} diff --git a/signal/poissondirichlet_test.go b/signal/poissondirichlet_test.go new file mode 100644 index 0000000..558c5ea --- /dev/null +++ b/signal/poissondirichlet_test.go @@ -0,0 +1,130 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// poissonManufactured builds f = −Δu for u = sin(πx)·sin(πy) on the +// unit square over an (n × n) grid including the boundary, so the +// exact solution is known everywhere. +func poissonManufactured(t *testing.T, n int) (f, exact *core.Array) { + t.Helper() + flatF := make([]float64, n*n) + flatU := make([]float64, n*n) + for r := range n { + for c := range n { + x := float64(c) / float64(n-1) + y := float64(r) / float64(n-1) + i := r*n + c + flatU[i] = math.Sin(math.Pi*x) * math.Sin(math.Pi*y) + flatF[i] = 2 * math.Pi * math.Pi * flatU[i] + } + } + f = mustFloats(t, flatF, n, n) + exact = mustFloats(t, flatU, n, n) + return f, exact +} + +// TestSolvePoissonDirichletManufactured checks the sine-solve against +// the manufactured solution: second-order convergence, with the +// coarse grid already inside 5e-3. +func TestSolvePoissonDirichletManufactured(t *testing.T) { + for _, n := range []int{9, 17, 33} { + f, exact := poissonManufactured(t, n) + u, err := SolvePoissonDirichlet(f, 1, 1) + if err != nil { + t.Fatalf("SolvePoissonDirichlet(%d): %v", n, err) + } + worst := 0.0 + for i := range exact.Len() { + if e := math.Abs(u.FloatAt(i) - exact.FloatAt(i)); e > worst { + worst = e + } + } + h := 1 / float64(n-1) + if worst > 3*h*h { + t.Fatalf("grid %d: worst error %.3g above the O(h²) budget %.3g", n, worst, 3*h*h) + } + // The boundary is exactly zero by construction. + for c := range n { + if u.FloatAt(c) != 0 || u.FloatAt((n-1)*n+c) != 0 { + t.Fatalf("grid %d: boundary came back nonzero", n) + } + } + } +} + +// TestSolvePoissonNeumannManufactured checks the cosine solve with a +// zero-mean source whose solution is u = cos(πx)·cos(πy), the +// derivative-free data the Neumann solve exists for. +func TestSolvePoissonNeumannManufactured(t *testing.T) { + const n = 17 + flatF := make([]float64, n*n) + flatU := make([]float64, n*n) + for r := range n { + for c := range n { + x := float64(c) / float64(n-1) + y := float64(r) / float64(n-1) + i := r*n + c + flatU[i] = math.Cos(math.Pi*x) * math.Cos(math.Pi*y) + flatF[i] = 2 * math.Pi * math.Pi * flatU[i] + } + } + // Zero-mean shift: subtract the mean, which also shifts u by a + // constant the Neumann problem cannot see. + meanF := 0.0 + for _, v := range flatF { + meanF += v + } + meanF /= float64(n * n) + for i := range flatF { + flatF[i] -= meanF + } + f := mustFloats(t, flatF, n, n) + exact := mustFloats(t, flatU, n, n) + u, err := SolvePoissonNeumann(f, 1, 1) + if err != nil { + t.Fatalf("SolvePoissonNeumann: %v", err) + } + // Compare up to the free constant: recentre both to zero mean. + got := 0.0 + for i := range u.Len() { + got += u.FloatAt(i) + } + got /= float64(u.Len()) + worst := 0.0 + for i := range u.Len() { + if e := math.Abs(u.FloatAt(i) - got - exact.FloatAt(i)); e > worst { + worst = e + } + } + h := 1 / float64(n-1) + if worst > 3*h*h { + t.Fatalf("worst error %.3g above the O(h²) budget %.3g", worst, 3*h*h) + } +} + +// TestSolvePoissonNeumannMeanRefusal checks the compatibility +// condition: a nonzero-mean source has no Neumann solution. +func TestSolvePoissonNeumannMeanRefusal(t *testing.T) { + f := mustFloats(t, []float64{ + 1, 1, 1, + 1, 1, 1, + 1, 1, 1, + }, 3, 3) + if _, err := SolvePoissonNeumann(f, 1, 1); err == nil { + t.Fatal("nonzero-mean source accepted") + } + if _, err := SolvePoissonDirichlet(mustFloats(t, []float64{1, 1}, 1, 2), 1, 1); err == nil { + t.Fatal("grid under 3×3 accepted") + } + if _, err := SolvePoissonDirichlet(f, 0, 1); err == nil { + t.Fatal("non-positive length accepted") + } +} diff --git a/signal/pool.go b/signal/pool.go new file mode 100644 index 0000000..7a2d3a1 --- /dev/null +++ b/signal/pool.go @@ -0,0 +1,815 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Pooling operations for 2-D feature maps. All functions take +// NCHW input (N, C, H, W) and return (N, C, H_out, W_out). Padding is +// a single int applied to both spatial dims; stride is a single int. +// Padding must stay below the kernel: a window entirely inside the +// padding holds no data, and a max would answer −Inf there, an +// average 0. +// +// The kernels hold two constraints fixed: +// +// - Element order. Average pooling adds its window elements in the +// (kernel height, kernel width) nesting and divides once, with the +// countIncludePad divisor chosen exactly as before; any other +// order or divisor changes bits. Max pooling is order free but +// keeps the same walk for a single shared code path, and a NaN in +// a window wins: v > best is false for NaN, so the comparison +// would silently drop it and answer the largest finite neighbour +// instead. +// - Parallel split. Work is distributed over output rows, which own +// disjoint output elements, so no split can influence a result. +// +// Non-float inputs widen once into a scratch payload (see +// widenFloats): FloatAt widens exactly, so the values +// compared and summed are bit-identical to the per-element accessor +// reads. + +// parallelWorkBudget is the number of element updates a worker should +// carry before splitting a kernel pays for the worker's spawn. Every +// scheduling floor in the package is derived from it: an item body +// costing a few hundred nanoseconds still leaves a worker below the +// budget on a tiny input, where the split would spawn goroutines to do +// less work than the spawn itself costs. +const parallelWorkBudget = 1 << 12 + +// workFloorFor returns the smallest number of items a worker should +// carry when one item costs about workPerItem element updates. It is +// the floor argument engine.ParallelMin wants: below it the kernel runs +// whole on the calling goroutine, and above it the split is the one +// engine.Parallel computes, so the worker choice and the chunk +// boundaries of a split kernel are unchanged. +func workFloorFor(workPerItem int) int { + if workPerItem < 1 { + return parallelWorkBudget + } + return max(1, parallelWorkBudget/workPerItem) +} + +// MaxPool2D returns the maximum over each kernel-sized window. A NaN +// in a window propagates: the maximum answers NaN. +func MaxPool2D(input *core.Array, kernel, stride, padding int) (*core.Array, error) { + return pool2D(input, kernel, stride, padding, true, false) +} + +// AvgPool2D returns the average over each kernel-sized window. When +// countIncludePad is true the divisor is kH*kW; otherwise it's the +// number of non-padded elements in the window. +func AvgPool2D(input *core.Array, kernel, stride, padding int, countIncludePad bool) (*core.Array, error) { + return pool2D(input, kernel, stride, padding, false, countIncludePad) +} + +// MaxPool1D returns the maximum over each kernel-sized window of a +// 3-D tensor (N, C, L). A NaN in a window propagates: the maximum +// answers NaN. +func MaxPool1D(input *core.Array, kernel, stride, padding int) (*core.Array, error) { + return pool1D(input, kernel, stride, padding, true, false) +} + +// AvgPool1D returns the average over each kernel-sized window of a +// 3-D tensor (N, C, L). When countIncludePad is true the divisor is +// the full kernel size; otherwise it is the number of non-padded +// elements in the window. +func AvgPool1D(input *core.Array, kernel, stride, padding int, countIncludePad bool) (*core.Array, error) { + return pool1D(input, kernel, stride, padding, false, countIncludePad) +} + +// MaxPool3D returns the maximum over each kernel-sized window of a +// 5-D tensor (N, C, D, H, W). A NaN in a window propagates: the +// maximum answers NaN. +func MaxPool3D(input *core.Array, kernel [3]int, stride [3]int, padding [3]int) (*core.Array, error) { + return pool3D(input, kernel, stride, padding, true, false) +} + +// AvgPool3D returns the average over each kernel-sized window of a +// 5-D tensor (N, C, D, H, W). When countIncludePad is true the +// divisor is the full kernel size; otherwise it is the number of +// non-padded elements in the window. +func AvgPool3D(input *core.Array, kernel [3]int, stride [3]int, padding [3]int, countIncludePad bool) (*core.Array, error) { + return pool3D(input, kernel, stride, padding, false, countIncludePad) +} + +// AdaptiveMaxPool2D pools to the requested output size by taking the +// maximum over each output window. Window o covers the input rows +// from floor(o·H_in/H_out) to ceil((o+1)·H_in/H_out), so every input +// sample lands in some window. A NaN in a window propagates: the +// maximum answers NaN. +func AdaptiveMaxPool2D(input *core.Array, outputH, outputW int) (*core.Array, error) { + if input.NDim() != 4 { + return nil, base.Errf("AdaptiveMaxPool2D: input must be 4-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("AdaptiveMaxPool2D: complex arrays are not supported") + } + if outputH < 1 || outputW < 1 { + return nil, base.Errf("AdaptiveMaxPool2D: output size must be at least 1") + } + n, c, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3] + // A zero spatial dimension leaves every window empty: max would + // answer -Inf, so it is an error like every other shape refusal. + if hIn < 1 || wIn < 1 { + return nil, base.Errf("AdaptiveMaxPool2D: the spatial dimensions must not be empty, got shape %s", base.ShapeText(input.Shape())) + } + out := core.New(core.Float, []int{n, c, outputH, outputW}...) + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per output row; rows are disjoint. A row walks + // about one window height of input rows per window column. + items := n * c * outputH + floor := workFloorFor(max(1, (hIn+outputH-1)/outputH) * wIn) + engine.ParallelMin(items, floor, func(bs, be int) { + for item := bs; item < be; item++ { + oy := item % outputH + row := item / outputH + ch := row % c + batch := row / c + chanBase := batch*c*hIn*wIn + ch*hIn*wIn + yStart := oy * hIn / outputH + yEnd := ((oy+1)*hIn + outputH - 1) / outputH + for ox := range outputW { + xStart := ox * wIn / outputW + xEnd := ((ox+1)*wIn + outputW - 1) / outputW + best := math.Inf(-1) + for y := yStart; y < yEnd; y++ { + inRow := chanBase + y*wIn + for x := xStart; x < xEnd; x++ { + if v := inF[inRow+x]; v > best || math.IsNaN(v) { + best = v + } + } + } + outF[batch*c*outputH*outputW+ch*outputH*outputW+oy*outputW+ox] = best + } + } + }) + return out, nil +} + +// GlobalAvgPool2D returns AdaptiveAvgPool2D reduced to (1, 1), the +// global average pooling operation common in modern CNNs. +func GlobalAvgPool2D(input *core.Array) (*core.Array, error) { + return AdaptiveAvgPool2D(input, 1, 1) +} + +// GlobalMaxPool2D returns the global maximum of a 4-D feature map. +func GlobalMaxPool2D(input *core.Array) (*core.Array, error) { + return AdaptiveMaxPool2D(input, 1, 1) +} + +// GlobalAvgPool1D reduces a 3-D input (N, C, L) to shape (N, C, 1). +func GlobalAvgPool1D(input *core.Array) (*core.Array, error) { + if input.NDim() != 3 { + return nil, base.Errf("GlobalAvgPool1D: input must be 3-D (N, C, L), got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("GlobalAvgPool1D: complex arrays are not supported") + } + if input.Shape()[2] == 0 { + return nil, base.Errf("GlobalAvgPool1D: the spatial dimension must not be empty") + } + return pool1D(input, input.Shape()[2], 1, 0, false, false) +} + +// GlobalMaxPool1D reduces a 3-D input (N, C, L) to shape (N, C, 1). +func GlobalMaxPool1D(input *core.Array) (*core.Array, error) { + if input.NDim() != 3 { + return nil, base.Errf("GlobalMaxPool1D: input must be 3-D (N, C, L), got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("GlobalMaxPool1D: complex arrays are not supported") + } + if input.Shape()[2] == 0 { + return nil, base.Errf("GlobalMaxPool1D: the spatial dimension must not be empty") + } + return pool1D(input, input.Shape()[2], 1, 0, true, false) +} + +// GlobalAvgPool3D reduces a 5-D input (N, C, D, H, W) to (N, C, 1, 1, 1). +func GlobalAvgPool3D(input *core.Array) (*core.Array, error) { + if input.NDim() != 5 { + return nil, base.Errf("GlobalAvgPool3D: input must be 5-D (N, C, D, H, W), got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("GlobalAvgPool3D: complex arrays are not supported") + } + if input.Shape()[2] == 0 || input.Shape()[3] == 0 || input.Shape()[4] == 0 { + return nil, base.Errf("GlobalAvgPool3D: the spatial dimensions must not be empty") + } + return pool3D(input, [3]int{input.Shape()[2], input.Shape()[3], input.Shape()[4]}, [3]int{1, 1, 1}, [3]int{0, 0, 0}, false, false) +} + +// GlobalMaxPool3D reduces a 5-D input (N, C, D, H, W) to (N, C, 1, 1, 1). +func GlobalMaxPool3D(input *core.Array) (*core.Array, error) { + if input.NDim() != 5 { + return nil, base.Errf("GlobalMaxPool3D: input must be 5-D (N, C, D, H, W), got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("GlobalMaxPool3D: complex arrays are not supported") + } + if input.Shape()[2] == 0 || input.Shape()[3] == 0 || input.Shape()[4] == 0 { + return nil, base.Errf("GlobalMaxPool3D: the spatial dimensions must not be empty") + } + return pool3D(input, [3]int{input.Shape()[2], input.Shape()[3], input.Shape()[4]}, [3]int{1, 1, 1}, [3]int{0, 0, 0}, true, false) +} + +// AdaptiveMaxPool1D pools the input to the requested output size by +// taking the maximum over each output window. Window o covers the +// input positions from floor(o·L_in/L_out) to ceil((o+1)·L_in/L_out), +// so every input sample lands in some window. Input is (N, C, L). A +// NaN in a window propagates: the maximum answers NaN. +func AdaptiveMaxPool1D(input *core.Array, outputL int) (*core.Array, error) { + if input.NDim() != 3 { + return nil, base.Errf("AdaptiveMaxPool1D: input must be 3-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("AdaptiveMaxPool1D: complex arrays are not supported") + } + if outputL < 1 { + return nil, base.Errf("AdaptiveMaxPool1D: output size must be at least 1") + } + n, c, lIn := input.Shape()[0], input.Shape()[1], input.Shape()[2] + if lIn < 1 { + return nil, base.Errf("AdaptiveMaxPool1D: the spatial dimension must not be empty, got shape %s", base.ShapeText(input.Shape())) + } + out := core.New(core.Float, []int{n, c, outputL}...) + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per (batch, channel) row; rows are disjoint. + engine.ParallelMin(n*c, workFloorFor(lIn), func(bs, be int) { + for row := bs; row < be; row++ { + ch := row % c + batch := row / c + chanBase := batch*c*lIn + ch*lIn + for ol := range outputL { + start := ol * lIn / outputL + end := ((ol+1)*lIn + outputL - 1) / outputL + best := math.Inf(-1) + for l := start; l < end; l++ { + if v := inF[chanBase+l]; v > best || math.IsNaN(v) { + best = v + } + } + outF[batch*c*outputL+ch*outputL+ol] = best + } + } + }) + return out, nil +} + +// AdaptiveMaxPool3D pools the input to the requested output size by +// taking the maximum over each output window, with the same +// floor-start, ceil-end convention as AdaptiveMaxPool1D. Input is +// (N, C, D, H, W). A NaN in a window propagates: the maximum answers +// NaN. +func AdaptiveMaxPool3D(input *core.Array, outputD, outputH, outputW int) (*core.Array, error) { + if input.NDim() != 5 { + return nil, base.Errf("AdaptiveMaxPool3D: input must be 5-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("AdaptiveMaxPool3D: complex arrays are not supported") + } + if outputD < 1 || outputH < 1 || outputW < 1 { + return nil, base.Errf("AdaptiveMaxPool3D: output sizes must be at least 1") + } + n, c, dIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3], input.Shape()[4] + if dIn < 1 || hIn < 1 || wIn < 1 { + return nil, base.Errf("AdaptiveMaxPool3D: the spatial dimensions must not be empty, got shape %s", base.ShapeText(input.Shape())) + } + out := core.New(core.Float, []int{n, c, outputD, outputH, outputW}...) + strideHW := wIn + strideDHW := hIn * wIn + strideCDHW := dIn * hIn * wIn + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per output depth-height row; rows are disjoint. + items := n * c * outputD * outputH + floor := workFloorFor(wIn * max(1, (hIn+outputH-1)/outputH) * max(1, (dIn+outputD-1)/outputD)) + engine.ParallelMin(items, floor, func(bs, be int) { + for item := bs; item < be; item++ { + oh := item % outputH + t := item / outputH + od := t % outputD + t /= outputD + ch := t % c + batch := t / c + chanBase := batch*c*strideCDHW + ch*strideCDHW + dStart := od * dIn / outputD + dEnd := ((od+1)*dIn + outputD - 1) / outputD + hStart := oh * hIn / outputH + hEnd := ((oh+1)*hIn + outputH - 1) / outputH + for ow := range outputW { + wStart := ow * wIn / outputW + wEnd := ((ow+1)*wIn + outputW - 1) / outputW + best := math.Inf(-1) + for d := dStart; d < dEnd; d++ { + for h := hStart; h < hEnd; h++ { + inRow := chanBase + d*strideDHW + h*strideHW + for w := wStart; w < wEnd; w++ { + if v := inF[inRow+w]; v > best || math.IsNaN(v) { + best = v + } + } + } + } + outF[batch*c*outputD*outputH*outputW+ch*outputD*outputH*outputW+od*outputH*outputW+oh*outputW+ow] = best + } + } + }) + return out, nil +} + +// AdaptiveAvgPool1D pools to the requested output size by averaging +// each output window, with the same floor-start, ceil-end convention +// as AdaptiveMaxPool1D. Input is (N, C, L). +func AdaptiveAvgPool1D(input *core.Array, outputL int) (*core.Array, error) { + if input.NDim() != 3 { + return nil, base.Errf("AdaptiveAvgPool1D: input must be 3-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("AdaptiveAvgPool1D: complex arrays are not supported") + } + if outputL < 1 { + return nil, base.Errf("AdaptiveAvgPool1D: output size must be at least 1") + } + n, c, lIn := input.Shape()[0], input.Shape()[1], input.Shape()[2] + if lIn < 1 { + return nil, base.Errf("AdaptiveAvgPool1D: the spatial dimension must not be empty, got shape %s", base.ShapeText(input.Shape())) + } + out := core.New(core.Float, []int{n, c, outputL}...) + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per (batch, channel) row; rows are disjoint and + // every window sums its elements in ascending order. + engine.ParallelMin(n*c, workFloorFor(lIn), func(bs, be int) { + for row := bs; row < be; row++ { + ch := row % c + batch := row / c + chanBase := batch*c*lIn + ch*lIn + for ol := range outputL { + start := ol * lIn / outputL + end := ((ol+1)*lIn + outputL - 1) / outputL + var sum float64 + for l := start; l < end; l++ { + sum += inF[chanBase+l] + } + outF[batch*c*outputL+ch*outputL+ol] = sum / float64(end-start) + } + } + }) + return out, nil +} + +// AdaptiveAvgPool3D pools to the requested output size by averaging +// each output window, with the same floor-start, ceil-end convention +// as AdaptiveMaxPool1D. Input is (N, C, D, H, W). +func AdaptiveAvgPool3D(input *core.Array, outputD, outputH, outputW int) (*core.Array, error) { + if input.NDim() != 5 { + return nil, base.Errf("AdaptiveAvgPool3D: input must be 5-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("AdaptiveAvgPool3D: complex arrays are not supported") + } + if outputD < 1 || outputH < 1 || outputW < 1 { + return nil, base.Errf("AdaptiveAvgPool3D: output sizes must be at least 1") + } + n, c, dIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3], input.Shape()[4] + if dIn < 1 || hIn < 1 || wIn < 1 { + return nil, base.Errf("AdaptiveAvgPool3D: the spatial dimensions must not be empty, got shape %s", base.ShapeText(input.Shape())) + } + out := core.New(core.Float, []int{n, c, outputD, outputH, outputW}...) + strideHW := wIn + strideDHW := hIn * wIn + strideCDHW := dIn * hIn * wIn + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per output element here: the 3-D windows share no + // row structure worth exploiting, and every element sums its own + // window in ascending order. + items := n * c * outputD * outputH * outputW + floor := workFloorFor(max(1, (dIn+outputD-1)/outputD) * max(1, (hIn+outputH-1)/outputH) * max(1, (wIn+outputW-1)/outputW)) + engine.ParallelMin(items, floor, func(bs, be int) { + for item := bs; item < be; item++ { + ow := item % outputW + t := item / outputW + oh := t % outputH + t /= outputH + od := t % outputD + t /= outputD + ch := t % c + batch := t / c + chanBase := batch*c*strideCDHW + ch*strideCDHW + dStart := od * dIn / outputD + dEnd := ((od+1)*dIn + outputD - 1) / outputD + hStart := oh * hIn / outputH + hEnd := ((oh+1)*hIn + outputH - 1) / outputH + wStart := ow * wIn / outputW + wEnd := ((ow+1)*wIn + outputW - 1) / outputW + var sum float64 + count := 0 + for d := dStart; d < dEnd; d++ { + for h := hStart; h < hEnd; h++ { + inRow := chanBase + d*strideDHW + h*strideHW + for w := wStart; w < wEnd; w++ { + sum += inF[inRow+w] + count++ + } + } + } + if count > 0 { + outF[batch*c*outputD*outputH*outputW+ch*outputD*outputH*outputW+od*outputH*outputW+oh*outputW+ow] = sum / float64(count) + } + } + }) + return out, nil +} + +// AdaptiveAvgPool2D pools the input to the requested output size by +// averaging each output window, with the same floor-start, ceil-end +// convention as AdaptiveMaxPool2D. outputH and outputW must both be +// at least 1. +func AdaptiveAvgPool2D(input *core.Array, outputH, outputW int) (*core.Array, error) { + if input.NDim() != 4 { + return nil, base.Errf("AdaptiveAvgPool2D: input must be 4-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("AdaptiveAvgPool2D: complex arrays are not supported") + } + if outputH < 1 || outputW < 1 { + return nil, base.Errf("AdaptiveAvgPool2D: output size must be at least 1, got %dx%d", outputH, outputW) + } + n, c, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3] + // A zero spatial dimension leaves every window empty: avg would + // divide by zero, so it is an error like every other shape refusal. + if hIn < 1 || wIn < 1 { + return nil, base.Errf("AdaptiveAvgPool2D: the spatial dimensions must not be empty, got shape %s", base.ShapeText(input.Shape())) + } + out := core.New(core.Float, []int{n, c, outputH, outputW}...) + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per output row; rows are disjoint and every window + // sums its elements in ascending order. + items := n * c * outputH + engine.ParallelMin(items, workFloorFor(max(1, (hIn+outputH-1)/outputH)*wIn), func(bs, be int) { + for item := bs; item < be; item++ { + oy := item % outputH + row := item / outputH + ch := row % c + batch := row / c + chanBase := batch*c*hIn*wIn + ch*hIn*wIn + // Window in the input: floor start, ceil end, so + // every input sample lands in some window. + yStart := oy * hIn / outputH + yEnd := ((oy+1)*hIn + outputH - 1) / outputH + for ox := range outputW { + xStart := ox * wIn / outputW + xEnd := ((ox+1)*wIn + outputW - 1) / outputW + var sum float64 + for y := yStart; y < yEnd; y++ { + inRow := chanBase + y*wIn + for x := xStart; x < xEnd; x++ { + sum += inF[inRow+x] + } + } + count := float64((yEnd - yStart) * (xEnd - xStart)) + off := batch*c*outputH*outputW + ch*outputH*outputW + oy*outputW + ox + outF[off] = sum / count + } + } + }) + return out, nil +} + +// pool2D is the shared implementation behind MaxPool2D and AvgPool2D. +func pool2D(input *core.Array, kernel, stride, padding int, isMax, countIncludePad bool) (*core.Array, error) { + if input.NDim() != 4 { + return nil, base.Errf("Pool2D: input must be 4-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("Pool2D: complex arrays are not supported") + } + if kernel < 1 { + return nil, base.Errf("Pool2D: kernel must be at least 1, got %d", kernel) + } + if stride < 1 { + return nil, base.Errf("Pool2D: stride must be at least 1, got %d", stride) + } + if padding < 0 { + return nil, base.Errf("Pool2D: padding must be non-negative, got %d", padding) + } + // Padding from the kernel upward leaves output windows entirely + // inside the padding (the first window covers −padding .. + // −padding+kernel−1): a max would answer −Inf and an average 0, + // so such configurations are refused, in the spirit of the + // common pooling implementations' padding bound. + if padding >= kernel { + return nil, base.Errf("Pool2D: padding %d must stay below the kernel %d, some windows would hold no data", padding, kernel) + } + n, c, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3] + // Guard the numerators before the division: a negative numerator + // truncates toward zero and masquerades as a 1-wide output. + hNum := hIn + 2*padding - kernel + wNum := wIn + 2*padding - kernel + if hNum < 0 || wNum < 0 { + return nil, base.Errf("Pool2D: kernel %d with padding %d does not fit the %dx%d input", kernel, padding, hIn, wIn) + } + hOut := hNum/stride + 1 + wOut := wNum/stride + 1 + if hOut < 1 || wOut < 1 { + return nil, base.Errf("Pool2D: output is empty (H_out=%d, W_out=%d)", hOut, wOut) + } + out := core.New(core.Float, []int{n, c, hOut, wOut}...) + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per output row (batch, channel, oy). Rows are + // disjoint. The max and average variants walk the same windows in + // the same (kernel height, kernel width) order, but as two loop + // bodies: the max/sum choice leaves the innermost loop, which would + // otherwise pay for a perfectly predicted branch on every tap, and + // neither variant's arithmetic changes a bit. + items := n * c * hOut + engine.ParallelMin(items, workFloorFor(wOut*kernel*kernel), func(bs, be int) { + if isMax { + for item := bs; item < be; item++ { + oy := item % hOut + row := item / hOut + ch := row % c + batch := row / c + chanBase := batch*c*hIn*wIn + ch*hIn*wIn + for ox := range wOut { + best := math.Inf(-1) + for kh := range kernel { + iy := oy*stride + kh - padding + if iy < 0 || iy >= hIn { + continue + } + inRow := chanBase + iy*wIn + for kw := range kernel { + ix := ox*stride + kw - padding + if ix < 0 || ix >= wIn { + continue + } + if v := inF[inRow+ix]; v > best || math.IsNaN(v) { + best = v + } + } + } + off := batch*c*hOut*wOut + ch*hOut*wOut + oy*wOut + ox + outF[off] = best + } + } + return + } + for item := bs; item < be; item++ { + oy := item % hOut + row := item / hOut + ch := row % c + batch := row / c + chanBase := batch*c*hIn*wIn + ch*hIn*wIn + for ox := range wOut { + var acc float64 + var count int + for kh := range kernel { + iy := oy*stride + kh - padding + if iy < 0 || iy >= hIn { + continue + } + inRow := chanBase + iy*wIn + for kw := range kernel { + ix := ox*stride + kw - padding + if ix < 0 || ix >= wIn { + continue + } + acc += inF[inRow+ix] + count++ + } + } + if countIncludePad { + acc /= float64(kernel * kernel) + } else if count > 0 { + acc /= float64(count) + } + off := batch*c*hOut*wOut + ch*hOut*wOut + oy*wOut + ox + outF[off] = acc + } + } + }) + return out, nil +} + +// pool1D is the shared implementation behind MaxPool1D and AvgPool1D. +// countIncludePad is honoured only by the average variant. +func pool1D(input *core.Array, kernel, stride, padding int, isMax, countIncludePad bool) (*core.Array, error) { + if input.NDim() != 3 { + return nil, base.Errf("Pool1D: input must be 3-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("Pool1D: complex arrays are not supported") + } + if kernel < 1 || stride < 1 || padding < 0 { + return nil, base.Errf("Pool1D: kernel/stride ≥1, padding ≥0") + } + // Padding from the kernel upward leaves output windows entirely + // inside the padding (the first window covers −padding .. + // −padding+kernel−1): a max would answer −Inf and an average 0, + // so such configurations are refused, in the spirit of the + // common pooling implementations' padding bound. + if padding >= kernel { + return nil, base.Errf("Pool1D: padding %d must stay below the kernel %d, some windows would hold no data", padding, kernel) + } + n, c, lIn := input.Shape()[0], input.Shape()[1], input.Shape()[2] + lNum := lIn + 2*padding - kernel + if lNum < 0 { + return nil, base.Errf("Pool1D: kernel %d with padding %d does not fit the length-%d input", kernel, padding, lIn) + } + lOut := lNum/stride + 1 + if lOut < 1 { + return nil, base.Errf("Pool1D: output is empty (L_out=%d)", lOut) + } + out := core.New(core.Float, []int{n, c, lOut}...) + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per (batch, channel) row; rows are disjoint and + // every window walks its elements in ascending order. + engine.ParallelMin(n*c, workFloorFor(lOut*kernel), func(bs, be int) { + for row := bs; row < be; row++ { + ch := row % c + batch := row / c + chanBase := batch*c*lIn + ch*lIn + for ol := range lOut { + var acc float64 + var best float64 + if isMax { + best = math.Inf(-1) + } + var count int + for kl := range kernel { + il := ol*stride + kl - padding + if il < 0 || il >= lIn { + continue + } + v := inF[chanBase+il] + if isMax { + if v > best || math.IsNaN(v) { + best = v + } + } else { + acc += v + } + count++ + } + if isMax { + outF[batch*c*lOut+ch*lOut+ol] = best + } else if countIncludePad { + outF[batch*c*lOut+ch*lOut+ol] = acc / float64(kernel) + } else if count > 0 { + outF[batch*c*lOut+ch*lOut+ol] = acc / float64(count) + } + } + } + }) + return out, nil +} + +// pool3D is the shared implementation behind MaxPool3D and AvgPool3D. +// countIncludePad is honoured only by the average variant. +func pool3D(input *core.Array, kernel, stride, padding [3]int, isMax, countIncludePad bool) (*core.Array, error) { + if input.NDim() != 5 { + return nil, base.Errf("Pool3D: input must be 5-D, got shape %s", base.ShapeText(input.Shape())) + } + if input.Dtype() == core.Complex { + return nil, base.Errf("Pool3D: complex arrays are not supported") + } + for d := range kernel { + if kernel[d] < 1 || stride[d] < 1 || padding[d] < 0 { + return nil, base.Errf("Pool3D: kernel/stride ≥1, padding ≥0") + } + // Padding from the kernel upward leaves output windows + // entirely inside the padding: a max would answer −Inf and + // an average 0, so such configurations are refused, in the + // spirit of the common pooling implementations' padding + // bound. + if padding[d] >= kernel[d] { + return nil, base.Errf("Pool3D: padding %d must stay below the kernel %d in dimension %d, some windows would hold no data", padding[d], kernel[d], d) + } + } + n, c, dIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3], input.Shape()[4] + // Guard the numerators before the division: a negative numerator + // truncates toward zero and masquerades as a 1-wide output. + dNum := dIn + 2*padding[0] - kernel[0] + hNum := hIn + 2*padding[1] - kernel[1] + wNum := wIn + 2*padding[2] - kernel[2] + if dNum < 0 || hNum < 0 || wNum < 0 { + return nil, base.Errf("Pool3D: kernel %v with padding %v does not fit the input", kernel, padding) + } + dOut := dNum/stride[0] + 1 + hOut := hNum/stride[1] + 1 + wOut := wNum/stride[2] + 1 + if dOut < 1 || hOut < 1 || wOut < 1 { + return nil, base.Errf("Pool3D: output is empty") + } + out := core.New(core.Float, []int{n, c, dOut, hOut, wOut}...) + strideHW := wIn + strideDHW := hIn * wIn + strideCDHW := dIn * hIn * wIn + inF := input.RawFloats() + if inF == nil || input.Strided() { + inF = widenFloats(input) + } + outF := out.RawFloats() + // One work item per output depth-height row (batch, channel, od, + // oh); rows are disjoint and each window walks its elements in the + // (kernel depth, kernel height, kernel width) order. + items := n * c * dOut * hOut + engine.ParallelMin(items, workFloorFor(wOut*kernel[0]*kernel[1]*kernel[2]), func(bs, be int) { + for item := bs; item < be; item++ { + oh := item % hOut + t := item / hOut + od := t % dOut + t /= dOut + ch := t % c + batch := t / c + chanBase := batch*c*strideCDHW + ch*strideCDHW + for ow := range wOut { + var acc float64 + var best float64 + if isMax { + best = math.Inf(-1) + } + var count int + for kd := range kernel[0] { + id := od*stride[0] + kd - padding[0] + if id < 0 || id >= dIn { + continue + } + inDepth := chanBase + id*strideDHW + for kh := range kernel[1] { + ih := oh*stride[1] + kh - padding[1] + if ih < 0 || ih >= hIn { + continue + } + inRow := inDepth + ih*strideHW + for kw := range kernel[2] { + iw := ow*stride[2] + kw - padding[2] + if iw < 0 || iw >= wIn { + continue + } + v := inF[inRow+iw] + if isMax { + if v > best || math.IsNaN(v) { + best = v + } + } else { + acc += v + } + count++ + } + } + } + off := batch*c*dOut*hOut*wOut + ch*dOut*hOut*wOut + od*hOut*wOut + oh*wOut + ow + if isMax { + outF[off] = best + } else if countIncludePad { + outF[off] = acc / float64(kernel[0]*kernel[1]*kernel[2]) + } else if count > 0 { + outF[off] = acc / float64(count) + } + } + } + }) + return out, nil +} diff --git a/signal/pooling2d_test.go b/signal/pooling2d_test.go new file mode 100644 index 0000000..55c095b --- /dev/null +++ b/signal/pooling2d_test.go @@ -0,0 +1,194 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import "sourcedock.dev/petrbalvin/tensor/internal/core" + +import ( + "math" + "testing" +) + +func TestConv2D(t *testing.T) { + // 1×1×3×3 input, 1×1×2×2 kernel (basic case). + in, _ := core.FromFloats([]float64{ + 1, 2, 3, + 4, 5, 6, + 7, 8, 9, + }, 1, 1, 3, 3) + ker, _ := core.FromFloats([]float64{ + 1, 0, + 0, 1, + }, 1, 1, 2, 2) + got, err := Conv2D(in, ker, nil, 1, 0) + if err != nil { + t.Fatal(err) + } + // 2×2 output: conv at (oh, ow) = sum(in[oh..oh+2, ow..ow+2] * ker). + // (0,0) = 1*1 + 2*0 + 4*0 + 5*1 = 6 + // (0,1) = 2*1 + 3*0 + 5*0 + 6*1 = 8 + // (1,0) = 4*1 + 5*0 + 7*0 + 8*1 = 12 + // (1,1) = 5*1 + 6*0 + 8*0 + 9*1 = 14 + for i, w := range []float64{6, 8, 12, 14} { + v, _ := core.FloatAt(got, 0, 0, i/2, i%2) + if v != w { + t.Errorf("conv2d [%d]: got %v, want %v", i, v, w) + } + } + // With stride=2, output is 1×1. + got, err = Conv2D(in, ker, nil, 2, 0) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 1 || got.Shape()[3] != 1 { + t.Errorf("conv2d stride 2: shape = %v", got.Shape()) + } + // With bias. + bias, _ := core.FromFloats([]float64{1}, 1) + got, err = Conv2D(in, ker, bias, 1, 0) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0, 0, 0, 0) + if v != 7 { + t.Errorf("conv2d+bias [0,0,0,0]: got %v, want 7", v) + } + // Rank error. + badIn, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := Conv2D(badIn, ker, nil, 1, 0); err == nil { + t.Error("conv2d: expected error for wrong input rank") + } +} + +func TestMaxPool2D(t *testing.T) { + in, _ := core.FromFloats([]float64{ + 1, 3, 2, 4, + 5, 6, 7, 8, + 9, 2, 3, 1, + 4, 0, 6, 2, + }, 1, 1, 4, 4) + got, err := MaxPool2D(in, 2, 2, 0) + if err != nil { + t.Fatal(err) + } + // 2×2 output. Each 2×2 window picks the max. + // TL: max(1,3,5,6) = 6; TR: max(2,4,7,8) = 8 + // BL: max(9,2,4,0) = 9; BR: max(3,1,6,2) = 6 + for i, w := range []float64{6, 8, 9, 6} { + v, _ := core.FloatAt(got, 0, 0, i/2, i%2) + if v != w { + t.Errorf("maxpool2d [%d]: got %v, want %v", i, v, w) + } + } + // With stride=1 and kernel=3: 2×2 output, overlapping. + got, err = MaxPool2D(in, 3, 1, 0) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 2 || got.Shape()[3] != 2 { + t.Errorf("maxpool2d 3x1: shape = %v", got.Shape()) + } + // With padding=1, kernel=3, stride=1: output is (H+2*1-3)/1+1 = H. + got, err = MaxPool2D(in, 3, 1, 1) + if err != nil { + t.Fatal(err) + } + if got.Shape()[2] != 4 || got.Shape()[3] != 4 { + t.Errorf("maxpool2d 3x1+pad: shape = %v", got.Shape()) + } + // Rank error. + bad, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := MaxPool2D(bad, 2, 2, 0); err == nil { + t.Error("maxpool2d: expected error for wrong rank") + } +} + +func TestAvgPool2D(t *testing.T) { + in, _ := core.FromFloats([]float64{ + 1, 2, + 3, 4, + }, 1, 1, 2, 2) + got, err := AvgPool2D(in, 2, 2, 0, true) + if err != nil { + t.Fatal(err) + } + v, _ := core.FloatAt(got, 0, 0, 0, 0) + if v != 2.5 { + t.Errorf("avgpool2d: got %v, want 2.5", v) + } + // Without counting padding: same result when no padding. + got, err = AvgPool2D(in, 2, 2, 0, false) + if err != nil { + t.Fatal(err) + } + v, _ = core.FloatAt(got, 0, 0, 0, 0) + if v != 2.5 { + t.Errorf("avgpool2d (no pad): got %v, want 2.5", v) + } +} + +func TestAvgPool2DCountIncludePad(t *testing.T) { + // 1×1×2×2 with kernel 3, stride 1, padding 1. The output is 2×2. + // Window (0,0) covers input rows/cols -1..1, so rows/cols 0..1 are + // real, the rest padded: 4 real cells of sum 10. includePad divides + // by 9 (10/9 ≈ 1.111), excludePad by 4 (10/4 = 2.5). + in, _ := core.FromFloats([]float64{ + 1, 2, + 3, 4, + }, 1, 1, 2, 2) + inc, err := AvgPool2D(in, 3, 1, 1, true) + if err != nil { + t.Fatal(err) + } + exc, err := AvgPool2D(in, 3, 1, 1, false) + if err != nil { + t.Fatal(err) + } + for i := range 4 { + vi, _ := core.FloatAt(inc, 0, 0, i/2, i%2) + ve, _ := core.FloatAt(exc, 0, 0, i/2, i%2) + // The includePad divisor is the full kernel (9); the excludePad + // divisor is the number of real cells, which is 4 for all four + // windows of a 2×2 input with kernel 3 and padding 1. + if math.Abs(vi-10.0/9) > 1e-9 { + t.Errorf("includePad [%d]: got %v, want %v", i, vi, 10.0/9) + } + if math.Abs(ve-2.5) > 1e-9 { + t.Errorf("excludePad [%d]: got %v, want 2.5", i, ve) + } + } +} + +func TestAdaptiveAvgPool2D(t *testing.T) { + // 1×1×4×4 gives 1×1×2×2. + in, _ := core.FromFloats([]float64{ + 1, 2, 3, 4, + 5, 6, 7, 8, + 9, 10, 11, 12, + 13, 14, 15, 16, + }, 1, 1, 4, 4) + got, err := AdaptiveAvgPool2D(in, 2, 2) + if err != nil { + t.Fatal(err) + } + // Top-left 2x2: avg(1,2,5,6) = 3.5 + v, _ := core.FloatAt(got, 0, 0, 0, 0) + if math.Abs(v-3.5) > 1e-9 { + t.Errorf("adaptive avgpool [0,0,0,0]: got %v, want 3.5", v) + } + // Top-right 2x2: avg(3,4,7,8) = 5.5 + v, _ = core.FloatAt(got, 0, 0, 0, 1) + if math.Abs(v-5.5) > 1e-9 { + t.Errorf("adaptive avgpool [0,0,0,1]: got %v, want 5.5", v) + } + // Invalid output size. + if _, err := AdaptiveAvgPool2D(in, 0, 1); err == nil { + t.Error("adaptive avgpool: expected error for outputH=0") + } + // Rank error. + bad, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := AdaptiveAvgPool2D(bad, 2, 2); err == nil { + t.Error("adaptive avgpool: expected error for wrong rank") + } +} diff --git a/signal/rankfilter.go b/signal/rankfilter.go new file mode 100644 index 0000000..c227986 --- /dev/null +++ b/signal/rankfilter.go @@ -0,0 +1,258 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Rank-statistic filters. Every output sample is one order +// statistic of the window around it: the median for the middle rank, +// the minimum for rank 0, the maximum for the last rank. Order +// statistics ignore the shape of the tail inside the window, which is +// why the median removes an impulsive spike outright where a linear +// filter would smear it, and why a rank filter holds an edge where a +// moving average would round it off. +// +// The kernels hold two conventions fixed: +// +// - Edges. Windows near a boundary are truncated to the samples +// that exist (the SavitzkyGolay convention), so the output is +// complete without padding, and the requested rank scales to the +// truncation: rank k of a full window becomes rank k·m/window of +// an m-sample truncation, so a median stays the middle of +// whatever samples the edge left and an extremum stays an +// extremum. +// - NaN. As in the pooling kernels, a NaN in a window propagates: +// the window's order statistic answers NaN, rather than sorting +// it silently into some position. +// +// Non-complex inputs widen once through widenFloats, so the values +// compared are exactly the elements the accessors report; the outputs +// are float64, as for the pooling kernels. + +// rankMinPointsPerWorker is the smallest per-worker chunk of output +// points the 1-D rank sweep splits for: below it a chunk's window +// sorts no longer pay the worker spawn cost. +const rankMinPointsPerWorker = 1 << 10 + +// MedianFilter returns the running median of the rank-1 signal x +// over an odd window: each sample becomes the middle order statistic +// of its neighbourhood. A polynomial-free smoother that removes +// isolated spikes whole and holds monotone ramps and edges still. +// The window must be odd and at least 3, must not exceed the signal, +// and a NaN in a window propagates into that output sample. +func MedianFilter(x *core.Array, window int) (*core.Array, error) { + return RankFilter(x, window, window/2) +} + +// RankFilter returns the k-th order statistic (ascending, 0-based) of +// each window of the rank-1 signal x: rank 0 is the running minimum, +// window−1 the running maximum, window/2 the median MedianFilter +// takes. The window must be odd and at least 3 and must not exceed +// the signal; k must lie in [0, window). +func RankFilter(x *core.Array, window, k int) (*core.Array, error) { + const name = "RankFilter" + if err := rankGates(name, x, 1, window, k, window); err != nil { + return nil, err + } + n := x.Len() + if window > n { + return nil, base.Errf("%s: window %d exceeds the signal length %d", name, window, n) + } + src := widenFloats(x) + out := core.New(core.Float, n) + dst := out.RawFloats() + // Output points split across workers: each owns disjoint slots + // and its own window scratch, so the split cannot move a bit. + engine.ParallelMin(n, rankMinPointsPerWorker, func(ps, pe int) { + win := make([]float64, window) + // sorted mirrors the current window in ascending order, so the + // rank's element is a direct read instead of a fresh sort of the + // same values. Neighbouring windows share all but one element, so + // the mirror slides: one removal and one insertion per output + // point. + // + // The mirror answers the rank only while the window holds no zero. + // A zero may sit next to a −0, its equal under comparison but not + // in bits, and which of the two the sort leaves at the rank is the + // sort's own tie choice, which no multiset can predict. Those + // windows take the exact path below, the sort this filter has + // always run. A NaN entering or leaving the window likewise drops + // the mirror, which is rebuilt from the window on the next point. + sorted := make([]float64, 0, window) + zeros := 0 + mirror := false + half := window / 2 + prevLo, prevHi := 0, 0 + for i := ps; i < pe; i++ { + lo := max(i-half, 0) + hi := min(i+half+1, n) + if i > ps && mirror { + // Leave behind what the window no longer covers and take + // up what it now does. Both bounds advance by at most one + // element, so at most one leaves and one enters. + drop := false + for j := prevLo; j < lo && !drop; j++ { + if v := src[j]; !math.IsNaN(v) { + if v == 0 { + zeros-- + } + pos, _ := slices.BinarySearch(sorted, v) + sorted = slices.Delete(sorted, pos, pos+1) + continue + } + drop = true + } + for j := prevHi; j < hi && !drop; j++ { + if v := src[j]; !math.IsNaN(v) { + if v == 0 { + zeros++ + } + pos, _ := slices.BinarySearch(sorted, v) + sorted = slices.Insert(sorted, pos, v) + continue + } + drop = true + } + mirror = !drop + } + prevLo, prevHi = lo, hi + if !mirror { + // A NaN is not part of the mirror, so a window holding one + // answers NaN and leaves the mirror to be rebuilt later. + sorted, zeros = sorted[:0], 0 + nan := false + for j := lo; j < hi; j++ { + v := src[j] + if math.IsNaN(v) { + nan = true + break + } + if v == 0 { + zeros++ + } + sorted = append(sorted, v) + } + if nan { + dst[i] = math.NaN() + continue + } + slices.Sort(sorted) + mirror = true + } + count := hi - lo + // The rank scales to the truncation; k below the full + // window's capacity keeps the product below count and + // the index inside the collected prefix. + rank := k * count / window + if zeros > 0 { + for j := lo; j < hi; j++ { + win[j-lo] = src[j] + } + slices.Sort(win[:count]) + dst[i] = win[rank] + continue + } + dst[i] = sorted[rank] + } + }) + return out, nil +} + +// MedianFilter2D returns the running median of the rank-2 image img +// over a square odd window, the standard impulse-noise cleaner: each +// pixel becomes the middle of the values in its neighbourhood, so +// salt-and-pepper dots vanish while steps between regions keep their +// corners. The window must be odd and at least 3 and must fit both +// dimensions; a NaN in a window propagates into that output pixel. +func MedianFilter2D(img *core.Array, window int) (*core.Array, error) { + return RankFilter2D(img, window, window*window/2) +} + +// RankFilter2D returns the k-th order statistic (ascending, 0-based) +// of each square window of the rank-2 image img: rank 0 is the +// running minimum (erosion), window²−1 the running maximum (dilation), +// the middle rank the median. The window must be odd and at least 3 +// and must fit both dimensions; k must lie in [0, window²). +func RankFilter2D(img *core.Array, window, k int) (*core.Array, error) { + const name = "RankFilter2D" + if err := rankGates(name, img, 2, window, k, window*window); err != nil { + return nil, err + } + h, w := img.Shape()[0], img.Shape()[1] + if window > h || window > w { + return nil, base.Errf("%s: window %d does not fit the %dx%d image", name, window, h, w) + } + src := widenFloats(img) + out := core.New(core.Float, h, w) + dst := out.RawFloats() + // One work item per image row; rows own disjoint output pixels + // and their own window scratch, and every window walks its + // elements in row-major order. + engine.ParallelMin(h, workFloorFor(w*window*window), func(ys, ye int) { + win := make([]float64, window*window) + half := window / 2 + for y := ys; y < ye; y++ { + yLo := max(y-half, 0) + yHi := min(y+half+1, h) + for x := range w { + xLo := max(x-half, 0) + xHi := min(x+half+1, w) + count := 0 + nan := false + for yy := yLo; yy < yHi && !nan; yy++ { + row := yy * w + for xx := xLo; xx < xHi; xx++ { + v := src[row+xx] + if math.IsNaN(v) { + nan = true + break + } + win[count] = v + count++ + } + } + if nan { + dst[y*w+x] = math.NaN() + continue + } + slices.Sort(win[:count]) + // The rank scales to the truncation; k below the + // full window's capacity window² keeps the + // product below count and the index inside the + // collected prefix. + dst[y*w+x] = win[k*count/(window*window)] + } + } + }) + return out, nil +} + +// rankGates checks the contract the four public rank filters share: +// the promised rank, a real dtype, a non-empty input, an odd window +// of at least 3, and a rank inside the window's capacity kMax. +func rankGates(name string, a *core.Array, ndim, window, k, kMax int) error { + if a.NDim() != ndim { + return base.Errf("%s: needs a rank-%d input, got shape %s", name, ndim, base.ShapeText(a.Shape())) + } + if a.Dtype() == core.Complex { + return base.Errf("%s: complex arrays are not supported", name) + } + if a.Len() == 0 { + return base.Errf("%s: the input must not be empty", name) + } + if window < 3 || window%2 == 0 { + return base.Errf("%s: window must be an odd number ≥ 3, got %d", name, window) + } + if k < 0 || k >= kMax { + return base.Errf("%s: k must lie in [0, %d) for window %d, got %d", name, kMax, window, k) + } + return nil +} diff --git a/signal/rankfilter_test.go b/signal/rankfilter_test.go new file mode 100644 index 0000000..c88ea8e --- /dev/null +++ b/signal/rankfilter_test.go @@ -0,0 +1,353 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestMedianFilterBruteForce pins the 1-D median against a naive +// sort-per-sample reference on randomised fixtures, edges included +// under the documented truncation policy. +func TestMedianFilterBruteForce(t *testing.T) { + g := core.NewGenerator(7) + for _, c := range []struct{ n, window int }{ + {16, 3}, {31, 5}, {64, 7}, {50, 9}, {7, 7}, + } { + vals := make([]float64, c.n) + for i := range vals { + vals[i] = g.NormalUnit() + } + got, err := MedianFilter(mustFloats(t, vals), c.window) + if err != nil { + t.Fatalf("MedianFilter(n=%d, window=%d): %v", c.n, c.window, err) + } + for i := range c.n { + lo := max(i-c.window/2, 0) + hi := min(i+c.window/2+1, c.n) + ref := slices.Clone(vals[lo:hi]) + slices.Sort(ref) + want := ref[(c.window/2)*len(ref)/c.window] + if got.FloatAt(i) != want { + t.Fatalf("n=%d window=%d sample %d: %v, want %v", c.n, c.window, i, got.FloatAt(i), want) + } + } + } +} + +// TestRankFilterMatchesMedianAndExtremes pins the general rank +// filter: the middle rank reproduces the median sample for sample, +// rank 0 and the last rank reproduce the running minimum and maximum, +// each against the brute-force reference. +func TestRankFilterMatchesMedianAndExtremes(t *testing.T) { + g := core.NewGenerator(21) + const n = 40 + vals := make([]float64, n) + for i := range vals { + vals[i] = g.NormalUnit() + } + x := mustFloats(t, vals) + for _, window := range []int{3, 5, 9} { + med, err := MedianFilter(x, window) + if err != nil { + t.Fatalf("MedianFilter(window=%d): %v", window, err) + } + asRank, err := RankFilter(x, window, window/2) + if err != nil { + t.Fatalf("RankFilter(window=%d): %v", window, err) + } + for i := range n { + if med.FloatAt(i) != asRank.FloatAt(i) { + t.Fatalf("window %d: median and middle rank disagree at %d", window, i) + } + } + for _, c := range []struct{ rank int }{{0}, {window - 1}} { + got, err := RankFilter(x, window, c.rank) + if err != nil { + t.Fatalf("RankFilter(window=%d, rank=%d): %v", window, c.rank, err) + } + for i := range n { + lo := max(i-window/2, 0) + hi := min(i+window/2+1, n) + ref := slices.Clone(vals[lo:hi]) + slices.Sort(ref) + want := ref[c.rank*len(ref)/window] + if got.FloatAt(i) != want { + t.Fatalf("window %d rank %d at %d: %v, want %v", + window, c.rank, i, got.FloatAt(i), want) + } + } + } + } +} + +// TestMedianFilterEdges pins the documented edge policy by hand: the +// window truncates at the boundary and the rank scales to the +// truncation, so the two-sample edge windows answer their lower +// median. +func TestMedianFilterEdges(t *testing.T) { + // i=0: [1,3] at rank 1·2/3 = 0 → 1; i=1: [1,2,3] → 2; + // i=2: [1,2] at rank 0 → 1. + got, err := MedianFilter(mustFloats(t, []float64{3, 1, 2}), 3) + if err != nil { + t.Fatalf("MedianFilter: %v", err) + } + for i, want := range []float64{1, 2, 1} { + if got.FloatAt(i) != want { + t.Fatalf("sample %d: %v, want %v", i, got.FloatAt(i), want) + } + } +} + +// TestMedianFilterDeepTruncation pins the edge policy where the +// truncation runs deeper than the shallow two-sample cut: a window of +// 7 keeps only 4 samples at the boundary, and the scaled rank +// 3·4/7 = 1 answers the second-smallest of the four, not the median +// of 7 the full window would carry. +func TestMedianFilterDeepTruncation(t *testing.T) { + vals := []float64{50, 10, 40, 20, 90, 30, 80, 60, 70} + got, err := MedianFilter(mustFloats(t, vals), 7) + if err != nil { + t.Fatalf("MedianFilter: %v", err) + } + // i=0: {50,10,40,20} sorted [10,20,40,50], rank 3·4/7 = 1 → 20. + // i=4: the full window {10,40,20,90,30,80,60} sorted, rank 3 → 40. + // i=8: {30,80,60,70} sorted [30,60,70,80], rank 1 → 60. Every + // other position hand-computed the same way. + for i, want := range []float64{20, 40, 30, 40, 40, 60, 60, 70, 60} { + if got.FloatAt(i) != want { + t.Fatalf("sample %d: %v, want %v", i, got.FloatAt(i), want) + } + } +} + +// TestMedianFilterNaNWindow pins the NaN contract: a NaN anywhere in +// a window makes that output sample NaN, and windows clear of it +// stay exact. +func TestMedianFilterNaNWindow(t *testing.T) { + vals := []float64{1, math.NaN(), 3, 4, 5} + got, err := MedianFilter(mustFloats(t, vals), 3) + if err != nil { + t.Fatalf("MedianFilter: %v", err) + } + for i, want := range []bool{true, true, true, false, false} { + isNaN := math.IsNaN(got.FloatAt(i)) + if isNaN != want { + t.Fatalf("sample %d: NaN = %v, want %v", i, isNaN, want) + } + } + // The last window truncates to [4,5] and the scaled rank takes + // its lower median, so both clean samples read 4. + if got.FloatAt(3) != 4 || got.FloatAt(4) != 4 { + t.Fatalf("clean windows moved: %g, %g", got.FloatAt(3), got.FloatAt(4)) + } +} + +// TestMedianFilter2DNaNWindow pins the NaN contract in two +// dimensions: windows touching the NaN answer NaN, windows clear of +// it stay exact under the truncation policy. +func TestMedianFilter2DNaNWindow(t *testing.T) { + vals := make([]float64, 16) + for i := range vals { + vals[i] = float64(i) + } + vals[1*4+1] = math.NaN() + got, err := MedianFilter2D(mustFloats(t, vals, 4, 4), 3) + if err != nil { + t.Fatalf("MedianFilter2D: %v", err) + } + for y := range 4 { + for x := range 4 { + v := got.FloatAt(y*4 + x) + near := max(absInt(y-1), absInt(x-1)) <= 1 + if near && !math.IsNaN(v) { + t.Fatalf("window touching the NaN answered %v at (%d,%d)", v, y, x) + } + if !near && math.IsNaN(v) { + t.Fatalf("clean window answered NaN at (%d,%d)", y, x) + } + } + } + // Hand check one clean corner: window {[10,11],[14,15]}, rank + // 4·4/9 = 1. + if got.FloatAt(3*4+3) != 11 { + t.Fatalf("corner median %v, want 11", got.FloatAt(3*4+3)) + } +} + +// TestMedianFilter2DConstant pins the flat case: a constant image +// stays constant under any window, edges included. +func TestMedianFilter2DConstant(t *testing.T) { + const h, w = 9, 11 + vals := make([]float64, h*w) + for i := range vals { + vals[i] = 3.25 + } + for _, window := range []int{3, 5} { + got, err := MedianFilter2D(mustFloats(t, vals, h, w), window) + if err != nil { + t.Fatalf("MedianFilter2D(window=%d): %v", window, err) + } + for i := range vals { + if got.FloatAt(i) != 3.25 { + t.Fatalf("window %d: pixel %d = %v, want 3.25", window, i, got.FloatAt(i)) + } + } + } +} + +// TestMedianFilter2DRemovesSpikes pins the impulse-noise promise on a +// strictly monotone ramp: every full window clear of a spike comes +// back exact (the ramp's centre value is its window's median), every +// spike collapses to within one ramp step of the truth, where the +// input error was a thousand, and the truncated border windows stay +// equally close under the documented rank scaling. +func TestMedianFilter2DRemovesSpikes(t *testing.T) { + const h, w = 12, 15 + clean := make([]float64, h*w) + for y := range h { + for x := range w { + clean[y*w+x] = float64(10*y + x) + } + } + // Every spike sits far enough inside that each window holding + // one still clears the median rank: a border window truncates to + // its m live samples and scales the rank to k·m/window, so a + // spike planted in one is not judged by the interior's rank at + // all. + spikes := [][2]int{{3, 4}, {3, 5}, {8, 11}, {2, 2}, {10, 2}} + noisy := slices.Clone(clean) + for _, c := range spikes { + noisy[c[0]*w+c[1]] += 1000 + } + got, err := MedianFilter2D(mustFloats(t, noisy, h, w), 3) + if err != nil { + t.Fatalf("MedianFilter2D: %v", err) + } + for i := range clean { + d := math.Abs(got.FloatAt(i) - clean[i]) + if d > 11 { + t.Fatalf("pixel %d: %v, off the clean ramp by %g", i, got.FloatAt(i), d) + } + } + const half = 1 + for y := half; y < h-half; y++ { + for x := half; x < w-half; x++ { + near := false + for _, c := range spikes { + if max(absInt(y-c[0]), absInt(x-c[1])) <= 1 { + near = true + break + } + } + if !near && got.FloatAt(y*w+x) != clean[y*w+x] { + t.Fatalf("pixel (%d,%d) clear of every spike: %v, want %v", + y, x, got.FloatAt(y*w+x), clean[y*w+x]) + } + } + } +} + +// absInt is the small integer absolute value the spike proximity +// check needs. +func absInt(v int) int { + if v < 0 { + return -v + } + return v +} + +// TestRankFilter2DBruteForce pins the 2-D rank statistics, medians +// and erosions both, against a naive sort-per-pixel reference on a +// randomised image, border windows included. +func TestRankFilter2DBruteForce(t *testing.T) { + g := core.NewGenerator(33) + const h, w = 8, 9 + vals := make([]float64, h*w) + for i := range vals { + vals[i] = g.NormalUnit() + } + for _, c := range []struct{ window, rank int }{ + {3, 4}, {3, 0}, {5, 12}, {5, 24}, + } { + got, err := RankFilter2D(mustFloats(t, vals, h, w), c.window, c.rank) + if err != nil { + t.Fatalf("RankFilter2D(window=%d, rank=%d): %v", c.window, c.rank, err) + } + half := c.window / 2 + for y := range h { + for x := range w { + var ref []float64 + for yy := max(y-half, 0); yy < min(y+half+1, h); yy++ { + for xx := max(x-half, 0); xx < min(x+half+1, w); xx++ { + ref = append(ref, vals[yy*w+xx]) + } + } + slices.Sort(ref) + want := ref[c.rank*len(ref)/(c.window*c.window)] + if got.FloatAt(y*w+x) != want { + t.Fatalf("window %d rank %d at (%d,%d): %v, want %v", + c.window, c.rank, y, x, got.FloatAt(y*w+x), want) + } + } + } + } +} + +// TestMedianFilterWidening pins the dtype behaviour: an int input +// widens to a float output, exactly as pooling widens. +func TestMedianFilterWidening(t *testing.T) { + ints, err := core.FromInts([]int64{5, 3, 4, 1, 2}, 5) + if err != nil { + t.Fatal(err) + } + got, err := MedianFilter(ints, 3) + if err != nil { + t.Fatalf("MedianFilter: %v", err) + } + if got.Dtype() != core.Float { + t.Fatalf("int input produced dtype %s", got.Dtype()) + } + if got.FloatAt(1) != 4 || got.FloatAt(2) != 3 { + t.Fatalf("median of ints wrong: %g, %g", got.FloatAt(1), got.FloatAt(2)) + } +} + +// TestRankFilterErrors pins the input gates: rank and dtype, the odd +// window, the rank range, and the fit of the window to the input. +func TestRankFilterErrors(t *testing.T) { + x := mustFloats(t, []float64{1, 2, 3, 4, 5}) + for _, c := range []struct { + name string + call func() error + }{ + {"even window", func() error { _, err := MedianFilter(x, 4); return err }}, + {"window 1", func() error { _, err := MedianFilter(x, 1); return err }}, + {"window 0", func() error { _, err := MedianFilter(x, 0); return err }}, + {"negative k", func() error { _, err := RankFilter(x, 3, -1); return err }}, + {"k at window", func() error { _, err := RankFilter(x, 3, 3); return err }}, + {"window over length", func() error { _, err := MedianFilter(x, 7); return err }}, + {"rank 2", func() error { _, err := MedianFilter(mustFloats(t, []float64{1, 2, 3, 4}, 2, 2), 3); return err }}, + {"empty", func() error { _, err := MedianFilter(mustFloats(t, nil), 3); return err }}, + {"complex", func() error { _, err := MedianFilter(mustComplexes(t, []complex128{1, 2, 3, 4, 5}, 5), 3); return err }}, + } { + if err := c.call(); err == nil { + t.Errorf("%s: want an error", c.name) + } + } + img := mustFloats(t, make([]float64, 25), 5, 5) + if _, err := MedianFilter2D(img, 7); err == nil { + t.Error("window wider than the image accepted") + } + if _, err := RankFilter2D(img, 3, 9); err == nil { + t.Error("k at window² accepted") + } + if _, err := MedianFilter2D(mustFloats(t, []float64{1, 2, 3, 4}), 3); err == nil { + t.Error("rank-1 input accepted by the 2-D filter") + } +} diff --git a/signal/resample.go b/signal/resample.go new file mode 100644 index 0000000..cb140e3 --- /dev/null +++ b/signal/resample.go @@ -0,0 +1,249 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Sample-rate conversion. Decimation and rational resampling run a +// linear-phase FIR anti-alias filter in the time domain, so they suit +// aperiodic streams; ResampleFourier is the exact band-limited +// resample of the Fourier definition and suits whole records. + +// kaiserI0 evaluates the modified Bessel function of the first kind +// at a real order zero by its power series, which converges in a few +// dozen terms across the window betas any reasonable design uses. +func kaiserI0(x float64) float64 { + sum, term := 1.0, 1.0 + half := x / 2 + for k := 1; k < 64; k++ { + term *= (half * half) / float64(k*k) + sum += term + if term < 1e-18*sum { + break + } + } + return sum +} + +// kaiserSinc builds the odd-length FIR low-pass of a windowed sinc: +// gain 1 at DC, cutoff fc normalised to the sampling rate (half the +// output rate after conversion), Kaiser taper with beta 8.6, which +// puts the stopband around 80 dB. The cutoff rides the middle of the +// transition band, so the design slightly attenuates the top of the +// passband by construction; tests pin that below a dB. +func kaiserSinc(taps int, fc float64) []float64 { + if taps%2 == 0 { + taps++ + } + // A 1-tap kernel has no window to apply: the ratio inside r + // divides by taps−1 = 0 and turns every coefficient NaN. The + // kernel is the identity, gain 1 at DC, which is exactly what a + // one-tap low-pass means: keep the sample, filter nothing. + if taps == 1 { + return []float64{1} + } + const beta = 8.6 + i0b := kaiserI0(beta) + h := make([]float64, taps) + centre := (taps - 1) / 2 + sum := 0.0 + for i := range taps { + r := 2*float64(i)/float64(taps-1) - 1 + w := kaiserI0(beta*math.Sqrt(math.Max(0, 1-r*r))) / i0b + arg := 2 * fc * float64(i-centre) + s := 1.0 + if arg != 0 { + s = math.Sin(math.Pi*arg) / (math.Pi * arg) + } + h[i] = 2 * fc * s * w + sum += h[i] + } + for i := range h { + h[i] /= sum + } + return h +} + +// Decimate reduces the sample rate by the integer factor: a Kaiser +// tapered FIR anti-alias filter runs first, the filter's group delay +// is compensated, and every factor-th sample of the compensated +// series is kept. The passband ends at nine tenths of the new +// Nyquist; the transition band then reaches its stopband floor of +// about 80 dB before the new Nyquist, so everything that would fold +// into the output band is suppressed. Taps sets the filter length; a +// non-positive value means 32·factor+1, which is the length that +// fits that transition. Output starts once the filter has full +// context, so the result holds about (n − taps)/factor samples of an +// input of n; when the tap count leaves no sample with full context, +// the result is empty rather than read past the end of the signal. +func Decimate(data *core.Array, factor, taps int) (*core.Array, error) { + const name = "Decimate" + if data.NDim() != 1 { + return nil, base.Errf("%s: the series must be a vector, got shape %s", name, base.ShapeText(data.Shape())) + } + if data.Dtype() == core.Complex { + return nil, base.Errf("%s: complex series are not supported", name) + } + if factor < 2 { + return nil, base.Errf("%s: the factor must be at least 2, got %d", name, factor) + } + n := data.Len() + if n == 0 { + return nil, base.Errf("%s: the series must not be empty", name) + } + if taps <= 0 { + taps = 32*factor + 1 + } + if taps%2 == 0 { + taps++ + } + if taps >= n { + return nil, base.Errf("%s: %d taps against %d samples leaves nothing after the filter delay", name, taps, n) + } + h := kaiserSinc(taps, 0.45/float64(factor)) + filtered, err := FilterApply(h, []float64{1}, data) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + delay := (taps - 1) / 2 + // Sample the compensated series (compensated index m lives at + // filtered[m + delay]) from where its filter context is complete: + // m + delay ≥ taps − 1 means m ≥ delay, so the first kept index + // is the first multiple of factor at or past delay, and the last + // is the last multiple whose compensated read still lies inside + // the signal. When no multiple qualifies, the result is empty: + // a negative numerator must not be left to truncate toward zero, + // which used to turn "no fully filtered sample" into one read + // past the end of the signal. + skip := (delay + factor - 1) / factor + outLen := 0 + if span := filtered.Len() - 1 - delay; span >= 0 { + outLen = max(span/factor-skip+1, 0) + } + out := core.New(core.Float, outLen) + vals := out.RawFloats() + // FilterApply keeps the input dtype, so a float32 series comes + // back with a float32 payload: read it through widenFloats (the + // FloatAt widening the rest of the package reads with) instead of + // the nil float64 payload. + src := widenFloats(filtered) + for i := range outLen { + vals[i] = src[skip*factor+delay+i*factor] + } + return out, nil +} + +// Resample converts the sample rate by the rational factor up/down: +// the series is filtered on the up-sampled grid by a Kaiser tapered +// FIR at the tighter of the two Nyquists with the up gain folded in, +// and every down-th sample of the compensated result is kept. Up and +// down must be at least 1 and not both 1; taps sets the kernel length +// on the up-sampled grid, a non-positive value meaning +// 32·max(up, down)+1, the length that fits the anti-alias transition. +// At the series ends the filter reads past the data: the up-sampled +// grid counts as zeros outside the input, so the head and the tail +// after the delay compensation are the sums over those zeros, not +// wrapped or skipped samples. +func Resample(data *core.Array, up, down, taps int) (*core.Array, error) { + const name = "Resample" + if data.NDim() != 1 { + return nil, base.Errf("%s: the series must be a vector, got shape %s", name, base.ShapeText(data.Shape())) + } + if data.Dtype() == core.Complex { + return nil, base.Errf("%s: complex series are not supported", name) + } + if up < 1 || down < 1 || (up == 1 && down == 1) { + return nil, base.Errf("%s: the rate change up/down must not be the identity %d/%d", name, up, down) + } + n := data.Len() + if n == 0 { + return nil, base.Errf("%s: the series must not be empty", name) + } + if taps <= 0 { + taps = 32*max(up, down) + 1 + } + if taps%2 == 0 { + taps++ + } + if taps > up*n { + return nil, base.Errf("%s: %d taps against %d up-sampled samples leaves nothing after the filter delay", name, taps, up*n) + } + // The source is read through widenFloats: a float32 or int series + // carries no float64 payload, and widenFloats is the exact FloatAt + // widening the rest of the package reads with. + src := widenFloats(data) + // Cutoff on the up-sampled grid: the smaller Nyquist of input and + // output, in cycles per up-sampled sample. + fc := 0.5 * math.Min(1/float64(up), 1/float64(down)) + h := kaiserSinc(taps, fc) + delay := (taps - 1) / 2 + outLen := (n*up + down - 1) / down + out := core.New(core.Float, outLen) + vals := out.RawFloats() + for m := range outLen { + // Output sample m is up-sample index m·down; the filter + // centred there sums the inputs within its support. Only k + // with centre − k·up inside [0, taps) contributes, so the loop + // runs over that window alone, in the same ascending order. + centre := delay + m*down + kmin := max(0, (centre-taps+1+up-1)/up) + kmax := min(n-1, centre/up) + total := 0.0 + for k := kmin; k <= kmax; k++ { + total += h[centre-k*up] * src[k] + } + // j runs over the kernel, so the compensated sample sits at + // up-index centre − delay; scale by up for the zero-stuffed + // grid's unit gain. + vals[m] = float64(up) * total + } + return out, nil +} + +// ResampleFourier resamples a series to exactly size samples by the +// Fourier (band-limited) definition: the spectrum's bins are kept, +// padded with zeros or truncated at the fold, and scaled so the +// result carries the same tone amplitudes as the input. It is exact +// for series band-limited below the new Nyquist and treats the input +// as one period, the same convention AnalyticSignal uses. Signals +// with energy past the new Nyquist lose it, which is the brick-wall +// anti-alias this resample implies. +func ResampleFourier(data *core.Array, size int) (*core.Array, error) { + const name = "ResampleFourier" + if data.NDim() != 1 { + return nil, base.Errf("%s: the series must be a vector, got shape %s", name, base.ShapeText(data.Shape())) + } + n := data.Len() + if n == 0 { + return nil, base.Errf("%s: the series must not be empty", name) + } + if size < 1 { + return nil, base.Errf("%s: the target size must be at least 1, got %d", name, size) + } + spec, err := RFFT(data) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + oldHalf := spec.Len() + newHalf := size/2 + 1 + scaled := make([]complex128, newHalf) + ratio := float64(size) / float64(n) + for k := range min(oldHalf, newHalf) { + scaled[k] = complex(ratio, 0) * spec.RawComplexes()[k] + } + target, err := core.FromComplexes(scaled, newHalf) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + out, err := IRFFT(target, size) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + return out, nil +} diff --git a/signal/resample_test.go b/signal/resample_test.go new file mode 100644 index 0000000..6ec0052 --- /dev/null +++ b/signal/resample_test.go @@ -0,0 +1,184 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// resampleTone builds n samples of a·cos(2π·cycles·i/n) and its +// analytic continuation for comparisons. +func resampleTone(t *testing.T, a float64, cycles, n int) *core.Array { + t.Helper() + vals := make([]float64, n) + for i := range n { + vals[i] = a * math.Cos(2*math.Pi*float64(cycles)*float64(i)/float64(n)) + } + out, err := core.FromFloats(vals, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return out +} + +// TestDecimateTone checks that an in-band tone survives decimation +// with its amplitude and phase, sampled on the new grid. Output +// starts once the filter has full context, at input index +// ceil(delay/factor)·factor + delay. +func TestDecimateTone(t *testing.T) { + const cycles, n, factor = 4, 256, 4 + out, err := Decimate(resampleTone(t, 1, cycles, n), factor, 0) + if err != nil { + t.Fatalf("Decimate: %v", err) + } + // taps = 32·factor+1 by default, delay = 16·factor, so the first + // kept compensated index is 64 and sample i sits at input index + // 64 + i·factor. + const first = 64 + vals := out.RawFloats()[:out.Len()] + for i, got := range vals { + want := math.Cos(2 * math.Pi * float64(cycles) * float64(first+i*factor) / float64(n)) + if math.Abs(got-want) > 0.01 { + t.Fatalf("sample %d = %.6g, want %.6g", i, got, want) + } + } +} + +// TestDecimateAliasedTone checks the anti-alias job: a tone past the +// new Nyquist must come out suppressed, not folded down. The first +// and last taps samples carry filter transients and are skipped. +func TestDecimateAliasedTone(t *testing.T) { + // 44 cycles over 256 samples sits past the new Nyquist of 32 when + // decimating by 4; the stopband is around 80 dB, so a small leak + // is honest, a fold-back to a visible tone is not. + out, err := Decimate(resampleTone(t, 1, 44, 256), 4, 0) + if err != nil { + t.Fatalf("Decimate: %v", err) + } + vals := out.RawFloats()[:out.Len()] + peak := 0.0 + for i, v := range vals { + if i*4 < 128 || i*4 > 256-128 { + continue + } + if a := math.Abs(v); a > peak { + peak = a + } + } + if peak > 0.02 { + t.Fatalf("aliased tone survived at peak %.4g, the anti-alias filter leaked", peak) + } +} + +// TestResampleUpThenDown checks the rational path on a tone: up by 3 +// lands on a finer grid with the tone intact, and the round trip back +// down recovers the input samples. +func TestResampleUpThenDown(t *testing.T) { + const cycles, n = 5, 60 + in := resampleTone(t, 1, cycles, n) + up, err := Resample(in, 3, 1, 0) + if err != nil { + t.Fatalf("Resample up: %v", err) + } + if up.Len() != 3*n { + t.Fatalf("up-sampled length %d, want %d", up.Len(), 3*n) + } + vals := up.RawFloats()[:up.Len()] + // Skip the edges, where the filter context is partial. + for i := 24; i < up.Len()-24; i++ { + want := math.Cos(2 * math.Pi * float64(cycles) * float64(i) / float64(3*n)) + if math.Abs(vals[i]-want) > 0.02 { + t.Fatalf("up sample %d = %.6g, want %.6g", i, vals[i], want) + } + } + down, err := Resample(up, 1, 3, 0) + if err != nil { + t.Fatalf("Resample down: %v", err) + } + if down.Len() != n { + t.Fatalf("round-trip length %d, want %d", down.Len(), n) + } + inVals := in.RawFloats()[:n] + dVals := down.RawFloats()[:down.Len()] + for i := 12; i < n-12; i++ { + if math.Abs(inVals[i]-dVals[i]) > 0.02 { + t.Fatalf("round trip sample %d = %.6g, want %.6g", i, dVals[i], inVals[i]) + } + } +} + +// TestResampleFourierTone checks the exact band-limited resample: a +// tone stays the same amplitude on a doubled grid, point for point. +func TestResampleFourierTone(t *testing.T) { + const cycles, n, size = 4, 64, 192 + out, err := ResampleFourier(resampleTone(t, 0.75, cycles, n), size) + if err != nil { + t.Fatalf("ResampleFourier: %v", err) + } + if out.Len() != size { + t.Fatalf("length %d, want %d", out.Len(), size) + } + vals := out.RawFloats()[:out.Len()] + for i, got := range vals { + want := 0.75 * math.Cos(2*math.Pi*float64(cycles)*float64(i)/float64(size)) + if math.Abs(got-want) > 1e-12 { + t.Fatalf("sample %d = %.12g, want %.12g", i, got, want) + } + } +} + +// TestResampleFourierDown checks the truncating direction: a mixed +// two-tone series resampled to a third of its length keeps the low +// tone and drops the one past the new Nyquist. +func TestResampleFourierDown(t *testing.T) { + const n = 96 + vals := make([]float64, n) + for i := range n { + low := math.Cos(2 * math.Pi * 3 * float64(i) / float64(n)) + high := 0.5 * math.Cos(2*math.Pi*30*float64(i)/float64(n)) + vals[i] = low + high + } + in, _ := core.FromFloats(vals, n) + out, err := ResampleFourier(in, n/3) + if err != nil { + t.Fatalf("ResampleFourier: %v", err) + } + outVals := out.RawFloats()[:out.Len()] + for i, got := range outVals { + want := math.Cos(2 * math.Pi * 3 * float64(i) / float64(n/3)) + if math.Abs(got-want) > 1e-12 { + t.Fatalf("sample %d = %.12g, want %.12g", i, got, want) + } + } +} + +// TestResampleRefusals checks the shape and argument guards. +func TestResampleRefusals(t *testing.T) { + bad := core.New(core.Float, 2, 2) + if _, err := Decimate(bad, 2, 0); err == nil { + t.Fatal("matrix accepted by Decimate") + } + if _, err := Resample(bad, 2, 1, 0); err == nil { + t.Fatal("matrix accepted by Resample") + } + if _, err := ResampleFourier(bad, 8); err == nil { + t.Fatal("matrix accepted by ResampleFourier") + } + one := core.New(core.Float, 16) + if _, err := Decimate(one, 1, 0); err == nil { + t.Fatal("identity factor accepted by Decimate") + } + if _, err := Resample(one, 1, 1, 0); err == nil { + t.Fatal("identity rate accepted by Resample") + } + if _, err := Resample(one, 3, 1, 200); err == nil { + t.Fatal("taps larger than the series accepted by Resample") + } + if _, err := ResampleFourier(one, 0); err == nil { + t.Fatal("zero size accepted by ResampleFourier") + } +} diff --git a/signal/stencil.go b/signal/stencil.go new file mode 100644 index 0000000..972caa8 --- /dev/null +++ b/signal/stencil.go @@ -0,0 +1,218 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import "slices" + +import "sourcedock.dev/petrbalvin/tensor/internal/engine" + +// Numerically robust scalar reduction and finite-difference stencils +// for grid data. SumKahan applies Kahan's compensated summation to the +// whole array: the same sum `Sum` computes, but with the running +// rounding error carried alongside, so long chains of similar +// magnitudes stop bleeding digits. The finite-difference stencils +// produce the central-difference gradient and Laplacian of grid data, +// the raw material of explicit PDE steppers. + +// SumKahan returns the sum of all elements by Kahan's compensated +// summation. For arrays with mixed magnitudes the result carries more +// of the small terms than a plain left-to-right accumulation. +// core.Complex arrays are refused. +func SumKahan(a *core.Array) (float64, error) { + if a.Dtype() == core.Complex { + return 0, base.Errf("SumKahan: complex arrays are not supported") + } + src := widenFloats(a) + var sum, comp float64 + // Bounded by the element count: a rebased view's aliased payload + // runs past its extent. + for i := range a.Len() { + v := src[i] + v -= comp + t := sum + v + comp = (t - sum) - v + sum = t + } + return sum, nil +} + +// Gradient1D returns the central-difference derivative of a 1-D +// uniform grid signal y with spacing dx. Interior points use the +// second-order central stencil; the endpoints fall back to the +// first-order one-sided stencil, so the result has the same length as +// the input. +func Gradient1D(y *core.Array, dx float64) (*core.Array, error) { + if y.NDim() != 1 { + return nil, base.Errf("Gradient1D: needs a rank-1 signal, got shape %s", base.ShapeText(y.Shape())) + } + if y.Dtype() == core.Complex { + return nil, base.Errf("Gradient1D: complex arrays are not supported") + } + n := y.Len() + if n < 2 { + return nil, base.Errf("Gradient1D: needs at least 2 points, got %d", n) + } + if dx == 0 { + return nil, base.Errf("Gradient1D: spacing must not be zero") + } + src := widenFloats(y) + out := core.New(core.Float, n) + dst := out.RawFloats() + dst[0] = (src[1] - src[0]) / dx + for i := 1; i < n-1; i++ { + dst[i] = (src[i+1] - src[i-1]) / (2 * dx) + } + dst[n-1] = (src[n-1] - src[n-2]) / dx + return out, nil +} + +// Laplacian returns the second-derivative Laplacian of a grid signal. +// 1-D input takes one spacing, 2-D input two (dx, dy for a row-major +// (rows, cols) grid, 5-point stencil), 3-D input three (7-point +// stencil). Boundary points copy their nearest interior value; the +// caller owns the boundary condition. A rank outside 1..3, a zero +// spacing, a missing spacing, or a degenerate extent is an error: +// every axis needs at least 2 points, and every axis of a 2-D or 3-D +// grid at least 3, so the interior stencil and the boundary copies +// have somewhere to stand. +func Laplacian(a *core.Array, spacings ...float64) (*core.Array, error) { + const name = "Laplacian" + if a.Dtype() == core.Complex { + return nil, base.Errf("%s: complex arrays are not supported", name) + } + if len(spacings) < a.NDim() { + return nil, base.Errf("%s: needs %d spacings for rank %d, got %d", + name, a.NDim(), a.NDim(), len(spacings)) + } + if slices.Contains(spacings, 0) { + return nil, base.Errf("%s: spacing must not be zero", name) + } + switch a.NDim() { + case 1: + if a.Len() < 2 { + return nil, base.Errf("%s: needs at least 2 points, got %d", name, a.Len()) + } + case 2, 3: + for i, d := range a.Shape() { + if d < 3 { + return nil, base.Errf("%s: axis %d has extent %d, the stencil needs at least 3 points", + name, i, d) + } + } + default: + return nil, base.Errf("%s: needs a 1-D, 2-D or 3-D array, got rank %d", name, a.NDim()) + } + switch a.NDim() { + case 1: + return laplacian1D(a, spacings[0]), nil + case 2: + return laplacian2D(a, spacings[0], spacings[1]), nil + default: + return laplacian3D(a, spacings[0], spacings[1], spacings[2]), nil + } +} + +// laplacian1D applies the 3-point second-difference stencil along the +// single axis. Interior: (y[i+1] − 2y[i] + y[i−1]) / dx². For n = 2 +// there is no interior point, and the edge copy leaves both entries +// zero. +func laplacian1D(y *core.Array, dx float64) *core.Array { + n := y.Len() + out := core.New(core.Float, n) + inv := 1 / (dx * dx) + src := widenFloats(y) + dst := out.RawFloats() + for i := 1; i < n-1; i++ { + dst[i] = (src[i+1] - 2*src[i] + src[i-1]) * inv + } + dst[0] = dst[1] + dst[n-1] = dst[n-2] + return out +} + +// laplacian2D applies the 5-point stencil on a row-major grid of +// shape (ny, nx) with spacings dx, dy. +func laplacian2D(y *core.Array, dx, dy float64) *core.Array { + ny, nx := y.Shape()[0], y.Shape()[1] + out := core.New(core.Float, ny, nx) + invx := 1 / (dx * dx) + invy := 1 / (dy * dy) + src := widenFloats(y) + dst := out.RawFloats() + engine.ParallelMin(ny, workFloorFor(nx*6), func(s, e int) { + for j := s; j < e; j++ { + for i := 1; i < nx-1; i++ { + idx := j*nx + i + c := src[idx] + acc := (src[idx+1] + src[idx-1] - 2*c) * invx + if j > 0 && j < ny-1 { + acc += (src[(j-1)*nx+i] + src[(j+1)*nx+i] - 2*c) * invy + } + dst[idx] = acc + } + dst[j*nx] = dst[j*nx+1] + dst[j*nx+nx-1] = dst[j*nx+nx-2] + } + }) + for i := range nx { + dst[i] = dst[nx+i] + dst[(ny-1)*nx+i] = dst[(ny-2)*nx+i] + } + return out +} + +// laplacian3D applies the 7-point stencil on a row-major grid of +// shape (nz, ny, nx) with spacings dx, dy, dz. Boundary planes copy +// their nearest interior plane, the caller owns the boundary +// condition. +func laplacian3D(y *core.Array, dx, dy, dz float64) *core.Array { + nz, ny, nx := y.Shape()[0], y.Shape()[1], y.Shape()[2] + out := core.New(core.Float, nz, ny, nx) + invx := 1 / (dx * dx) + invy := 1 / (dy * dy) + invz := 1 / (dz * dz) + src := widenFloats(y) + dst := out.RawFloats() + engine.ParallelMin(nz, workFloorFor(ny*nx*9), func(zs, ze int) { + for k := zs; k < ze; k++ { + for j := range ny { + for i := 1; i < nx-1; i++ { + idx := (k*ny+j)*nx + i + c := src[idx] + acc := (src[idx+1] + src[idx-1] - 2*c) * invx + if j > 0 && j < ny-1 { + acc += (src[(k*ny+j-1)*nx+i] + src[(k*ny+j+1)*nx+i] - 2*c) * invy + } + if k > 0 && k < nz-1 { + acc += (src[((k-1)*ny+j)*nx+i] + src[((k+1)*ny+j)*nx+i] - 2*c) * invz + } + dst[idx] = acc + } + // The x boundaries copy their nearest interior value. + dst[(k*ny+j)*nx] = dst[(k*ny+j)*nx+1] + dst[(k*ny+j)*nx+nx-1] = dst[(k*ny+j)*nx+nx-2] + } + } + }) + // The y boundaries copy their nearest interior row. + for k := range nz { + for i := range nx { + dst[k*ny*nx+i] = dst[(k*ny+1)*nx+i] + dst[((k+1)*ny-1)*nx+i] = dst[((k+1)*ny-2)*nx+i] + } + } + // The z boundaries copy their nearest interior plane. + for j := range ny { + for i := range nx { + dst[j*nx+i] = dst[(ny+j)*nx+i] + dst[((nz-1)*ny+j)*nx+i] = dst[((nz-2)*ny+j)*nx+i] + } + } + return out +} diff --git a/signal/stencil_test.go b/signal/stencil_test.go new file mode 100644 index 0000000..366c620 --- /dev/null +++ b/signal/stencil_test.go @@ -0,0 +1,256 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestSumKahanRecoversSmallTerms checks the compensation: small terms +// between huge ones survive, where a plain left-to-right sum loses +// them entirely (the naive sum of this sequence is 0). +func TestSumKahanRecoversSmallTerms(t *testing.T) { + vals := mustFloats(t, []float64{1e16, 1, 2, -1e16}) + got, err := SumKahan(vals) + if err != nil { + t.Fatalf("SumKahan: %v", err) + } + if got != 4 { + t.Fatalf("SumKahan = %g, want 4", got) + } +} + +// TestGradient1DDifferentiatesLinesExactly: both the central stencil +// and the one-sided endpoints are exact on linear signals. +func TestGradient1DDifferentiatesLinesExactly(t *testing.T) { + const n = 16 + const dx = 0.25 + y := make([]float64, n) + for i := range n { + y[i] = 3*float64(i)*dx + 7 + } + got, err := Gradient1D(mustFloats(t, y, n), dx) + if err != nil { + t.Fatalf("Gradient1D: %v", err) + } + for i := range n { + if math.Abs(got.FloatAt(i)-3) > 1e-12 { + t.Fatalf("gradient[%d] = %.14g, want 3", i, got.FloatAt(i)) + } + } +} + +// TestGradient1DKnownDifferentiatesQuadratics: the central interior +// stencil is exact on quadratics; the first-order endpoint stencils +// carry an O(dx) offset. +func TestGradient1DKnownDifferentiatesQuadratics(t *testing.T) { + const n = 16 + const dx = 0.25 + y := make([]float64, n) + for i := range n { + x := float64(i) * dx + y[i] = x * x + } + got, err := Gradient1D(mustFloats(t, y, n), dx) + if err != nil { + t.Fatalf("Gradient1D: %v", err) + } + for i := 1; i < n-1; i++ { + want := 2 * float64(i) * dx + if math.Abs(got.FloatAt(i)-want) > 1e-12 { + t.Fatalf("gradient[%d] = %.14g, want %.14g", i, got.FloatAt(i), want) + } + } + if math.Abs(got.FloatAt(0)-dx) > 1e-12 { + t.Fatalf("left endpoint = %.14g, want the one-sided 2x+h = %.14g", got.FloatAt(0), dx) + } +} + +// TestGradient1DErrors pins the validation contract. +func TestGradient1DErrors(t *testing.T) { + if _, err := Gradient1D(mustFloats(t, []float64{1}), 0.1); err == nil { + t.Fatal("expected an error for a single-point signal") + } + if _, err := Gradient1D(mustFloats(t, []float64{1, 2, 3}), 0); err == nil { + t.Fatal("expected an error for a zero spacing") + } + rank2, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) + if _, err := Gradient1D(rank2, 0.1); err == nil { + t.Fatal("expected an error for a rank-2 signal") + } +} + +// TestLaplacian1DKnownDifferentiatesSines checks the 3-point stencil +// against −sin on a full period and the boundary copy at both ends. +func TestLaplacian1DKnownDifferentiatesSines(t *testing.T) { + const n = 64 + const dx = 2 * math.Pi / n + y := make([]float64, n) + for i := range n { + y[i] = math.Sin(float64(i) * dx) + } + got, err := Laplacian(mustFloats(t, y, n), dx) + if err != nil { + t.Fatalf("Laplacian: %v", err) + } + // The interior matches −sin; the boundary copies the neighbour, so + // it is checked separately below. + for i := 1; i < n-1; i++ { + want := -math.Sin(float64(i) * dx) + if math.Abs(got.FloatAt(i)-want) > 1e-3 { + t.Fatalf("laplacian[%d] = %.6g, want %.6g", i, got.FloatAt(i), want) + } + } + // The boundaries copy their nearest interior value. + if got.FloatAt(0) != got.FloatAt(1) || got.FloatAt(n-1) != got.FloatAt(n-2) { + t.Fatal("the 1-D boundaries did not copy their interior values") + } +} + +// TestLaplacian2DKnownDifferentiatesSines checks the 5-point stencil +// against −2·sin x·cos y and the boundary copy on every edge. +func TestLaplacian2DKnownDifferentiatesSines(t *testing.T) { + const n = 32 + const d = 2 * math.Pi / n + grid := make([]float64, n*n) + for r := range n { + for c := range n { + grid[r*n+c] = math.Sin(float64(c)*d) * math.Cos(float64(r)*d) + } + } + got, err := Laplacian(mustFloats(t, grid, n, n), d, d) + if err != nil { + t.Fatalf("Laplacian: %v", err) + } + // The interior matches −2·sin x·cos y to the stencil's own + // truncation error 2h²/12 ≈ 0.0064 at h = 2π/32; the boundary + // copies its nearest interior value, checked separately below. + for r := 1; r < n-1; r++ { + for c := 1; c < n-1; c++ { + want := -2 * math.Sin(float64(c)*d) * math.Cos(float64(r)*d) + if math.Abs(got.FloatAt(r*n+c)-want) > 8e-3 { + t.Fatalf("laplacian[%d,%d] = %.6g, want %.6g", r, c, got.FloatAt(r*n+c), want) + } + } + } + // Every boundary point must equal its nearest interior neighbour. + for i := range n { + if got.FloatAt(i) != got.FloatAt(n+i) { + t.Fatalf("top row %d did not copy the row below", i) + } + if got.FloatAt((n-1)*n+i) != got.FloatAt((n-2)*n+i) { + t.Fatalf("bottom row %d did not copy the row above", i) + } + if got.FloatAt(i*n) != got.FloatAt(i*n+1) { + t.Fatalf("left column %d did not copy its interior neighbour", i) + } + if got.FloatAt(i*n+n-1) != got.FloatAt(i*n+n-2) { + t.Fatalf("right column %d did not copy its interior neighbour", i) + } + } +} + +// TestLaplacian3DKnownDifferentiatesSines checks the 7-point stencil +// against −3·sin x·sin y·sin z, and pins the boundary behaviour: all +// six faces copy their nearest interior plane (the doc's promise). +func TestLaplacian3DKnownDifferentiatesSines(t *testing.T) { + const n = 16 + const d = 2 * math.Pi / n + at := func(z, y, x int) float64 { + return math.Sin(float64(x)*d) * math.Sin(float64(y)*d) * math.Sin(float64(z)*d) + } + grid := make([]float64, n*n*n) + for k := range n { + for j := range n { + for i := range n { + grid[(k*n+j)*n+i] = at(k, j, i) + } + } + } + got, err := Laplacian(mustFloats(t, grid, n, n, n), d, d, d) + if err != nil { + t.Fatalf("Laplacian: %v", err) + } + // Interior accuracy first; the six boundary faces are copies, so + // they are excluded from the truncation-error bound and checked + // against their interior neighbours below. + worst := 0.0 + for k := 1; k < n-1; k++ { + for j := 1; j < n-1; j++ { + for i := 1; i < n-1; i++ { + want := -3 * at(k, j, i) + d := math.Abs(got.FloatAt((k*n+j)*n+i) - want) + if d > worst { + worst = d + } + } + } + } + if worst > 5e-2 { + // The bound is the stencil's own truncation error 3h²/12 ≈ 0.039 + // at h = 2π/16; a wrong stencil is orders of magnitude worse. + t.Fatalf("7-point stencil error %.3g, want under 5e-2", worst) + } + // All six boundary faces copy their nearest interior plane. + for j := range n { + for i := range n { + if got.FloatAt(j*n+i) != got.FloatAt((n+j)*n+i) { + t.Fatalf("front plane (%d,%d) did not copy the plane behind it", j, i) + } + if got.FloatAt(((n-1)*n+j)*n+i) != got.FloatAt(((n-2)*n+j)*n+i) { + t.Fatalf("back plane (%d,%d) did not copy the plane before it", j, i) + } + } + } + for k := range n { + for i := range n { + if got.FloatAt((k*n)*n+i) != got.FloatAt((k*n+1)*n+i) { + t.Fatalf("top face (%d,%d) did not copy the row below", k, i) + } + if got.FloatAt(((k+1)*n-1)*n+i) != got.FloatAt(((k+1)*n-2)*n+i) { + t.Fatalf("bottom face (%d,%d) did not copy the row above", k, i) + } + } + for j := range n { + if got.FloatAt((k*n+j)*n) != got.FloatAt((k*n+j)*n+1) { + t.Fatalf("left face (%d,%d) did not copy its interior neighbour", k, j) + } + if got.FloatAt((k*n+j)*n+n-1) != got.FloatAt((k*n+j)*n+n-2) { + t.Fatalf("right face (%d,%d) did not copy its interior neighbour", k, j) + } + } + } +} + +// TestLaplacianDegenerateErrors pins the extent contract: every rank-1 +// signal needs at least 2 points and every axis of a rank-2 or rank-3 +// grid at least 3, so the stencil and the boundary copies have +// somewhere to stand. +func TestLaplacianDegenerateErrors(t *testing.T) { + if _, err := Laplacian(mustFloats(t, []float64{1}), 0.1); err == nil { + t.Fatal("expected an error for a single-point 1-D signal") + } + tiny2D, _ := core.FromFloats(make([]float64, 16), 2, 8) + if _, err := Laplacian(tiny2D, 0.1, 0.1); err == nil { + t.Fatal("expected an error for a 2-D grid with an extent of 2") + } + tiny3D, _ := core.FromFloats(make([]float64, 32), 2, 4, 4) + if _, err := Laplacian(tiny3D, 0.1, 0.1, 0.1); err == nil { + t.Fatal("expected an error for a 3-D grid with an extent of 2") + } + twoD, _ := core.FromFloats(make([]float64, 9), 3, 3) + if _, err := Laplacian(twoD, 0.1); err == nil { + t.Fatal("expected an error for a missing spacing") + } + if _, err := Laplacian(mustFloats(t, []float64{1, 2, 3}), 0); err == nil { + t.Fatal("expected an error for a zero spacing") + } + rank4, _ := core.FromFloats(make([]float64, 16), 2, 2, 2, 2) + if _, err := Laplacian(rank4, 0.1, 0.1, 0.1, 0.1); err == nil { + t.Fatal("expected an error for a rank-4 array") + } +} diff --git a/signal/stft.go b/signal/stft.go new file mode 100644 index 0000000..2129d94 --- /dev/null +++ b/signal/stft.go @@ -0,0 +1,158 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The short-time Fourier transform: the signal is cut into +// overlapping segments, each segment is windowed and transformed, and +// the frames stack into a time-by-frequency picture. The spectrogram +// is the magnitude squared of it, scaled per frame exactly as Welch's +// estimate scales its periodograms, so averaging a spectrogram over +// its frames reproduces WelchPSD on the same segmentation. + +// STFTOptions tunes the transform. Segment is the length of one +// frame in samples, Overlap the samples two neighbouring frames share +// (segment/2 is the usual choice) and Window names the taper: "hann", +// "hamming" or "box", with "hann" implied by the zero value. +type STFTOptions struct { + Segment int + Overlap int + Window string +} + +// framesOf computes the frame count of a signal of n samples under +// the segment and overlap geometry, the count Welch's estimate uses. +func framesOf(n, segment, overlap int) int { + return (n - overlap) / (segment - overlap) +} + +// STFT transforms the vector x frame by frame and returns the +// complex frames as a (frames × segment) array in row order: row t +// holds the spectrum of the t-th frame, bin k of frame t is the +// transform of x[t·hop : t·hop+segment] under the window, with hop = +// segment − overlap. The Fourier definition treats each windowed +// segment as one period, which is the standard convention here. +func STFT(x *core.Array, opts STFTOptions) (*core.Array, error) { + const name = "STFT" + if x.NDim() != 1 { + return nil, base.Errf("%s: the signal must be a vector, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, base.Errf("%s: complex arrays are not supported", name) + } + n := x.Len() + if opts.Segment < 2 || opts.Segment > n { + return nil, base.Errf("%s: segment must be between 2 and %d, got %d", name, n, opts.Segment) + } + if opts.Overlap < 0 || opts.Overlap >= opts.Segment { + return nil, base.Errf("%s: overlap must be in [0, %d), got %d", name, opts.Segment, opts.Overlap) + } + window := opts.Window + if window == "" { + window = "hann" + } + taper, err := windowTaper(window, opts.Segment) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + hop := opts.Segment - opts.Overlap + frames := framesOf(n, opts.Segment, opts.Overlap) + if frames < 1 { + return nil, base.Errf("%s: %d samples fill only one %d-sample segment with overlap %d", + name, n, opts.Segment, opts.Overlap) + } + src := widenFloats(x) + segment := opts.Segment + out, oerr := core.Zeros(core.Complex, frames, segment) + if oerr != nil { + return nil, oerr + } + // A frame's windowed samples go straight into its output row and are + // transformed there: the row is a private, fully overwritten buffer, + // so the copy a separate scratch segment needed is one pass fewer and + // the bits land exactly where the copy would have put them. + rows := out.RawComplexes()[:out.Len()] + engine.ParallelMin(frames, welchMinSegmentsPerWorker, func(fs, fe int) { + for f := fs; f < fe; f++ { + start := f * hop + z := rows[f*segment : (f+1)*segment] + for i := range segment { + z[i] = complex(src[start+i]*taper[i], 0) + } + transform(z, -1) + } + }) + return out, nil +} + +// Spectrogram returns the one-sided power spectrogram of the vector x +// sampled at fs hertz as a (frames × segment/2+1) array: every frame +// carries the periodogram of its windowed segment with the same +// one-sided doubling and window-power normalisation Welch's estimate +// applies, in units of x²/Hz. Averaging the frames reproduces +// WelchPSD on the same geometry. +func Spectrogram(x *core.Array, fs float64, opts STFTOptions) (*core.Array, error) { + const name = "Spectrogram" + if !(fs > 0) || math.IsInf(fs, 0) { + return nil, base.Errf("%s: fs must be positive and finite, got %g", name, fs) + } + // The segment geometry is checked before the taper is built: the + // taper allocates Segment samples, so a hostile segment must be + // refused here rather than after the allocation. STFT repeats the + // same checks on its own terms. + n := x.Len() + if opts.Segment < 2 || opts.Segment > n { + return nil, base.Errf("%s: segment must be between 2 and %d, got %d", name, n, opts.Segment) + } + if opts.Overlap < 0 || opts.Overlap >= opts.Segment { + return nil, base.Errf("%s: overlap must be in [0, %d), got %d", name, opts.Segment, opts.Overlap) + } + window := opts.Window + if window == "" { + window = "hann" + } + taper, err := windowTaper(window, opts.Segment) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + wPower := 0.0 + for i := range taper { + wPower += taper[i] * taper[i] + } + if wPower == 0 { + return nil, base.Errf("%s: the window has zero power", name) + } + framesSpec, err := STFT(x, opts) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + frames := framesSpec.Shape()[0] + segment := opts.Segment + bins := segment/2 + 1 + scale := 1 / (fs * wPower) + out, oerr := core.Zeros(core.Float, frames, bins) + if oerr != nil { + return nil, oerr + } + spec := framesSpec.RawComplexes()[:framesSpec.Len()] + vals := out.RawFloats() + for f := range frames { + for k := range bins { + mag := base.AbsComplex(spec[f*segment+k]) + p := mag * mag * scale + if k > 0 && k < segment-k { + p *= 2 + } + vals[f*bins+k] = p + } + } + return out, nil +} diff --git a/signal/stft_test.go b/signal/stft_test.go new file mode 100644 index 0000000..a7d7a5a --- /dev/null +++ b/signal/stft_test.go @@ -0,0 +1,128 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// stftTone builds n samples of cos(2π·cycles·i/n). +func stftTone(t *testing.T, cycles, n int) *core.Array { + t.Helper() + vals := make([]float64, n) + for i := range n { + vals[i] = math.Cos(2 * math.Pi * float64(cycles) * float64(i) / float64(n)) + } + a, err := core.FromFloats(vals, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestSTFTFrameSpectra checks the frame geometry on a tone with an +// integer number of cycles per segment under the box window: every +// frame is the same spectrum with its peak in the tone's bin. +func TestSTFTFrameSpectra(t *testing.T) { + const n, segment, cycles = 256, 64, 16 + x, err := STFT(stftTone(t, cycles, n), STFTOptions{Segment: segment, Overlap: 0, Window: "box"}) + if err != nil { + t.Fatalf("STFT: %v", err) + } + if x.Shape()[0] != n/segment || x.Shape()[1] != segment { + t.Fatalf("shape %v, want (%d, %d)", x.Shape(), n/segment, segment) + } + rows := x.RawComplexes()[:x.Len()] + // 16 cycles over 256 samples is 4 cycles per 64-sample segment. + perSegment := cycles * segment / n + for f := range x.Shape()[0] { + peak, peakMag := 0, 0.0 + for k := range segment { + if mag := math.Hypot(real(rows[f*segment+k]), imag(rows[f*segment+k])); mag > peakMag { + peak, peakMag = k, mag + } + } + if peak != perSegment && peak != segment-perSegment { + t.Fatalf("frame %d peaks at bin %d, want %d or %d", f, peak, perSegment, segment-perSegment) + } + } +} + +// TestSpectrogramMatchesWelch pins the scaling contract: averaging +// the spectrogram over its frames must reproduce WelchPSD on the same +// segmentation, because the spectrogram is that estimate unaveraged. +func TestSpectrogramMatchesWelch(t *testing.T) { + const n = 512 + vals := make([]float64, n) + for i := range n { + vals[i] = math.Cos(2*math.Pi*13*float64(i)/float64(n)) + 0.25*math.Sin(2*math.Pi*40*float64(i)/float64(n)) + } + x, err := core.FromFloats(vals, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + const fs, segment, overlap, window = 256.0, 128, 64, "hann" + opts := STFTOptions{Segment: segment, Overlap: overlap, Window: window} + spec, err := Spectrogram(x, fs, opts) + if err != nil { + t.Fatalf("Spectrogram: %v", err) + } + freqs, psd, err := WelchPSD(x, fs, segment, overlap, window) + if err != nil { + t.Fatalf("WelchPSD: %v", err) + } + if freqs.Len() != spec.Shape()[1] { + t.Fatalf("bin counts disagree: spectrogram %d, Welch %d", spec.Shape()[1], freqs.Len()) + } + frames := float64(spec.Shape()[0]) + vals = spec.RawFloats()[:spec.Len()] + bins := spec.Shape()[1] + for k := range bins { + mean := 0.0 + for f := range spec.Shape()[0] { + mean += vals[f*bins+k] + } + mean /= frames + if math.Abs(mean-psd.RawFloats()[k]) > 1e-12+1e-9*math.Abs(psd.RawFloats()[k]) { + t.Fatalf("bin %d: spectrogram mean %.12g vs Welch %.12g", k, mean, psd.RawFloats()[k]) + } + } +} + +// TestSTFTRefusals checks the geometry and window guards. +func TestSTFTRefusals(t *testing.T) { + bad := core.New(core.Float, 2, 2) + if _, err := STFT(bad, STFTOptions{Segment: 2}); err == nil { + t.Fatal("matrix accepted") + } + one := core.New(core.Float, 16) + if _, err := STFT(one, STFTOptions{Segment: 32}); err == nil { + t.Fatal("segment past the signal accepted") + } + if _, err := STFT(one, STFTOptions{Segment: 8, Overlap: 8}); err == nil { + t.Fatal("overlap at the segment length accepted") + } + if _, err := STFT(one, STFTOptions{Segment: 8, Window: "gauss"}); err == nil { + t.Fatal("unknown window accepted") + } + if _, err := Spectrogram(one, -1, STFTOptions{Segment: 8}); err == nil { + t.Fatal("negative fs accepted") + } +} + +// TestSpectrogramRefusesHostileSegment pins the validation order: the +// taper is allocated from Segment, so a hostile segment must be +// refused before any allocation happens, not after it. +func TestSpectrogramRefusesHostileSegment(t *testing.T) { + x := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 8) + for _, segment := range []int{-1, 0, 1, 9} { + if _, err := Spectrogram(x, 8, STFTOptions{Segment: segment}); err == nil || !strings.Contains(err.Error(), "segment") { + t.Errorf("Spectrogram with segment %d: %v", segment, err) + } + } +} diff --git a/signal/transform_cache_pins_test.go b/signal/transform_cache_pins_test.go new file mode 100644 index 0000000..d4e961b --- /dev/null +++ b/signal/transform_cache_pins_test.go @@ -0,0 +1,112 @@ +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The fit cache is keyed by the window and the order together; a key +// that lost the order would answer one order's weights for another. +// Each reference is computed with an empty cache, so a collision +// between the two keys cannot hide behind a shared fit row. +func TestSavitzkyGolayCacheKeyCarriesOrder(t *testing.T) { + n := 64 + vals := make([]float64, n) + for i := range vals { + x := float64(i) / float64(n-1) + vals[i] = x*x*x - 2*x + } + src, err := core.FromFloats(vals, n) + if err != nil { + t.Fatal(err) + } + run := func() (map[int][]float64, error) { + sgFitKeys = map[sgFitKey]*sgFits{} + got := map[int][]float64{} + for _, order := range []int{1, 2, 3} { + out, err := SavitzkyGolay(src, 7, order) + if err != nil { + return nil, err + } + raw := out.RawFloats() + got[order] = append([]float64(nil), raw...) + } + return got, nil + } + want, err := run() + if err != nil { + t.Fatal(err) + } + got, err := run() + if err != nil { + t.Fatal(err) + } + // The orders must differ from one another on a curved signal, + // otherwise the comparison below pins nothing. + for a := 1; a <= 2; a++ { + for i := range want[a] { + if want[a][i] != want[a+1][i] { + break + } + if i == len(want[a])-1 { + t.Fatalf("orders %d and %d agree everywhere; the fixture does not separate them", a, a+1) + } + } + } + for order := 1; order <= 3; order++ { + for i := range want[order] { + if want[order][i] != got[order][i] { + t.Fatalf("order %d, sample %d: cached run answered %v, fresh-cache run answered %v", order, i, got[order][i], want[order][i]) + } + } + } +} + +// Sweeping more scales than the cache's entry cap must leave every +// answer unchanged and the cache bounded. +func TestCWTSpectrumCacheCap(t *testing.T) { + n := 32 + vals := make([]float64, n) + for i := range vals { + vals[i] = math.Sin(2*math.Pi*float64(i)/8) + 0.25*math.Sin(2*math.Pi*float64(i)/3) + } + src, err := core.FromFloats(vals, n) + if err != nil { + t.Fatal(err) + } + const scales = cwtSpectrumCacheMax + 8 + scaleList := make([]float64, scales) + for k := range scaleList { + scaleList[k] = 2 + float64(k)*0.5 + } + build := func() [][]complex128 { + cwtSpectrumKeys = map[cwtSpectrumKey][]complex128{} + out, err := CWT(src, Morlet, scaleList, 1) + if err != nil { + t.Fatalf("CWT: %v", err) + } + raw, err := out.ComplexValues("CWT") + if err != nil { + t.Fatalf("ComplexValues: %v", err) + } + rows := make([][]complex128, len(scaleList)) + for k := range rows { + rows[k] = append([]complex128(nil), raw[k*n:(k+1)*n]...) + } + return rows + } + want := build() + got := build() + if len(cwtSpectrumKeys) > cwtSpectrumCacheMax { + t.Fatalf("the spectrum cache holds %d entries over the cap %d", len(cwtSpectrumKeys), cwtSpectrumCacheMax) + } + for k := range want { + for i := range want[k] { + if want[k][i] != got[k][i] { + t.Fatalf("scale %v, sample %d: %v vs %v", scaleList[k], i, got[k][i], want[k][i]) + } + } + } +} diff --git a/signal/wave_bench_test.go b/signal/wave_bench_test.go new file mode 100644 index 0000000..51a855b --- /dev/null +++ b/signal/wave_bench_test.go @@ -0,0 +1,446 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "math/big" + "sync" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Accuracy cross-check for the chirp-z line transforms the Poisson +// solves run, against a reference that is not the implementation under +// test: the defining trigonometric sums evaluated in 256-bit floating +// point, with π from Machin's formula at the same precision. The +// padded full-length route the solves previously ran survives here as +// the public DST/DCT pair, untouched, so both routes are measured +// against the same exact reference. + +const refPrec = 256 + +// refPiFloor is the magnitude below which a series term is past the +// reference's own precision, so every series here stops on a proven +// bound rather than on an exact zero the wide exponent range may never +// produce. +var refPiFloor = new(big.Float).SetPrec(uint(refPrec)).SetMantExp(big.NewFloat(1).SetPrec(uint(refPrec)), -refPrec-8) + +// refPi returns π at refPrec bits, computed once per test binary. +var refPi = sync.OnceValue(func() *big.Float { + prec := uint(refPrec) + atanInv := func(x int64) *big.Float { + xb := new(big.Float).SetPrec(prec).SetInt64(x) + term := new(big.Float).SetPrec(prec).Quo(big.NewFloat(1).SetPrec(prec), xb) + x2 := new(big.Float).SetPrec(prec).Mul(xb, xb) + sum := new(big.Float).SetPrec(prec).Set(term) + for m := 1; m < 4*refPrec; m++ { + term.Quo(term, x2) + d := big.NewFloat(float64(2*m + 1)).SetPrec(prec) + t := new(big.Float).SetPrec(prec).Quo(term, d) + if t.Cmp(refPiFloor) < 0 { + break + } + if m%2 == 0 { + sum.Add(sum, t) + } else { + sum.Sub(sum, t) + } + } + return sum + } + pi := new(big.Float).SetPrec(prec) + pi.Mul(atanInv(5), big.NewFloat(16).SetPrec(prec)) + four := new(big.Float).SetPrec(prec).Mul(atanInv(239), big.NewFloat(4).SetPrec(prec)) + return pi.Sub(pi, four) +}) + +// refSin returns sin(z) at refPrec bits by its Taylor series; callers +// keep |z| ≤ 2π, where the terms z^{2m+1}/(2m+1)! pass the precision +// floor in a few dozen steps. The series term carries the factorial +// through its own recurrence, t_m = t_{m−1}·z²/((2m)(2m+1)), so it +// decays factorially rather than as a bare power. +func refSin(z *big.Float) *big.Float { + prec := uint(refPrec) + term := new(big.Float).SetPrec(prec).Set(z) + sum := new(big.Float).SetPrec(prec).Set(z) + z2 := new(big.Float).SetPrec(prec).Mul(z, z) + floor := new(big.Float).SetPrec(prec).SetMantExp(big.NewFloat(1).SetPrec(prec), -int(refPrec)-8) + for m := 1; m < 4*refPrec; m++ { + term.Mul(term, z2) + term.Quo(term, big.NewFloat(float64(2*m*(2*m+1))).SetPrec(prec)) + // term and the denominator are non-negative, so the term is; + // the floor test is a plain comparison. + if term.Sign() == 0 || term.Cmp(floor) < 0 { + break + } + if m%2 == 0 { + sum.Add(sum, term) + } else { + sum.Sub(sum, term) + } + } + return sum +} + +// refCos returns cos(z) at refPrec bits by its Taylor series, the +// terms carrying the factorial through t_m = t_{m−1}·z²/((2m−1)(2m)). +func refCos(z *big.Float) *big.Float { + prec := uint(refPrec) + term := big.NewFloat(1).SetPrec(prec) + sum := big.NewFloat(1).SetPrec(prec) + z2 := new(big.Float).SetPrec(prec).Mul(z, z) + floor := new(big.Float).SetPrec(prec).SetMantExp(big.NewFloat(1).SetPrec(prec), -int(refPrec)-8) + for m := 1; m < 4*refPrec; m++ { + term.Mul(term, z2) + term.Quo(term, big.NewFloat(float64((2*m-1)*(2*m))).SetPrec(prec)) + if term.Sign() == 0 || term.Cmp(floor) < 0 { + break + } + if m%2 == 0 { + sum.Add(sum, term) + } else { + sum.Sub(sum, term) + } + } + return sum +} + +// refLineTransform evaluates the defining sum of the orthonormal +// DST-I (sine) or DCT-I (cosine) of x at refPrec bits. The trig +// argument ((j+1)(k+1)) mod 2(n+1) is reduced in exact integers +// first, so no series argument exceeds 2π. +func refLineTransform(x []float64, sine bool) []*big.Float { + prec := uint(refPrec) + n := len(x) + pi := refPi() + out := make([]*big.Float, n) + xb := make([]*big.Float, n) + for j := range n { + xb[j] = big.NewFloat(x[j]).SetPrec(prec) + } + if sine { + den := float64(n + 1) + sqrt2 := new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(2).SetPrec(prec)) + norm := new(big.Float).SetPrec(prec).Quo(sqrt2, + new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(den).SetPrec(prec))) + period := 2 * (n + 1) + for k := range n { + sum := new(big.Float).SetPrec(prec) + for j := range n { + r := ((j + 1) * (k + 1)) % period + arg := new(big.Float).SetPrec(prec).Quo( + new(big.Float).SetPrec(prec).Mul(pi, big.NewFloat(float64(r)).SetPrec(prec)), + big.NewFloat(den).SetPrec(prec)) + sum.Add(sum, new(big.Float).SetPrec(prec).Mul(xb[j], refSin(arg))) + } + out[k] = new(big.Float).SetPrec(prec).Mul(norm, sum) + } + return out + } + den := float64(n - 1) + sqrt2 := new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(2).SetPrec(prec)) + half := new(big.Float).SetPrec(prec).Quo(sqrt2, big.NewFloat(2).SetPrec(prec)) + norm := new(big.Float).SetPrec(prec).Quo(sqrt2, + new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(den).SetPrec(prec))) + period := 2 * (n - 1) + wj := make([]*big.Float, n) + for j := range n { + if j == 0 || j == n-1 { + wj[j] = half + } else { + wj[j] = big.NewFloat(1).SetPrec(prec) + } + } + for k := range n { + sum := new(big.Float).SetPrec(prec) + for j := range n { + r := (j * k) % period + arg := new(big.Float).SetPrec(prec).Quo( + new(big.Float).SetPrec(prec).Mul(pi, big.NewFloat(float64(r)).SetPrec(prec)), + big.NewFloat(den).SetPrec(prec)) + term := new(big.Float).SetPrec(prec).Mul(xb[j], wj[j]) + sum.Add(sum, term.Mul(term, refCos(arg))) + } + wk := big.NewFloat(1).SetPrec(prec) + if k == 0 || k == n-1 { + wk = half + } + out[k] = new(big.Float).SetPrec(prec).Mul(norm, + new(big.Float).SetPrec(prec).Mul(wk, sum)) + } + return out +} + +// lineFixture builds a deterministic line of length n with mixed +// magnitudes and signs. +func lineFixture(n int) []float64 { + x := make([]float64, n) + state := uint64(0x9e3779b97f4a7c15) ^ uint64(n) + for i := range n { + state ^= state << 13 + state ^= state >> 7 + state ^= state << 17 + v := float64(int64(state)%2000)/1000 - 1 + x[i] = v * math.Pow(2, float64((i%7)-3)) + } + return x +} + +// legacyLine runs the pre-chirp route for one line: the public DST or +// DCT at kind 1, which carries the padded full-length arithmetic the +// solves used to run per line. +func legacyLine(t *testing.T, x []float64, sine bool) []float64 { + arr, err := core.FromFloats(x, len(x)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + var out *core.Array + if sine { + out, err = DST(arr, 1) + } else { + out, err = DCT(arr, 1) + } + if err != nil { + t.Fatalf("public transform: %v", err) + } + vals := make([]float64, out.Len()) + for i := range vals { + vals[i] = out.FloatAt(i) + } + return vals +} + +// TestPoissonLineTransformAccuracy holds both line-transform routes +// against the exact trigonometric sums at 256 bits: the chirp-z route +// the solves run now, and the padded full-length route the public +// DST/DCT still runs. The verdict is printed for every length and +// family; the assertion keeps the chirp route at the rounding floor +// and never far behind the padded route it replaced. +func TestPoissonLineTransformAccuracy(t *testing.T) { + for _, tc := range []struct { + n int + sine bool + label string + }{ + {8, true, "dst1-8"}, + {31, true, "dst1-31"}, + {64, true, "dst1-64"}, + {8, false, "dct1-8"}, + {31, false, "dct1-31"}, + {64, false, "dct1-64"}, + } { + x := lineFixture(tc.n) + want := refLineTransform(x, tc.sine) + wantMax := 0.0 + for _, w := range want { + f, _ := w.Float64() + wantMax = max(wantMax, math.Abs(f)) + } + if wantMax == 0 { + t.Fatalf("%s: reference is identically zero", tc.label) + } + scale := 0.0 + for _, v := range x { + scale = max(scale, math.Abs(v)) + } + errOf := func(got []float64) float64 { + worst := 0.0 + for k := range tc.n { + g, _ := want[k].Float64() + worst = max(worst, math.Abs(got[k]-g)) + } + return worst / (wantMax) + } + legacy := legacyLine(t, x, tc.sine) + legacyErr := errOf(legacy) + chirp := make([]float64, tc.n) + newLineTransformPlan(tc.n, tc.sine).apply(chirp, x) + chirpErr := errOf(chirp) + t.Logf("%s (max|x| = %.3g): padded route %.3g, chirp route %.3g, relative to max|transform|", + tc.label, scale, legacyErr, chirpErr) + if chirpErr > 1e-13 { + t.Fatalf("%s: chirp route relative error %.3g above the 1e-13 floor", tc.label, chirpErr) + } + if chirpErr > max(4*legacyErr, 4e-15) { + t.Fatalf("%s: chirp route error %.3g against padded route %.3g, beyond the 4x margin", + tc.label, chirpErr, legacyErr) + } + } +} + +// TestPoissonLineTransformAgreementAtSolveSizes measures how far the +// chirp route sits from the padded route at the grid lengths the +// benchmarks and the solvers actually use. Two correct routes to the +// same sums disagree at rounding level; the bound pins that, and the +// log records the movement for the report. +func TestPoissonLineTransformAgreementAtSolveSizes(t *testing.T) { + for _, tc := range []struct { + n int + sine bool + label string + }{ + {255, true, "dst1-255"}, + {256, true, "dst1-256"}, + {510, true, "dst1-510"}, + {511, true, "dst1-511"}, + {256, false, "dct1-256"}, + {510, false, "dct1-510"}, + } { + x := lineFixture(tc.n) + legacy := legacyLine(t, x, tc.sine) + chirp := make([]float64, tc.n) + newLineTransformPlan(tc.n, tc.sine).apply(chirp, x) + legacyMax := 0.0 + for _, v := range legacy { + legacyMax = max(legacyMax, math.Abs(v)) + } + delta := 0.0 + for k := range tc.n { + delta = max(delta, math.Abs(chirp[k]-legacy[k])) + } + rel := delta / max(legacyMax, 1e-300) + t.Logf("%s: max|chirp − padded| = %.3g, relative %.3g", tc.label, delta, rel) + if rel > 1e-11 { + t.Fatalf("%s: routes disagree at relative %.3g, above the 1e-11 rounding bound", tc.label, rel) + } + } +} + +// poissonLegacyTransform runs the pre-chirp block transform through +// the public DST/DCT pair, exactly as the solves' old pipeline did. +func poissonLegacyTransform(t *testing.T, f *core.Array, row0, col0, rows, cols int, sine bool) []float64 { + srcVals := poissonFloats(f) + srcCols := f.Shape()[1] + work := make([]float64, rows*cols) + apply := func(in []float64) []float64 { return legacyLine(t, in, sine) } + row := make([]float64, cols) + for r := range rows { + for c := range cols { + row[c] = srcVals[(r+row0)*srcCols+(c+col0)] + } + if sine && cols == 1 { + work[r*cols] = row[0] + continue + } + copy(work[r*cols:(r+1)*cols], apply(row)) + } + col := make([]float64, rows) + for c := range cols { + for r := range rows { + col[r] = work[r*cols+c] + } + if sine && rows == 1 { + continue + } + out := apply(col) + for r := range rows { + work[r*cols+c] = out[r] + } + } + return work +} + +// TestPoissonSolveRoutesAgree compares the shipped solves against the +// same solves rebuilt on the public DST/DCT pipeline, on deterministic +// grids of two sizes. The routes answer the same mathematics through +// different rounding; the drift must stay at rounding level relative to +// the solution's own scale, and the measured value is logged. +func TestPoissonSolveRoutesAgree(t *testing.T) { + for _, n := range []int{32, 64} { + fVals := make([]float64, n*n) + state := uint64(0x123456789abcdef) ^ uint64(n) + for i := range n * n { + state ^= state << 13 + state ^= state >> 7 + state ^= state << 17 + fVals[i] = float64(int64(state)%2000)/1000 - 1 + } + f, err := core.FromFloats(fVals, n, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := SolvePoissonDirichlet(f, 1, 1) + if err != nil { + t.Fatalf("SolvePoissonDirichlet: %v", err) + } + // The legacy pipeline on the same input: interior DST-I pair + // with the same eigenvalue division the solve runs at lx = ly + // = 1, where hx = hy = 1/(n−1). + spectrum := poissonLegacyTransform(t, f, 1, 1, n-2, n-2, true) + interiorC, interiorR := n-2, n-2 + h := 1 / float64(n-1) + eigX := make([]float64, interiorC+1) + for kx := 1; kx <= interiorC; kx++ { + eigX[kx] = 2 / (h * h) * (1 - math.Cos(math.Pi*float64(kx)/float64(interiorC+1))) + } + eigY := make([]float64, interiorR+1) + for ky := 1; ky <= interiorR; ky++ { + eigY[ky] = 2 / (h * h) * (1 - math.Cos(math.Pi*float64(ky)/float64(interiorR+1))) + } + for ky := 1; ky <= interiorR; ky++ { + row := (ky - 1) * interiorC + for kx := 1; kx <= interiorC; kx++ { + spectrum[row+kx-1] /= eigX[kx] + eigY[ky] + } + } + back := poissonLegacyTransform(t, mustFloats(t, spectrum, interiorR, interiorC), 0, 0, interiorR, interiorC, true) + gotVals := poissonFloats(got) + scale := 0.0 + for _, v := range back { + scale = max(scale, math.Abs(v)) + } + delta := 0.0 + for r := range interiorR { + for c := range interiorC { + delta = max(delta, math.Abs(gotVals[(r+1)*n+(c+1)]-back[r*interiorC+c])) + } + } + rel := delta / max(scale, 1e-300) + t.Logf("Dirichlet %dx%d: max|chirp solve − padded solve| = %.3g, relative %.3g", n, n, delta, rel) + if rel > 1e-11 { + t.Fatalf("Dirichlet %dx%d: solve routes disagree at relative %.3g", n, n, rel) + } + // The solve is deterministic: the same input answers the same + // bits on a second call. + again, err := SolvePoissonDirichlet(f, 1, 1) + if err != nil { + t.Fatalf("second SolvePoissonDirichlet: %v", err) + } + av := poissonFloats(again) + for i := range gotVals { + if av[i] != gotVals[i] { + t.Fatalf("Dirichlet %dx%d: repeated solve moved a bit at %d", n, n, i) + } + } + } +} + +// BenchmarkPoissonLineTransform measures one line transform of the +// lengths the 512×512 Dirichlet and 256×256 Neumann solves carry, +// plan construction included, which is what every line of those grids +// pays. +func BenchmarkPoissonLineTransform(b *testing.B) { + for _, tc := range []struct { + n int + sine bool + name string + }{ + {510, true, "dst1-510"}, + {256, false, "dct1-256"}, + } { + x := make([]float64, tc.n) + for i := range x { + x[i] = math.Sin(float64(i)) + 0.25*math.Cos(3*float64(i)) + } + dst := make([]float64, tc.n) + b.Run(tc.name, func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + newLineTransformPlan(tc.n, tc.sine).apply(dst, x) + } + }) + } +} diff --git a/signal/wavelet.go b/signal/wavelet.go new file mode 100644 index 0000000..c5c5f4d --- /dev/null +++ b/signal/wavelet.go @@ -0,0 +1,310 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Wavelets. The discrete side carries Haar, the one wavelet +// whose filter coefficients are exactly ½ and whose reconstruction is +// exact by construction, at any level the length permits, in the +// packed [A_L, D_L, …, D_1] layout every wavelet text uses. The +// continuous side carries the analytic family (Morlet and +// the Mexican hat) evaluated through the FFT: no filter tables to +// trust, only formulas, which is the library's way. + +// DWT returns the discrete Haar wavelet transform of a rank-1 real +// signal over levels scales (periodic boundary): the coefficients pack +// as [A_levels, D_levels, D_{levels−1}, …, D_1], the approximation +// first and each detail band after it. levels must be at least 1 and +// at most log2(n); energy is conserved (Parseval holds for the +// orthonormal Haar basis). +func DWT(x *core.Array, levels int) (*core.Array, error) { + const name = "DWT" + n, err := waveletValidate(name, x, levels) + if err != nil { + return nil, err + } + coef := make([]float64, n) + for i := range n { + coef[i] = x.FloatAt(i) + } + // In each pass the live block at [0, m) halves: averages land in + // [0, m/2), details (scaled by 1/sqrt2 for orthonormality) in + // [m/2, m) of the LIVE block, then shuffle the details to their + // final resting slot at the end. + m := n + for range levels { + half := m / 2 + avg := make([]float64, half) + det := make([]float64, half) + s := math.Sqrt2 + for i := range half { + avg[i] = (coef[2*i] + coef[2*i+1]) / s + det[i] = (coef[2*i] - coef[2*i+1]) / s + } + copy(coef[:half], avg) + // Details of this level go to the tail slot reserved for D_{l+1}. + copy(coef[n-half-(n-m):n-(n-m)], det) + m = half + } + return core.FromFloats(coef, n) +} + +// IDWT inverts DWT over the same level count and layout. The +// coefficients are widened through FloatAt, the accessor DWT reads its +// input with, so every real dtype DWT accepts inverts here as well. +func IDWT(coef *core.Array, levels int) (*core.Array, error) { + const name = "IDWT" + n, err := waveletValidate(name, coef, levels) + if err != nil { + return nil, err + } + out := make([]float64, n) + for i := range n { + out[i] = coef.FloatAt(i) + } + // Undo level by level, innermost (shortest) first. + m := n >> levels + for l := levels; l >= 1; l-- { + half := m + m = 2 * m + avg := append([]float64(nil), out[:half]...) + det := append([]float64(nil), out[n-half-(n-m):n-(n-m)]...) + s := math.Sqrt2 + for i := range half { + out[2*i] = (avg[i] + det[i]) / s + out[2*i+1] = (avg[i] - det[i]) / s + } + } + return core.FromFloats(out, n) +} + +// waveletValidate checks the shared contract and returns n. +func waveletValidate(name string, x *core.Array, levels int) (int, error) { + if x.NDim() != 1 { + return 0, base.Errf("%s: needs a rank-1 signal, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return 0, base.Errf("%s: complex signals are not supported", name) + } + n := x.Len() + if n == 0 { + return 0, base.Errf("%s: an empty signal has no transform", name) + } + if levels < 1 || (n>>levels)< 0 and turns the wavelet's time axis into + // NaN with no error, so finiteness is part of the gate. + if !(dt > 0) || math.IsInf(dt, 0) { + return nil, base.Errf("%s: the sample spacing must be positive and finite, got %g", name, dt) + } + n := x.Len() + if n == 0 { + return nil, base.Errf("%s: an empty signal has no transform", name) + } + if len(scales) == 0 { + return nil, base.Errf("%s: at least one scale is required", name) + } + for _, s := range scales { + if !(s > 0) || math.IsInf(s, 0) { + return nil, base.Errf("%s: every scale must be positive and finite, got %g", name, s) + } + } + switch wavelet { + case Morlet, MexicanHat: + default: + return nil, base.Errf("%s: unknown wavelet %q (want morlet or mexicanhat)", name, wavelet) + } + + // The signal padded to 2n so the linear convolution has room. The + // transform runs in place: FFT and IFFT are pure functions of the + // payload slice, so driving transform directly on scratch the + // function owns produces exactly the bits the wrappers produced, + // without their per-call copies. + sig := make([]complex128, 2*n) + for i := range n { + sig[i] = complex(x.FloatAt(i), 0) + } + transform(sig, -1) + + out := make([]complex128, len(scales)*n) + // transformScale runs the whole per-scale pipeline: the wavelet's + // cached spectrum, the conjugated product with the signal's + // spectrum and the return to time. The wavelet of a scale is a pure + // function of the transform geometry, so its spectrum is served + // from the same kind of cache the transform's twiddle tables keep; + // every other scratch is the caller's w, fully rewritten per scale. + // The per-row arithmetic sequence is the serial one unchanged, so + // the split cannot move a bit. + transformScale := func(k int, a float64, w []complex128) { + spec := cwtWaveletSpectrum(wavelet, n, dt, a) + for i := range w { + w[i] = sig[i] * conj128(spec[i]) + } + transform(w, +1) + // The inverse scale, applied entry by entry exactly as IFFT + // applies it, lands straight in this scale's output row. + row := out[k*n : (k+1)*n] + scale := complex(float64(2*n), 0) + for i := range row { + row[i] = w[i] / scale + } + } + // The scales split across workers once a scale's wavelet build and + // FFT pair outgrow the worker spawn cost; below that floor the + // calling goroutine walks every scale itself. Each worker owns one + // scratch line and rewrites it whole for every scale it serves. + wlen := 2 * n + if n >= cwtParallelMinN { + engine.Parallel(len(scales), func(ks, ke int) { + w := make([]complex128, wlen) + for k := ks; k < ke; k++ { + transformScale(k, scales[k], w) + } + }) + } else { + w := make([]complex128, wlen) + for k, a := range scales { + transformScale(k, a, w) + } + } + arr, err := core.ComplexFromArray(out, len(scales), n) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + return arr, nil +} + +// cwtSpectrumKey names one cached wavelet spectrum: the analysing +// wavelet and the exact transform geometry it was built for. +type cwtSpectrumKey struct { + wavelet CWTWavelet + n int + dt float64 + scale float64 +} + +// cwtSpectra caches per-scale wavelet spectra after their forward +// transform, keyed by the geometry that determines them. The spectra +// are read-only once published. Two caps bound the retention: the +// twiddleCacheMax policy applies per entry (a transform whose padded +// length exceeds it still builds its spectrum, it just does not keep +// it), and the map holds at most cwtSpectrumCacheMax entries, so a +// caller sweeping unboundedly many distinct scales cannot pin +// unbounded memory. +var ( + cwtSpectrumMu sync.RWMutex + cwtSpectrumKeys = map[cwtSpectrumKey][]complex128{} +) + +// cwtSpectrumCacheMax is the entry cap of the wavelet-spectrum cache. +const cwtSpectrumCacheMax = 128 + +// cwtWaveletSpectrum returns the forward transform of the zero-padded, +// L1-normalised wavelet of the given scale, its DC bin zeroed for +// admissibility. Building it repeats the evaluation the per-scale path +// always ran, sample for sample. +func cwtWaveletSpectrum(wavelet CWTWavelet, n int, dt, a float64) []complex128 { + key := cwtSpectrumKey{wavelet: wavelet, n: n, dt: dt, scale: a} + cwtSpectrumMu.RLock() + spec, ok := cwtSpectrumKeys[key] + cwtSpectrumMu.RUnlock() + if ok { + return spec + } + wlen := 2 * n + wrap := wlen/2 + 1 + w := make([]complex128, wlen) + switch wavelet { + case Morlet: + // π^{-1/4}·e^{iω₀t/a}·e^{−t²/(2a²)}, L1-normalised by 1/a. + // Sincos answers the same bits the separate Sin and Cos + // produce (verified bit-for-bit), so the wavelet is + // unchanged while the trig work halves. + const omega0 = 5.0 + amp := math.Pow(math.Pi, -0.25) / a + for j := range wrap { + arg := dt * float64(j) / a + s, c := math.Sincos(omega0 * arg) + w[j] = complex(amp*c, amp*s) * + complex(math.Exp(-arg*arg/2), 0) + } + period := dt * float64(wlen) + for j := wrap; j < wlen; j++ { + arg := (dt*float64(j) - period) / a // wrap to the negative half + s, c := math.Sincos(omega0 * arg) + w[j] = complex(amp*c, amp*s) * + complex(math.Exp(-arg*arg/2), 0) + } + case MexicanHat: + // (2/√3)π^{-1/4}(1−t²/a²)e^{−t²/2a²}, scaled by 1/a. + c := 2 / math.Sqrt(3) * math.Pow(math.Pi, -0.25) + for j := range wrap { + arg := dt * float64(j) / a + w[j] = complex(c*(1-arg*arg)*math.Exp(-arg*arg/2)/a, 0) + } + period := dt * float64(wlen) + for j := wrap; j < wlen; j++ { + arg := (dt*float64(j) - period) / a + w[j] = complex(c*(1-arg*arg)*math.Exp(-arg*arg/2)/a, 0) + } + } + transform(w, -1) + // Zero the wavelet's DC bin: the truncated tails leave a + // rounding-level mean, and admissibility demands exactly zero. + w[0] = 0 + if wlen <= twiddleCacheMax { + cwtSpectrumMu.Lock() + if len(cwtSpectrumKeys) < cwtSpectrumCacheMax { + cwtSpectrumKeys[key] = w + } + cwtSpectrumMu.Unlock() + } + return w +} diff --git a/signal/wavelet_test.go b/signal/wavelet_test.go new file mode 100644 index 0000000..9d9091b --- /dev/null +++ b/signal/wavelet_test.go @@ -0,0 +1,267 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestHaarReconstruction pins exactness: IDWT(DWT(x)) = x, energy +// conserved, across level counts. +func TestHaarReconstruction(t *testing.T) { + g := core.NewGenerator(3) + vals := make([]float64, 64) + for i := range vals { + vals[i] = g.NormalUnit() + } + x, _ := core.FromFloats(vals, 64) + for _, levels := range []int{1, 3, 6} { + c, err := DWT(x, levels) + if err != nil { + t.Fatalf("DWT(%d): %v", levels, err) + } + eIn, eC := 0.0, 0.0 + for i := range 64 { + eIn += x.FloatAt(i) * x.FloatAt(i) + eC += c.FloatAt(i) * c.FloatAt(i) + } + if math.Abs(eIn-eC) > 1e-10*eIn { + t.Fatalf("levels %d: energy %g vs %g", levels, eIn, eC) + } + back, err := IDWT(c, levels) + if err != nil { + t.Fatalf("IDWT(%d): %v", levels, err) + } + for i := range 64 { + if math.Abs(back.FloatAt(i)-x.FloatAt(i)) > 1e-12 { + t.Fatalf("levels %d: reconstruction off at %d", levels, i) + } + } + } +} + +// TestHaarStepSignal pins the sparsity promise: a piecewise-constant +// signal has wavelet coefficients concentrated at the jumps. +func TestHaarStepSignal(t *testing.T) { + // The jump sits at index 5: off every dyadic boundary, because a + // step landing exactly on one is invisible to Haar at every level. + vals := make([]float64, 16) + for i := 5; i < 16; i++ { + vals[i] = 1 + } + x, _ := core.FromFloats(vals, 16) + c, err := DWT(x, 1) + if err != nil { + t.Fatalf("DWT: %v", err) + } + // Only detail 2 (the pair 4/5) carries the jump. + for i := range 8 { + if i != 2 && math.Abs(c.FloatAt(8+i)) > 1e-12 { + t.Fatalf("detail %d = %g, only the jump pair should fire", i, c.FloatAt(8+i)) + } + } + if math.Abs(math.Abs(c.FloatAt(10))-1/math.Sqrt2) > 1e-12 { + t.Fatalf("jump detail = %g, want ±1/√2", c.FloatAt(10)) + } +} + +// TestCWTMorletRidge pins the time-frequency map: a pure sinusoid's +// Morlet transform peaks at the scale carrying its frequency. +func TestCWTMorletRidge(t *testing.T) { + const ( + n = 1024 + dt = 0.01 + freq = 5.0 + omega = 5.0 // Morlet ω₀ + ) + vals := make([]float64, n) + for i := range n { + vals[i] = math.Sin(2 * math.Pi * freq * dt * float64(i)) + } + x, _ := core.FromFloats(vals, n) + scales := make([]float64, 30) + for i := range scales { + scales[i] = 0.01 + 0.01*float64(i) + } + w, err := CWT(x, Morlet, scales, dt) + if err != nil { + t.Fatalf("CWT: %v", err) + } + // Ridge: the scale maximising the mean magnitude. The plain peak + // estimate of the Morlet centre frequency, f ≈ ω₀/(2πa), is within + // the tolerance the test uses. + best := 0.0 + bestMag := -1.0 + for k := range scales { + m := 0.0 + for i := range n { + z := w.ComplexAt(k*n + i) + m += math.Hypot(real(z), imag(z)) + } + m /= float64(n) + if m > bestMag { + bestMag, best = m, scales[k] + } + } + want := omega / (2 * math.Pi * freq) + if math.Abs(best-want) > 0.5*want { + t.Fatalf("ridge scale %.4f, expected around %.4f", best, want) + } +} + +// TestCWTMexicanHatBump pins the zero-mean bump detector: the CWT of a +// Gaussian bump peaks at the scale matching its width, and a constant +// signal transforms to zero (the wavelet has no DC). +func TestCWTMexicanHatBump(t *testing.T) { + const n = 512 + vals := make([]float64, n) + for i := range n { + t0 := (float64(i) - n/2) * 0.05 + vals[i] = math.Exp(-t0 * t0 / 2) + } + x, _ := core.FromFloats(vals, n) + scales := []float64{0.5, 1.0, 2.0, 4.0, 8.0} + w, err := CWT(x, MexicanHat, scales, 0.05) + if err != nil { + t.Fatalf("CWT: %v", err) + } + mags := make([]float64, len(scales)) + for k := range scales { + m := 0.0 + for i := range n { + z := w.ComplexAt(k*n + i) + m = max(m, math.Hypot(real(z), imag(z))) + } + mags[k] = m + } + // The bump has unit width in time units: dt·(width samples); the + // peak response lands at the scale of order one, not the extremes. + if mags[1] < mags[0] || mags[1] < mags[len(scales)-1] { + t.Fatalf("bump response %v peaks at an extreme scale", mags) + } + // DC insensitivity: a constant gives a numerically zero transform + // in the interior, away from both wrap-around edges by more than + // the wavelet's ~4-scale reach (4a/dt = 40 samples here). + consts := make([]float64, 256) + for i := range consts { + consts[i] = 3.5 + } + cx, _ := core.FromFloats(consts, 256) + cw, err := CWT(cx, MexicanHat, []float64{1.0}, 0.1) + if err != nil { + t.Fatalf("CWT constant: %v", err) + } + for i := 64; i < 192; i++ { + z := cw.ComplexAt(i) + if math.Hypot(real(z), imag(z)) > 1e-5 { + t.Fatalf("constant leaked into the mexican-hat transform at %d: %g", i, math.Hypot(real(z), imag(z))) + } + } +} + +// TestWaveletErrors pins the input gates. +func TestWaveletErrors(t *testing.T) { + x, _ := core.FromFloats([]float64{1, 2, 3, 4}, 4) + if _, err := DWT(x, 0); err == nil { + t.Error("zero levels accepted") + } + if _, err := DWT(x, 3); err == nil { + t.Error("levels beyond log2(n) accepted") + } + odd, _ := core.FromFloats([]float64{1, 2, 3}, 3) + if _, err := DWT(odd, 1); err == nil { + t.Error("odd length accepted") + } + if _, err := CWT(x, "db4", []float64{1}, 0.1); err == nil { + t.Error("unknown wavelet accepted") + } + if _, err := CWT(x, Morlet, []float64{-1}, 0.1); err == nil { + t.Error("negative scale accepted") + } +} + +// cwtReference computes one scale of CWT directly in the time domain: +// the output at sample i is the cyclic correlation of the signal with +// the wavelet, sum over j of x[j]·conj(ψ[(j−i) mod 2n]). The sum is +// what the transform's product of spectra evaluates, stated without an +// FFT, and the wavelet is built from the documented rule: its sample j +// sits at t = dt·j and the samples past the midpoint wrap to +// t = dt·(j − 2n), so the wavelet is zero-padded and centred. The +// transform removes the wavelet's DC bin, which the subtraction of the +// mean does here. +func cwtReference(vals []float64, dt, a float64) []complex128 { + n := len(vals) + wlen := 2 * n + const omega0 = 5.0 + amp := math.Pow(math.Pi, -0.25) / a + w := make([]complex128, wlen) + mean := complex(0, 0) + for j := range wlen { + tt := dt * float64(j) + if j > wlen/2 { + tt -= dt * float64(wlen) // wrap to the negative half + } + arg := tt / a + w[j] = complex(amp*math.Cos(omega0*arg), amp*math.Sin(omega0*arg)) * + complex(math.Exp(-arg*arg/2), 0) + mean += w[j] + } + mean /= complex(float64(wlen), 0) + row := make([]complex128, n) + for i := range n { + s := complex(0, 0) + for j := range n { + s += complex(vals[j], 0) * conj128(w[((j-i)%wlen+wlen)%wlen]-mean) + } + row[i] = s + } + return row +} + +// TestCWTWrapBoundary pins the wavelet's wrap point against the direct +// correlation above. The transform's route differs from the reference +// only in the wavelet sample at j = wlen/2: it is the last sample of +// the non-negative half, at t = +dt·n, and a pass that wraps it too +// puts it at −dt·n instead. That single sample is invisible at small +// scales, where the Gaussian has long since decayed, and decides the +// answer at scales near dt·n, because zeroing the wavelet's DC bin +// leaves every sample's deviation from the mean in every output. The +// scales below span both regimes. +func TestCWTWrapBoundary(t *testing.T) { + const ( + n = 64 + dt = 1.0 + ) + vals := make([]float64, n) + for i := range n { + vals[i] = math.Sin(0.3*float64(i)) + 0.25*math.Cos(0.11*float64(i)) + } + x := mustFromFloats(t, vals, n) + scales := []float64{float64(n) * dt, float64(n) * dt / 2, 4, 0.5} + out, err := CWT(x, Morlet, scales, dt) + if err != nil { + t.Fatalf("CWT: %v", err) + } + for k, a := range scales { + want := cwtReference(vals, dt, a) + worst, mag := 0.0, 0.0 + for i := range n { + got := out.ComplexAt(k*n + i) + d := got - want[i] + if e := math.Hypot(real(d), imag(d)); e > worst { + worst = e + } + if m := math.Hypot(real(want[i]), imag(want[i])); m > mag { + mag = m + } + } + if worst > 1e-11*mag { + t.Fatalf("scale %g: worst deviation from the direct correlation %.6g (scale %.6g), want the wrapped wavelet sample to sit at +dt·n", + a, worst, mag) + } + } +} diff --git a/signal/waveletdwt.go b/signal/waveletdwt.go new file mode 100644 index 0000000..aba22c6 --- /dev/null +++ b/signal/waveletdwt.go @@ -0,0 +1,323 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Daubechies discrete wavelets: db2 through db8 beside the +// Haar transform the package already carries in DWT and IDWT. Each +// family is an orthonormal two-channel filter bank driven over levels +// of halving blocks, packed in the same [A_L, D_L, …, D_1] layout the +// Haar transform uses; each dbN pair has N vanishing moments and 2N +// taps, so higher families resolve smoother signals into sparser +// details but tie the boundary condition to longer blocks. +// +// The tables below are not trusted on authority: they were produced +// by the spectral factorisation of Daubechies' polynomial +// P(y) = Σ_k C(N−1+k, k)·y^k, and the tests certify them against the +// defining conditions, the unit norm, the shift-2 orthogonality, the +// N vanishing moments and the spectral identity +// |H(ω)|² = 2·cos^{2N}(ω/2)·P(sin²(ω/2)), so a mistyped digit cannot +// survive. + +// Daubechies names a member of the Daubechies wavelet family: dbN +// carries N vanishing moments and a 2N-tap filter pair. Haar remains +// the db1 case and keeps its own exact transform in DWT and IDWT. +type Daubechies string + +const ( + // DB2 is the 4-tap Daubechies wavelet with 2 vanishing moments. + DB2 Daubechies = "db2" + // DB3 is the 6-tap Daubechies wavelet with 3 vanishing moments. + DB3 Daubechies = "db3" + // DB4 is the 8-tap Daubechies wavelet with 4 vanishing moments. + DB4 Daubechies = "db4" + // DB5 is the 10-tap Daubechies wavelet with 5 vanishing moments. + DB5 Daubechies = "db5" + // DB6 is the 12-tap Daubechies wavelet with 6 vanishing moments. + DB6 Daubechies = "db6" + // DB7 is the 14-tap Daubechies wavelet with 7 vanishing moments. + DB7 Daubechies = "db7" + // DB8 is the 16-tap Daubechies wavelet with 8 vanishing moments. + DB8 Daubechies = "db8" +) + +// The scaling filters h, the low-pass analysis side, normalised to +// Σ h = √2. Treat them as read-only. +var ( + db2Coeffs = []float64{ + 0.48296291314453421, 0.83651630373780794, 0.22414386804201339, -0.1294095225512604, + } + db3Coeffs = []float64{ + 0.33267055295008258, 0.80689150931109266, 0.45987750211849149, -0.13501102001025453, + -0.085441273882026644, 0.035226291885709533, + } + db4Coeffs = []float64{ + 0.23037781330889651, 0.71484657055291556, 0.63088076792985903, -0.027983769416859948, + -0.18703481171909306, 0.030841381835560764, 0.032883011666885203, -0.010597401785069032, + } + db5Coeffs = []float64{ + 0.1601023979741929, 0.60382926979718965, 0.72430852843777283, 0.13842814590132091, + -0.24229488706638203, -0.032244869584638361, 0.077571493840045691, -0.0062414902127982726, + -0.012580751999081997, 0.0033357252854737704, + } + db6Coeffs = []float64{ + 0.11154074335010949, 0.49462389039845328, 0.75113390802109559, 0.31525035170919741, + -0.22626469396543983, -0.12976686756726177, 0.097501605587322904, 0.027522865530305803, + -0.031582039317486044, 0.00055384220116149461, 0.0047772575109455116, -0.0010773010853084798, + } + db7Coeffs = []float64{ + 0.07785205408500917, 0.39653931948191723, 0.72913209084623509, 0.46978228740519357, + -0.14390600392856487, -0.22403618499387515, 0.071309219266830329, 0.080612609151083051, + -0.0380299369350144, -0.016574541630666881, 0.012550998556099839, 0.00042957797292136738, + -0.0018016407040474906, 0.00035371379997452024, + } + db8Coeffs = []float64{ + 0.054415842243104001, 0.3128715909143, 0.67563073629728954, 0.58535468365420718, + -0.015829105256349327, -0.28401554296154741, 0.00047248457391386436, 0.12874742662047808, + -0.017369301001807447, -0.044088253930794685, 0.013981027917398216, 0.0087460940474057974, + -0.0048703529934515776, -0.00039174037337694672, 0.00067544940645056933, -0.00011747678412476953, + } +) + +// daubechiesTaps resolves a family name to its scaling filter. +func daubechiesTaps(name string, family Daubechies) ([]float64, error) { + switch family { + case DB2: + return db2Coeffs, nil + case DB3: + return db3Coeffs, nil + case DB4: + return db4Coeffs, nil + case DB5: + return db5Coeffs, nil + case DB6: + return db6Coeffs, nil + case DB7: + return db7Coeffs, nil + case DB8: + return db8Coeffs, nil + default: + return nil, base.Errf("%s: unknown Daubechies family %q (want db2 through db8)", name, string(family)) + } +} + +// DWTMode picks the boundary treatment of the Daubechies transform. +type DWTMode int + +const ( + // DWTPeriodic treats the signal as one period of a periodic + // sequence: the default. Every block keeps its exact energy and + // the length must offer every level a multiple of the filter + // length to work on. + DWTPeriodic DWTMode = iota + // DWTZeroPad extends the signal with zeros at the tail to the + // next length the level tree needs, so any length transforms; + // the coefficients past the signal's own span carry the + // response of the step to zero at the seam. + DWTZeroPad +) + +// DaubechiesDWT returns the discrete Daubechies wavelet transform of +// the rank-1 real signal x over levels scales, packed as +// [A_levels, D_levels, …, D_1]: the deepest approximation first, each +// detail band after it, the finest last. mode picks the boundary +// treatment, DWTPeriodic by default: it refuses a length that does +// not leave every live block a multiple of the 2N-tap filter (the +// deepest block must hold at least one filter length), while +// DWTZeroPad zero-pads the tail to the next such length first. The +// periodic transform conserves energy exactly; a zero-padded one +// conserves the energy of the padded signal. levels must be at +// least 1. +func DaubechiesDWT(x *core.Array, family Daubechies, levels int, mode DWTMode) (*core.Array, error) { + const name = "DaubechiesDWT" + src, h, levels, mode, err := dwtPrepare(name, x, family, levels, mode) + if err != nil { + return nil, err + } + L := len(h) + n := x.Len() + // The periodic contract gates the length; the zero-padded one + // extends the tail to the smallest length that offers every + // level a multiple of the filter length instead. + length := n + if mode == DWTPeriodic { + if err := dwtLengthGate(name, n, L, levels); err != nil { + return nil, err + } + } else { + // The padded block must stay inside the allocator's reach: the + // shift wraps past the addressable range first, and a span that + // survives the wrap but tops the makeslice ceiling of 2^45 + // float64 elements panics in the make below instead of refusing + // here. + span := L << (levels - 1) + if span <= 0 || span >= 1<<45 { + return nil, base.Errf("%s: %d levels need a padded length no machine could hold", name, levels) + } + length = ((n-1)/span + 1) * span + } + work := make([]float64, length) + copy(work, src) + g := daubechiesHighPass(h) + // Each pass halves the live block: the approximation and the + // detail land in scratch, then the block rearranges into + // [A, D] exactly as the packed layout wants. + scratch := make([]float64, length) + m := length + for range levels { + half := m / 2 + a, d := scratch[:half], scratch[half:m] + dwtForwardLevel(work[:m], a, d, h, g) + copy(work[:half], a) + copy(work[half:m], d) + m = half + } + return core.FromFloats(work, length) +} + +// DaubechiesIDWT inverts DaubechiesDWT over the same family, level +// count and mode, undoing the packed layout deepest band first. The +// coefficient vector must satisfy the length contract every forward +// transform produced: a multiple of the filter length at every level +// of the tree, whichever mode the forward ran. A zero-padded forward +// transform therefore inverts to the padded length, whose head is +// the original signal; the mode argument gates exactly this +// contract, the inverse bank itself is the periodic one. +func DaubechiesIDWT(coef *core.Array, family Daubechies, levels int, mode DWTMode) (*core.Array, error) { + const name = "DaubechiesIDWT" + src, h, levels, _, err := dwtPrepare(name, coef, family, levels, mode) + if err != nil { + return nil, err + } + if err := dwtLengthGate(name, coef.Len(), len(h), levels); err != nil { + return nil, err + } + g := daubechiesHighPass(h) + work := make([]float64, len(src)) + copy(work, src) + m := len(work) >> levels + for range levels { + half := m + m = 2 * m + a := slices.Clone(work[:half]) + d := slices.Clone(work[half:m]) + for j := range m { + var av, dv float64 + for k := range h { + // The synthesis taps land at even steps only: + // e/2 is the band index the tap reads. + e := j - k + if e < 0 { + e += m + } + if e&1 == 1 { + continue + } + av += h[k] * a[e>>1] + dv += g[k] * d[e>>1] + } + work[j] = av + dv + } + } + return core.FromFloats(work, len(work)) +} + +// dwtPrepare runs the shared contract of the Daubechies transforms +// ahead of the length question: rank-1 real non-empty input, a known +// family, at least one level, a known mode. It returns the widened +// input, the scaling filter and the validated level count and mode. +func dwtPrepare(name string, x *core.Array, family Daubechies, levels int, mode DWTMode) (src []float64, h []float64, levelsOut int, modeOut DWTMode, err error) { + if x.NDim() != 1 { + return nil, nil, 0, 0, base.Errf("%s: needs a rank-1 signal, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, nil, 0, 0, base.Errf("%s: complex signals are not supported", name) + } + if x.Len() == 0 { + return nil, nil, 0, 0, base.Errf("%s: an empty signal has no transform", name) + } + h, err = daubechiesTaps(name, family) + if err != nil { + return nil, nil, 0, 0, err + } + if levels < 1 { + return nil, nil, 0, 0, base.Errf("%s: levels must be at least 1, got %d", name, levels) + } + switch mode { + case DWTPeriodic, DWTZeroPad: + default: + return nil, nil, 0, 0, base.Errf("%s: unknown boundary mode %d", name, int(mode)) + } + return widenFloats(x), h, levels, mode, nil +} + +// dwtLengthGate enforces the tree contract on a length: the deepest +// block, n>>levels−1 samples, must be a whole multiple of the filter +// length and at least one filter long, which the halvings before it +// then inherit. Shifts past the word width answer 0 and fail the +// gate, which is exactly the refusal a too deep tree wants. +func dwtLengthGate(name string, n, L, levels int) error { + blk := n >> (levels - 1) + if levels-1 >= 64 || (blk<<(levels-1)) != n || blk < L || blk%L != 0 { + return base.Errf("%s: %d levels need every live block to hold a multiple of the %d-tap filter, which %d samples do not offer", + name, levels, L, n) + } + return nil +} + +// daubechiesHighPass derives the wavelet (high-pass) half of the bank +// from the scaling filter: g[k] = (−1)^k·h[L−1−k]. +func daubechiesHighPass(h []float64) []float64 { + L := len(h) + g := make([]float64, L) + for k := range L { + sign := 1.0 + if k%2 == 1 { + sign = -1 + } + g[k] = sign * h[L-1-k] + } + return g +} + +// dwtForwardLevel runs one periodic analysis level over the block +// x of even length m: a[i] = Σ h[k]·x[(2i+k) mod m] and the same with +// g for d. Outputs whose taps never wrap walk the block directly; +// the ones near the block's end correct the few negative indices. +func dwtForwardLevel(x, a, d, h, g []float64) { + m := len(x) + L := len(h) + half := m / 2 + limit := (m - L) / 2 + for i := 0; i <= limit; i++ { + base := 2 * i + var av, dv float64 + for k := range L { + xv := x[base+k] + av += h[k] * xv + dv += g[k] * xv + } + a[i], d[i] = av, dv + } + for i := limit + 1; i < half; i++ { + base := 2*i - m + var av, dv float64 + for k := range L { + idx := base + k + if idx < 0 { + idx += m + } + xv := x[idx] + av += h[k] * xv + dv += g[k] * xv + } + a[i], d[i] = av, dv + } +} diff --git a/signal/waveletdwt_test.go b/signal/waveletdwt_test.go new file mode 100644 index 0000000..16f00aa --- /dev/null +++ b/signal/waveletdwt_test.go @@ -0,0 +1,370 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// daubechiesFamilies lists every family the machinery carries, with +// its vanishing-moment count. +var daubechiesFamilies = []struct { + family Daubechies + moments int + coeffs []float64 +}{ + {DB2, 2, db2Coeffs}, + {DB3, 3, db3Coeffs}, + {DB4, 4, db4Coeffs}, + {DB5, 5, db5Coeffs}, + {DB6, 6, db6Coeffs}, + {DB7, 7, db7Coeffs}, + {DB8, 8, db8Coeffs}, +} + +// TestDaubechiesPerfectReconstruction pins the defining property of +// the orthonormal bank: inverse(forward(x)) == x to machine +// precision, for every family, at every level count the length +// allows. +func TestDaubechiesPerfectReconstruction(t *testing.T) { + for _, f := range daubechiesFamilies { + L := 2 * f.moments + for _, mult := range []int{4, 8} { + n := L * mult + g := core.NewGenerator(int64(f.moments)*100 + int64(mult)) + vals := make([]float64, n) + for i := range vals { + vals[i] = g.NormalUnit() + } + x := mustFloats(t, vals) + for _, levels := range []int{1, 2, 3} { + c, err := DaubechiesDWT(x, f.family, levels, DWTPeriodic) + if err != nil { + t.Fatalf("%s n=%d levels=%d: %v", f.family, n, levels, err) + } + back, err := DaubechiesIDWT(c, f.family, levels, DWTPeriodic) + if err != nil { + t.Fatalf("%s n=%d levels=%d inverse: %v", f.family, n, levels, err) + } + for i := range vals { + if math.Abs(back.FloatAt(i)-vals[i]) > 1e-12 { + t.Fatalf("%s n=%d levels=%d: reconstruction off at %d: %v vs %v", + f.family, n, levels, i, back.FloatAt(i), vals[i]) + } + } + } + } + } +} + +// TestDaubechiesEnergyPreservation pins Parseval: the periodic +// transform moves no energy between the signal and its coefficients. +func TestDaubechiesEnergyPreservation(t *testing.T) { + for _, f := range []struct { + family Daubechies + moments int + }{ + {DB4, 4}, + {DB7, 7}, + } { + L := 2 * f.moments + n := 8 * L + g := core.NewGenerator(int64(f.moments) + 50) + vals := make([]float64, n) + for i := range vals { + vals[i] = g.NormalUnit() + } + x := mustFloats(t, vals) + for _, levels := range []int{1, 2, 4} { + c, err := DaubechiesDWT(x, f.family, levels, DWTPeriodic) + if err != nil { + t.Fatalf("%s levels=%d: %v", f.family, levels, err) + } + eIn, eC := 0.0, 0.0 + for i := range n { + eIn += vals[i] * vals[i] + eC += c.FloatAt(i) * c.FloatAt(i) + } + if math.Abs(eIn-eC) > 1e-11*eIn { + t.Fatalf("%s levels=%d: energy %g vs %g", f.family, levels, eC, eIn) + } + } + } +} + +// TestDaubechiesFilterTables certifies every coefficient table +// against the defining conditions of the family: the √2 sum, the +// unit norm, the shift-2 orthogonality, the vanishing moments, and +// the spectral identity against Daubechies' polynomial P, so a +// mistyped digit cannot survive. +func TestDaubechiesFilterTables(t *testing.T) { + for _, f := range daubechiesFamilies { + h := f.coeffs + L := len(h) + sum, norm := 0.0, 0.0 + for _, v := range h { + sum += v + norm += v * v + } + if math.Abs(sum-math.Sqrt2) > 1e-14 { + t.Fatalf("%s: sum %v, want √2", f.family, sum) + } + if math.Abs(norm-1) > 1e-14 { + t.Fatalf("%s: norm %v, want 1", f.family, norm) + } + for l := 1; 2*l < L; l++ { + s := 0.0 + for k := 0; k+2*l < L; k++ { + s += h[k] * h[k+2*l] + } + if math.Abs(s) > 1e-13 { + t.Fatalf("%s: shift-2 orthogonality at l=%d: %v", f.family, l, s) + } + } + // The vanishing moments, judged relatively: the raw sums of + // k^j-weighted terms reach magnitudes where the arithmetic + // noise floor alone is around 1e-9. + for j := 0; j < f.moments; j++ { + var s, scale float64 + for k, v := range h { + s += math.Pow(-1, float64(k)) * math.Pow(float64(k), float64(j)) * v + scale += math.Pow(float64(k), float64(j)) * math.Abs(v) + } + if math.Abs(s) > 1e-12*scale { + t.Fatalf("%s: vanishing moment %d: %v (scale %v)", f.family, j, s, scale) + } + } + // The spectral identity |H(ω)|² = 2·cos^{2N}(ω/2)·P(sin²(ω/2)) + // with P(y) = Σ C(N−1+k, k)·y^k. + P := make([]float64, f.moments) + for k := range P { + P[k] = binomialCoeffs(f.moments - 1 + k)[k] + } + for i := 0; i <= 512; i++ { + w := math.Pi * float64(i) / 512 + var hr, hi float64 + for k, v := range h { + hr += v * math.Cos(float64(k)*w) + hi -= v * math.Sin(float64(k)*w) + } + s := math.Sin(w / 2) + y := s * s + pv := 0.0 + pow := 1.0 + for k := range P { + pv += P[k] * pow + pow *= y + } + c := math.Cos(w / 2) + want := 2 * math.Pow(c*c, float64(f.moments)) * pv + if got := hr*hr + hi*hi; math.Abs(got-want) > 1e-12 { + t.Fatalf("%s: spectral identity at ω %d/512: %v vs %v", f.family, i, got, want) + } + } + } +} + +// TestDaubechiesKnownTwoLevel pins a whole 2-level decomposition, +// db2 over [1..8], against independently computed reference values +// (an explicit analysis matrix built row by row), layout included. +func TestDaubechiesKnownTwoLevel(t *testing.T) { + x := mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}) + c, err := DaubechiesDWT(x, DB2, 2, DWTPeriodic) + if err != nil { + t.Fatalf("DaubechiesDWT: %v", err) + } + want := []float64{ + 5.9019237886466849, 12.098076211353316, + 0.36602540378444015, -3.8301270189221936, + 0, 0, 0, -2.8284271247461907, + } + for i, w := range want { + got := c.FloatAt(i) + if math.Abs(got-w) > 1e-12 { + t.Fatalf("coefficient %d: %.15g, want %.15g", i, got, w) + } + } +} + +// TestDaubechiesConstantDetails pins the vanishing moments end to +// end: the detail bands of a constant signal are zero at every level +// and the approximation scales by 2^(levels/2), the Σh = √2 gain. +func TestDaubechiesConstantDetails(t *testing.T) { + for _, f := range []struct { + family Daubechies + moments int + }{ + {DB2, 2}, + {DB5, 5}, + {DB8, 8}, + } { + L := 2 * f.moments + n := 4 * L + for _, levels := range []int{1, 2} { + vals := make([]float64, n) + for i := range vals { + vals[i] = 3.5 + } + c, err := DaubechiesDWT(mustFloats(t, vals), f.family, levels, DWTPeriodic) + if err != nil { + t.Fatalf("%s levels=%d: %v", f.family, levels, err) + } + block := n >> levels + wantA := 3.5 * math.Pow(math.Sqrt2, float64(levels)) + for i := range block { + if got := c.FloatAt(i); math.Abs(got-wantA) > 1e-12 { + t.Fatalf("%s levels=%d: approximation %d: %v, want %v", + f.family, levels, i, got, wantA) + } + } + for i := block; i < n; i++ { + if got := c.FloatAt(i); math.Abs(got) > 1e-12 { + t.Fatalf("%s levels=%d: detail coefficient %d = %v, want 0", + f.family, levels, i, got) + } + } + } + } +} + +// TestDaubechiesZeroPad pins the padding mode: a length no level +// tree would accept transforms on the next valid padded length, the +// padded signal's energy is conserved, and the inverse returns the +// padded length whose head is the original signal. +func TestDaubechiesZeroPad(t *testing.T) { + g := core.NewGenerator(77) + const n = 13 + vals := make([]float64, n) + var eIn float64 + for i := range vals { + vals[i] = g.NormalUnit() + eIn += vals[i] * vals[i] + } + x := mustFloats(t, vals) + // db2 has 4 taps: two levels need a multiple of 8, so 13 pads + // to 16. + c, err := DaubechiesDWT(x, DB2, 2, DWTZeroPad) + if err != nil { + t.Fatalf("DaubechiesDWT zero-pad: %v", err) + } + if c.Len() != 16 { + t.Fatalf("padded coefficient length %d, want 16", c.Len()) + } + eC := 0.0 + for i := range 16 { + eC += c.FloatAt(i) * c.FloatAt(i) + } + if math.Abs(eIn-eC) > 1e-11*eIn { + t.Fatalf("padded energy %g vs %g", eC, eIn) + } + back, err := DaubechiesIDWT(c, DB2, 2, DWTZeroPad) + if err != nil { + t.Fatalf("DaubechiesIDWT zero-pad: %v", err) + } + if back.Len() != 16 { + t.Fatalf("reconstruction length %d, want the padded 16", back.Len()) + } + for i := range vals { + if math.Abs(back.FloatAt(i)-vals[i]) > 1e-12 { + t.Fatalf("reconstruction off at %d: %v vs %v", i, back.FloatAt(i), vals[i]) + } + } +} + +// TestDaubechiesErrors pins the input gates: family names, level +// counts, the length contract of both modes, and the shapes. +func TestDaubechiesErrors(t *testing.T) { + x := mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}) + if _, err := DaubechiesDWT(x, "db1", 1, DWTPeriodic); err == nil { + t.Error("db1 accepted; Haar keeps its own transform") + } + if _, err := DaubechiesDWT(x, "db9", 1, DWTPeriodic); err == nil { + t.Error("db9 accepted") + } + if _, err := DaubechiesDWT(x, "", 1, DWTPeriodic); err == nil { + t.Error("an empty family name accepted") + } + if _, err := DaubechiesDWT(x, DB2, 0, DWTPeriodic); err == nil { + t.Error("zero levels accepted") + } + // db2 has 4 taps: 8 samples support 2 levels, not 3. + if _, err := DaubechiesDWT(x, DB2, 3, DWTPeriodic); err == nil { + t.Error("levels deeper than the filter length accepted") + } + // 12 samples hold three 4-tap blocks at level 1 but 6 are not a + // multiple of 4 at level 2. + twelve := mustFloats(t, make([]float64, 12)) + if _, err := DaubechiesDWT(twelve, DB2, 1, DWTPeriodic); err != nil { + t.Errorf("12 samples at 1 level refused: %v", err) + } + if _, err := DaubechiesDWT(twelve, DB2, 2, DWTPeriodic); err == nil { + t.Error("a block of 6 accepted against a 4-tap filter") + } + // 14 samples are not a multiple of 4 even at level 1. + if _, err := DaubechiesDWT(mustFloats(t, make([]float64, 14)), DB2, 1, DWTPeriodic); err == nil { + t.Error("14 samples accepted against a 4-tap filter") + } + // The zero-padded mode takes exactly those. + if _, err := DaubechiesDWT(mustFloats(t, make([]float64, 14)), DB2, 1, DWTZeroPad); err != nil { + t.Errorf("zero-pad refused a short length: %v", err) + } + if _, err := DaubechiesDWT(mustFloats(t, make([]float64, 14)), DB2, 61, DWTZeroPad); err == nil { + t.Error("61 zero-pad levels accepted") + } + // db8's 16 taps at 60 levels would need a padded length of + // 16·2^59, which wraps past the addressable range: refused. + if _, err := DaubechiesDWT(mustFloats(t, make([]float64, 8)), DB8, 60, DWTZeroPad); err == nil { + t.Error("a padded length past the addressable range accepted") + } + // The band between the wrapped corner and the shallow depths holds + // spans the shift survives but no allocator does: each must refuse + // rather than panic in the work-buffer make. + for _, tc := range []struct { + family Daubechies + levels int + why string + }{ + {DB2, 60, "2^61 elements"}, + {DB8, 59, "2^62 elements"}, + {DB8, 50, "2^53 elements"}, + {DB8, 43, "2^46 elements, one past the makeslice ceiling"}, + } { + if _, err := DaubechiesDWT(mustFloats(t, make([]float64, 8)), tc.family, tc.levels, DWTZeroPad); err == nil { + t.Errorf("%s at %d levels accepted a padded span of %s", tc.family, tc.levels, tc.why) + } + } + if _, err := DaubechiesIDWT(x, DB2, 0, DWTPeriodic); err == nil { + t.Error("zero levels accepted by the inverse") + } + if _, err := DaubechiesIDWT(x, DB2, 1, DWTMode(9)); err == nil { + t.Error("unknown boundary mode accepted by the inverse") + } + if _, err := DaubechiesDWT(x, DB2, 1, DWTMode(9)); err == nil { + t.Error("unknown boundary mode accepted") + } + if _, err := DaubechiesDWT(mustFloats(t, []float64{1, 2, 3, 4}, 2, 2), DB2, 1, DWTPeriodic); err == nil { + t.Error("rank 2 accepted") + } + if _, err := DaubechiesDWT(mustComplexes(t, []complex128{1, 2, 3, 4, 5, 6, 7, 8}, 8), DB2, 1, DWTPeriodic); err == nil { + t.Error("complex accepted") + } + if _, err := DaubechiesDWT(mustFloats(t, nil), DB2, 1, DWTPeriodic); err == nil { + t.Error("empty accepted") + } + // The inverse repeats the length contract on its own input. + c, err := DaubechiesDWT(x, DB2, 1, DWTPeriodic) + if err != nil { + t.Fatal(err) + } + bad := mustFloats(t, make([]float64, 10)) + if _, err := DaubechiesIDWT(bad, DB2, 1, DWTPeriodic); err == nil { + t.Error("a length-10 coefficient vector accepted") + } + if _, err := DaubechiesIDWT(c, "db7", 1, DWTPeriodic); err == nil { + t.Error("inverting db2 coefficients as db7 accepted") + } +} diff --git a/signal/welch.go b/signal/welch.go new file mode 100644 index 0000000..60595b1 --- /dev/null +++ b/signal/welch.go @@ -0,0 +1,169 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Welch's power spectral density estimate: the signal is cut into +// overlapping segments, each segment is windowed, transformed and +// reduced to a periodogram, and the periodograms are averaged. The +// overlap and the window taper trade variance against spectral +// leakage; the averaging is what separates Welch from a single +// periodogram on noisy data. + +// welchMinSegmentsPerWorker is the smallest per-worker chunk of +// segments the Welch estimate splits for: below it a chunk's windowing +// and transforms no longer pay the worker spawn cost. +const welchMinSegmentsPerWorker = 4 + +// welchScratchMax bounds the per-worker transform scratch the pool +// keeps, in complex128 entries (1 MiB of memory per buffer at the +// cap). A bigger segment's scratch is dropped on return and rebuilt by +// the next borrower, so one huge segment size cannot pin its scratch +// on every processor; sync.Pool additionally forgets everything at +// each garbage collection. The estimate allocates at most one buffer +// per worker per call either way. +const welchScratchMax = 1 << 16 + +// welchScratch recycles the per-worker transform buffer through the +// pool: the pointer form keeps the Put from boxing a slice header on +// every return. +type welchScratch struct{ z []complex128 } + +var welchScratchPool = sync.Pool{ + New: func() any { return new(welchScratch) }, +} + +// WelchPSD estimates the one-sided power spectral density of the +// vector x sampled at fs hertz. segment is the length of each segment +// in samples, overlap the number of samples two neighbouring segments +// share (0 means none; segment/2 is the usual choice), and window +// names the taper: "hann", "hamming" or "box". The estimate averages +// (segment/2 + 1) bins from the (n−overlap)/(segment−overlap) +// segments the signal fills (a signal exactly one segment long is a +// single-segment periodogram) with the one-sided doubling applied +// away from DC and the Nyquist bin and the window's power normalising +// the scale, so a white noise sequence of variance σ² estimates σ² +// across the band. A short signal, a segment that exceeds it, a +// negative overlap not below the segment, or an unknown window name +// is an error. +func WelchPSD(x *core.Array, fs float64, segment, overlap int, window string) (freqs, psd *core.Array, err error) { + const name = "WelchPSD" + if x.NDim() != 1 { + return nil, nil, base.Errf("%s: the signal must be a vector, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, nil, base.Errf("%s: complex arrays are not supported", name) + } + n := x.Len() + if segment < 2 || segment > n { + return nil, nil, base.Errf("%s: segment must be between 2 and %d, got %d", name, n, segment) + } + if overlap < 0 || overlap >= segment { + return nil, nil, base.Errf("%s: overlap must be in [0, %d), got %d", name, segment, overlap) + } + if !(fs > 0) || math.IsInf(fs, 0) { + return nil, nil, base.Errf("%s: fs must be positive and finite, got %g", name, fs) + } + taper, err := windowTaper(window, segment) + if err != nil { + return nil, nil, base.Errf("%s: %w", name, err) + } + wPower := 0.0 + for i := range segment { + wPower += taper[i] * taper[i] + } + if wPower == 0 { + return nil, nil, base.Errf("%s: the window has zero power", name) + } + // The segments start every segment−overlap samples, so the count is + // the quotient, not one below it: a signal of exactly one segment + // fills one segment, not zero. + segments := (n - overlap) / (segment - overlap) + if segments < 1 { + return nil, nil, base.Errf("%s: %d samples fill only one %d-sample segment with overlap %d", + name, n, segment, overlap) + } + bins := segment/2 + 1 + // The magnitude rows come from the pooled scratch: every element + // is written by the segment pass below before the averaging pass + // reads it, so a recycled buffer behaves exactly like a fresh one. + mags := engine.GetFloat64Buf(segments * bins) + defer engine.PutFloat64Buf(mags) + src := widenFloats(x) + // The segments split across workers: every segment writes its own + // row of per-bin magnitudes and touches no shared state, and the + // averaging pass below then adds the rows in segment order, so the + // accumulation sequence over the segments is the serial one + // unchanged and no addend moves. + engine.ParallelMin(segments, welchMinSegmentsPerWorker, func(ss, se int) { + // One scratch segment per worker, borrowed from the pool and + // fully rewritten by the windowing below for every one of the + // segments it serves; wrapping it in a core.Array for FFT would + // copy these same values into a fresh slice only to run the + // same transform. + sc := welchScratchPool.Get().(*welchScratch) + if cap(sc.z) < segment { + sc.z = make([]complex128, segment) + } else { + sc.z = sc.z[:segment] + } + z := sc.z + for s := ss; s < se; s++ { + start := s * (segment - overlap) + for i := range segment { + z[i] = complex(src[start+i]*taper[i], 0) + } + transform(z, -1) + row := mags[s*bins : (s+1)*bins] + for k := range bins { + mag := base.AbsComplex(z[k]) + row[k] = mag * mag + } + } + // Retention cap: scratch above welchScratchMax is dropped + // rather than pooled, so a huge segment size cannot pin its + // buffer on every processor. + if cap(sc.z) > welchScratchMax { + sc.z = nil + } + welchScratchPool.Put(sc) + }) + power := make([]float64, bins) + for s := range segments { + row := mags[s*bins : (s+1)*bins] + for k := range bins { + power[k] += row[k] + } + } + scale := 1 / (float64(segments) * fs * wPower) + out, oerr := core.Zeros(core.Float, []int{bins}...) + if oerr != nil { + return nil, nil, oerr + } + psdVals := out.RawFloats() + for k := range bins { + p := power[k] * scale + if k > 0 && k < segment-k { + p *= 2 // one-sided doubling between DC and Nyquist + } + psdVals[k] = p + } + freqsArr, oerr := core.Zeros(core.Float, []int{bins}...) + if oerr != nil { + return nil, nil, oerr + } + freqVals := freqsArr.RawFloats() + for k := range bins { + freqVals[k] = float64(k) * fs / float64(segment) + } + return freqsArr, out, nil +} diff --git a/signal/windows.go b/signal/windows.go new file mode 100644 index 0000000..fcae004 --- /dev/null +++ b/signal/windows.go @@ -0,0 +1,212 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// The public window catalogue. Every builder returns a fresh +// preallocated []float64 of n samples. Two length conventions exist and +// the periodic flag picks between them: the symmetric window divides +// its argument by n−1, so its first and last samples coincide (the +// right shape for a finite impulse-response design) and the periodic +// window divides by n, which makes the sequence one exact period of +// its underlying continuous shape (the right shape for spectral +// estimates, where the segment is treated as one period and the +// doubled lobes of the symmetric tail would leak). periodic is the +// last argument, false everywhere it does not matter; a one-sample +// window is the single value 1 in both conventions. Every builder +// refuses n below 1. + +// WindowBox returns the untapered box: n ones, the window that +// filters nothing. The periodic flag changes nothing here and exists +// only for signature uniformity across the catalogue. +func WindowBox(n int, periodic bool) ([]float64, error) { + w, _, one, err := windowSetup("WindowBox", n, periodic) + if err != nil || one { + return w, err + } + for i := range w { + w[i] = 1 + } + return w, nil +} + +// WindowHann returns the Hann window, the raised cosine +// 0.5 − 0.5·cos(2πx), the gentlest of the generalised cosines: zero +// at the edges in both conventions, −6 dB per octave sidelobe roll-off. +func WindowHann(n int, periodic bool) ([]float64, error) { + return generalCosine("WindowHann", n, hannCoeffs, periodic) +} + +// WindowHamming returns the Hamming window, the raised cosine on a +// pedestal 0.54 − 0.46·cos(2πx): the nonzero pedestal cancels the +// Hann window's first sidelobe, at the cost of a floor the outer +// sidelobes never drop below. +func WindowHamming(n int, periodic bool) ([]float64, error) { + return generalCosine("WindowHamming", n, hammingCoeffs, periodic) +} + +// WindowBlackman returns the (exact) Blackman window +// 0.42 − 0.5·cos(2πx) + 0.08·cos(4πx): two cosines instead of +// Hann's one buy sidelobes below −58 dB at the price of a doubled +// main lobe. +func WindowBlackman(n int, periodic bool) ([]float64, error) { + return generalCosine("WindowBlackman", n, blackmanCoeffs, periodic) +} + +// WindowBlackmanHarris returns the four-term Blackman-Harris window, +// the minimum-sidelobe member of the generalised-cosine family with +// four terms: sidelobes below −92 dB, a main lobe three Hann lobes +// wide. The usual choice when dynamic range matters more than +// resolution. +func WindowBlackmanHarris(n int, periodic bool) ([]float64, error) { + return generalCosine("WindowBlackmanHarris", n, blackmanHarrisCoeffs, periodic) +} + +// WindowFlatTop returns the flat-top window, the five-term generalised +// cosine whose main lobe is flat to within a hundredth of a decibel: +// the amplitude of a spectral line reads true to the window's ripple +// no matter where the line falls between bins, which is what the wide +// lobe buys. The edge samples are slightly negative, so the window is +// for amplitude metrology, not for filtering. +func WindowFlatTop(n int, periodic bool) ([]float64, error) { + return generalCosine("WindowFlatTop", n, flatTopCoeffs, periodic) +} + +// WindowBartlett returns the Bartlett window, the triangle 1 − |2x − 1|: +// the piecewise-linear taper, zero at both edges, whose sidelobes sit +// between the box's and Hann's. It is the Fejér kernel of the box and +// is non-negative everywhere, which the generalised cosines are not. +func WindowBartlett(n int, periodic bool) ([]float64, error) { + w, den, one, err := windowSetup("WindowBartlett", n, periodic) + if err != nil || one { + return w, err + } + for i := range w { + w[i] = 1 - math.Abs(2*float64(i)/den-1) + } + return w, nil +} + +// WindowKaiser returns the Kaiser window of parameter beta: the +// modified Bessel taper I0(beta·sqrt(1 − r²))/I0(beta) over the +// normalised radius r = 2x − 1, the adjustable compromise between main +// lobe width and sidelobe height. beta 0 is the box; near 5 the +// sidelobes sit around −30 dB, near 9 around −60 dB, and the usual +// rule of thumb spends about 2.2·beta decibels of stopband. beta must +// be finite and non-negative. +func WindowKaiser(n int, beta float64, periodic bool) ([]float64, error) { + const name = "WindowKaiser" + if beta < 0 || math.IsNaN(beta) || math.IsInf(beta, 0) { + return nil, base.Errf("%s: beta must be finite and non-negative, got %g", name, beta) + } + w, den, one, err := windowSetup(name, n, periodic) + if err != nil || one { + return w, err + } + i0b := kaiserI0(beta) + for i := range w { + r := 2*float64(i)/den - 1 + // The radius can leave the unit disk by a rounding step at + // the edges; the squared radius is clamped so the square + // root stays real. + w[i] = kaiserI0(beta*math.Sqrt(math.Max(0, 1-r*r))) / i0b + } + return w, nil +} + +// WindowCosine returns the cosine (sine) window sin(πx): one positive +// half-period whose derivative vanishes at neither edge, the taper of +// the MDCT and of the minimum-tap Blackman derivations. +func WindowCosine(n int, periodic bool) ([]float64, error) { + w, den, one, err := windowSetup("WindowCosine", n, periodic) + if err != nil || one { + return w, err + } + for i := range w { + w[i] = math.Sin(math.Pi * float64(i) / den) + } + return w, nil +} + +// The generalised-cosine coefficient sets, with the signs carried in +// the table: sample i of the symmetric n-window is +// Σ_k c_k·cos(2πk·i/(n−1)). The flat-top set is written as the exact +// fractions of 19 the amplitude-calibration standard defines it by. +var ( + hannCoeffs = []float64{0.5, -0.5} + hammingCoeffs = []float64{0.54, -0.46} + blackmanCoeffs = []float64{0.42, -0.5, 0.08} + blackmanHarrisCoeffs = []float64{0.35875, -0.48829, 0.14128, -0.01168} + flatTopCoeffs = []float64{4.096 / 19, -7.916 / 19, 5.268 / 19, -1.588 / 19, 0.132 / 19} +) + +// generalCosine evaluates the generalised cosine family: the sum of +// signed cosine terms c_k over x = i/den, the denominator chosen by +// the symmetric or periodic convention. The term arguments keep the +// exact shape 2πk·i/den the package's spectral estimates have always +// fed their Hann and Hamming windows, so the periodic two-term +// members reproduce the legacy windowTaper outputs bit for bit. +func generalCosine(name string, n int, coeffs []float64, periodic bool) ([]float64, error) { + w, den, one, err := windowSetup(name, n, periodic) + if err != nil || one { + return w, err + } + for i := range w { + acc := coeffs[0] + for k, c := range coeffs[1:] { + // The explicit conversion is the FMA fence: the v4 build + // contracts a bare product-plus-add and the window bits + // drift a ulp from the portable build's, which the + // bit-identity pin holds. + acc += float64(c * math.Cos(2*math.Pi*float64(k+1)*float64(i)/den)) + } + w[i] = acc + } + return w, nil +} + +// windowSetup validates the requested length, prepares the window +// buffer and returns the denominator the window argument divides by: +// n−1 for the symmetric convention, n for the periodic one. The +// one-sample window is the constant 1 in both, so the flag-one return +// hands back that finished window and no builder reaches a zero +// denominator. +func windowSetup(name string, n int, periodic bool) (w []float64, den float64, one bool, err error) { + if n < 1 { + return nil, 0, false, base.Errf("%s: n must be at least 1, got %d", name, n) + } + if n == 1 { + return []float64{1}, 1, true, nil + } + if periodic { + return make([]float64, n), float64(n), false, nil + } + return make([]float64, n), float64(n - 1), false, nil +} + +// windowTaper builds the named window of the given length for the +// spectral estimates, routing through the catalogue's periodic forms: +// "hann" and "hamming" are WindowHann and WindowHamming at +// periodic=true, "box" is WindowBox, and the outputs are the same +// bits the dedicated loops this replaces produced for every length +// the callers accept (they refuse n below 2, where the two +// conventions differ). The legacy names stay because WelchPSD, STFT +// and Spectrogram publish them. +func windowTaper(window string, n int) ([]float64, error) { + switch window { + case "box": + return WindowBox(n, true) + case "hann": + return WindowHann(n, true) + case "hamming": + return WindowHamming(n, true) + default: + return nil, base.Errf("unknown window %q, want hann, hamming or box", window) + } +} diff --git a/signal/windows_test.go b/signal/windows_test.go new file mode 100644 index 0000000..1ab52cb --- /dev/null +++ b/signal/windows_test.go @@ -0,0 +1,298 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package signal + +import ( + "math" + "testing" +) + +// TestWindowTaperPeriodicBits pins the exact bit patterns the legacy +// windowTaper loops produced for WelchPSD, STFT and Spectrogram: the +// catalogue's periodic forms must reproduce them sample for sample, +// not merely to a tolerance. +func TestWindowTaperPeriodicBits(t *testing.T) { + hann := []uint64{ + 0x0, 0x3fa37ca1866b95d0, 0x3fc2bec333018866, 0x3fd3c10eaca8ab4e, + 0x3fdfffffffffffff, 0x3fe61f78a9abaa58, 0x3feb504f333f9de6, 0x3feec835e79946a3, + 0x3ff0000000000000, 0x3feec835e79946a4, 0x3feb504f333f9de7, 0x3fe61f78a9abaa5b, + 0x3fe0000000000001, 0x3fd3c10eaca8ab4c, 0x3fc2bec333018868, 0x3fa37ca1866b95e0, + } + hamming := []uint64{ + 0x3fb47ae147ae147c, 0x3fbd71a676273fcc, 0x3fcb7c4d2eecee22, 0x3fd74b3675e2db0b, + 0x3fe147ae147ae148, 0x3fe6e9c0ee04550a, 0x3febb048dd3a8707, 0x3feee1275a30da96, + 0x3ff0000000000000, 0x3feee1275a30da97, 0x3febb048dd3a8708, 0x3fe6e9c0ee04550d, + 0x3fe147ae147ae149, 0x3fd74b3675e2db0a, 0x3fcb7c4d2eecee24, 0x3fbd71a676273fd4, + } + for _, c := range []struct { + name string + want []uint64 + }{ + {"hann", hann}, + {"hamming", hamming}, + } { + w, err := windowTaper(c.name, 16) + if err != nil { + t.Fatalf("windowTaper(%s): %v", c.name, err) + } + for i := range w { + if got := math.Float64bits(w[i]); got != c.want[i] { + t.Fatalf("%s[%d]: bits %x, want %x", c.name, i, got, c.want[i]) + } + } + } + w, err := windowTaper("box", 16) + if err != nil { + t.Fatalf("windowTaper(box): %v", err) + } + for i := range w { + if got := math.Float64bits(w[i]); got != 0x3ff0000000000000 { + t.Fatalf("box[%d]: bits %x, want 3ff0000000000000", i, got) + } + } +} + +// TestWindowSymmetry pins the two length conventions: symmetric +// windows read the same at both ends and vanish there where their +// shape says so, periodic windows keep the raised tail the spectral +// estimates treat as one period. +func TestWindowSymmetry(t *testing.T) { + const n = 16 + for _, c := range []struct { + build func(n int, periodic bool) ([]float64, error) + edge float64 // the symmetric window's edge value + }{ + {WindowHann, 0}, + {WindowHamming, 0.08}, + {WindowBlackman, 0}, + {WindowBartlett, 0}, + {WindowCosine, 0}, + } { + sym, err := c.build(n, false) + if err != nil { + t.Fatalf("symmetric: %v", err) + } + per, err := c.build(n, true) + if err != nil { + t.Fatalf("periodic: %v", err) + } + if math.Abs(sym[0]-c.edge) > 1e-12 || math.Abs(sym[n-1]-c.edge) > 1e-12 { + t.Fatalf("symmetric edges %v, %v, want %v", sym[0], sym[n-1], c.edge) + } + for i := range n { + // Symmetry holds to the rounding of the per-sample + // argument, not bitwise: i and n−1−i compute their + // cosines from independently rounded arguments. + if math.Abs(sym[i]-sym[n-1-i]) > 1e-14 { + t.Fatalf("symmetric window off at %d: %v vs %v", i, sym[i], sym[n-1-i]) + } + if per[i] == per[n-1-i] && i != n-1-i { + t.Fatalf("periodic window mirrors its symmetric twin at %d", i) + } + } + // The conventions share only the first sample (argument 0); + // the periodic window continues to the raised tail, the + // symmetric one closes to the edge value. + if per[0] != sym[0] { + t.Fatalf("conventions disagree at the first sample: %v vs %v", per[0], sym[0]) + } + if c.edge == 0 && per[n-1] <= 0 { + t.Fatalf("periodic tail %v not above the symmetric edge", per[n-1]) + } + } + // Blackman-Harris and flat top share the symmetry, with the + // flat top's slightly negative edge. + bh, err := WindowBlackmanHarris(n, false) + if err != nil { + t.Fatal(err) + } + for i := range n { + if math.Abs(bh[i]-bh[n-1-i]) > 1e-14 { + t.Fatalf("Blackman-Harris off symmetry at %d: %v vs %v", i, bh[i], bh[n-1-i]) + } + } + ft, err := WindowFlatTop(n, false) + if err != nil { + t.Fatal(err) + } + wantEdge := -0.008 / 19 + if math.Abs(ft[0]-wantEdge) > 1e-14 { + t.Fatalf("flat top edge %v, want %v", ft[0], wantEdge) + } + // The flat top's flatness is a frequency-domain property: a + // spectral line reads the same amplitude wherever it falls + // between bins. The DTFT of the periodic window, sampled at + // bin offsets, must stay flat to the window's hundredth of a + // decibel. + ftPer, err := WindowFlatTop(16, true) + if err != nil { + t.Fatal(err) + } + dtft := func(bins float64) float64 { + var re, im float64 + for i, v := range ftPer { + ang := 2 * math.Pi * bins * float64(i) / float64(len(ftPer)) + re += v * math.Cos(ang) + im -= v * math.Sin(ang) + } + return math.Hypot(re, im) + } + w0 := dtft(0) + for _, bins := range []float64{0.25, 0.5} { + dev := math.Abs(dtft(bins)-w0) / w0 + if dev > 2e-3 { + t.Fatalf("flat top scalloping at %.2f bins: %g relative, want under 2e-3", bins, dev) + } + } +} + +// TestWindowKaiser pins the Kaiser taper: beta 0 is the box, the +// window peaks at its centre, and the underlying Bessel I0 hits its +// tabulated values. +func TestWindowKaiser(t *testing.T) { + box, err := WindowKaiser(9, 0, false) + if err != nil { + t.Fatal(err) + } + for i := range box { + if box[i] != 1 { + t.Fatalf("beta 0 sample %d = %v, want 1", i, box[i]) + } + } + w, err := WindowKaiser(11, 8.6, false) + if err != nil { + t.Fatal(err) + } + if w[5] != 1 { + t.Fatalf("Kaiser centre %v, want 1", w[5]) + } + for i := range w { + if math.Abs(w[i]-w[10-i]) > 1e-14 { + t.Fatalf("Kaiser off symmetry at %d: %v vs %v", i, w[i], w[10-i]) + } + } + // I0 against tabulated values. + for _, c := range []struct{ x, want float64 }{ + {0, 1}, + {1, 1.2660658777520084}, + {5, 27.23987182360444}, + } { + if got := kaiserI0(c.x); math.Abs(got-c.want) > 1e-12 { + t.Fatalf("I0(%g) = %v, want %v", c.x, got, c.want) + } + } + // A bigger beta is a stricter taper: lower at the same offset. + hard, err := WindowKaiser(11, 14, false) + if err != nil { + t.Fatal(err) + } + if hard[1] >= w[1] { + t.Fatalf("beta 14 edge %v not below beta 8.6 edge %v", hard[1], w[1]) + } +} + +// TestWindowShapes pins a few hand-computed sample values per +// builder. +func TestWindowShapes(t *testing.T) { + hann, err := WindowHann(5, false) + if err != nil { + t.Fatal(err) + } + for i, want := range []float64{0, 0.5, 1, 0.5, 0} { + if math.Abs(hann[i]-want) > 1e-12 { + t.Fatalf("Hann(5)[%d] = %v, want %v", i, hann[i], want) + } + } + bart, err := WindowBartlett(5, false) + if err != nil { + t.Fatal(err) + } + for i, want := range []float64{0, 0.5, 1, 0.5, 0} { + if math.Abs(bart[i]-want) > 1e-12 { + t.Fatalf("Bartlett(5)[%d] = %v, want %v", i, bart[i], want) + } + } + cos, err := WindowCosine(5, false) + if err != nil { + t.Fatal(err) + } + for i, want := range []float64{0, math.Sqrt2 / 2, 1, math.Sqrt2 / 2, 0} { + if math.Abs(cos[i]-want) > 1e-12 { + t.Fatalf("Cosine(5)[%d] = %v, want %v", i, cos[i], want) + } + } + black, err := WindowBlackman(5, false) + if err != nil { + t.Fatal(err) + } + if math.Abs(black[2]-1) > 1e-12 { + t.Fatalf("Blackman centre %v, want 1", black[2]) + } + // The periodic cosine window vanishes only at its first sample. + per, err := WindowCosine(8, true) + if err != nil { + t.Fatal(err) + } + if per[0] != 0 { + t.Fatalf("periodic cosine starts at %v, want 0", per[0]) + } + for i := 1; i < 8; i++ { + if per[i] <= 0 { + t.Fatalf("periodic cosine non-positive at %d: %v", i, per[i]) + } + } +} + +// TestWindowOneSample pins the one-sample convention: the constant 1 +// in both modes, for every builder. +func TestWindowOneSample(t *testing.T) { + for _, c := range []struct { + name string + build func(n int, periodic bool) ([]float64, error) + }{ + {"box", WindowBox}, + {"hann", WindowHann}, + {"hamming", WindowHamming}, + {"blackman", WindowBlackman}, + {"blackman-harris", WindowBlackmanHarris}, + {"flat-top", WindowFlatTop}, + {"bartlett", WindowBartlett}, + {"kaiser", func(n int, periodic bool) ([]float64, error) { return WindowKaiser(n, 6, periodic) }}, + {"cosine", WindowCosine}, + } { + for _, periodic := range []bool{false, true} { + w, err := c.build(1, periodic) + if err != nil { + t.Fatalf("%s periodic=%v: %v", c.name, periodic, err) + } + if len(w) != 1 || w[0] != 1 { + t.Fatalf("%s periodic=%v: one-sample window %v, want [1]", c.name, periodic, w) + } + } + } +} + +// TestWindowErrors pins the length and beta gates of the catalogue. +func TestWindowErrors(t *testing.T) { + for _, c := range []struct { + name string + build func() error + }{ + {"box n=0", func() error { _, err := WindowBox(0, true); return err }}, + {"hann n=-3", func() error { _, err := WindowHann(-3, false); return err }}, + {"hamming n=0", func() error { _, err := WindowHamming(0, true); return err }}, + {"blackman n=0", func() error { _, err := WindowBlackman(0, false); return err }}, + {"blackman-harris n=0", func() error { _, err := WindowBlackmanHarris(0, true); return err }}, + {"flat-top n=0", func() error { _, err := WindowFlatTop(0, false); return err }}, + {"bartlett n=0", func() error { _, err := WindowBartlett(0, true); return err }}, + {"kaiser n=0", func() error { _, err := WindowKaiser(0, 5, false); return err }}, + {"kaiser negative beta", func() error { _, err := WindowKaiser(8, -1, false); return err }}, + {"cosine n=0", func() error { _, err := WindowCosine(0, true); return err }}, + {"unknown taper name", func() error { _, err := windowTaper("hann2", 8); return err }}, + } { + if err := c.build(); err == nil { + t.Errorf("%s: want an error", c.name) + } + } +} diff --git a/spmd/arg.go b/spmd/arg.go new file mode 100644 index 0000000..65eb669 --- /dev/null +++ b/spmd/arg.go @@ -0,0 +1,248 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The arg reductions answer where an extremum lives in the global +// array: the same global index the single-array ArgMax and ArgMin +// answer, ties resolved by the earliest index, NaN elements skipped +// as missing. The shards compare values exactly, so the answer is the +// single-array answer's at any world size. + +// AllReduceArgShards answers the global index of the extremum (Max or +// Min) of one global array whose canonical pieces the ranks hold, +// every rank receiving the same index. Ties keep the earliest global +// index, the single-array walk's own rule; NaN elements never win, +// and an array with no candidate at all is an error, as the +// single-array reduction is. +func (w *World) AllReduceArgShards(local *core.Array, span Span, op Op) (int, error) { + ans, err := w.reduceArgShards(local, span, op, 0) + if err != nil { + return 0, err + } + s, err := w.broadcastScalar(ans, 0) + if err != nil { + return 0, err + } + return int(s.Int()), nil +} + +// ReduceArgShards is AllReduceArgShards with the answer on the root +// alone; every other rank receives nil. +func (w *World) ReduceArgShards(local *core.Array, span Span, op Op, root int) (*int, error) { + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + ans, err := w.reduceArgShards(local, span, op, root) + if err != nil { + return nil, err + } + if w.rank != root { + return nil, nil + } + i := int(ans.Int()) + return &i, nil +} + +// argCandidate is one rank's extremum candidate: the extremum's value +// and its global index. A slab with no candidate (all NaN, or no rows) +// carries ok=false. +type argCandidate struct { + ok bool + value float64 + iv int64 // the exact value for the integer dtypes + idx int + isInt bool +} + +// reduceArgShards runs the arg reduction; the answer exists at the +// root. +func (w *World) reduceArgShards(local *core.Array, span Span, op Op, root int) (core.Scalar, error) { + if err := w.status(); err != nil { + return core.Scalar{}, err + } + if op != Min && op != Max { + return core.Scalar{}, w.fail(base.Errf("spmd: arg reductions take Min or Max, got %s", op)) + } + if err := w.checkOneDimSpan(local, span); err != nil { + return core.Scalar{}, w.fail(err) + } + if local.Dtype() == core.Bool || local.Dtype() == core.Complex { + return core.Scalar{}, w.fail(base.Errf("spmd: %s of dtype %s has no ordering", op, local.Dtype())) + } + if span.Global == 0 { + return core.Scalar{}, w.fail(base.Errf("spmd: %s of an empty array has no answer", op)) + } + cand := argCandidateFor(local, span, op) + if w.rank != root { + if err := w.sendTo(root, tagShardsValues, encodeBlockValues(blockValues{ + first: cand.idx, kind: kindArg, i: []int64{cand.iv, int64(cand.idx)}, f: []float64{cand.value}, oks: []bool{cand.ok, cand.isInt}, + })); err != nil { + return core.Scalar{}, err + } + return core.Scalar{}, nil + } + all := make([]argCandidate, w.size) + all[w.rank] = cand + for r := range w.size { + if r == root { + continue + } + data, err := w.recvFrom(r, tagShardsValues) + if err != nil { + return core.Scalar{}, err + } + bv, err := decodeBlockValues(data) + if err != nil { + return core.Scalar{}, w.fail(err) + } + if bv.kind != kindArg || len(bv.i) != 2 || len(bv.f) != 1 || len(bv.oks) != 2 { + return core.Scalar{}, w.fail(base.Errf("spmd: arg shards disagree on the payload shape")) + } + all[r] = argCandidate{ok: bv.oks[0], value: bv.f[0], iv: bv.i[0], idx: int(bv.i[1]), isInt: bv.oks[1]} + } + greater := op == Max + best := -1 + for r, c := range all { + if !c.ok { + continue + } + if best < 0 { + best = r + continue + } + b := all[best] + better := false + if c.isInt == b.isInt { + if c.isInt { + better = (greater && c.iv > b.iv) || (!greater && c.iv < b.iv) + } else { + better = (greater && c.value > b.value) || (!greater && c.value < b.value) + } + } else { + return core.Scalar{}, w.fail(base.Errf("spmd: arg shards disagree on the payload kind")) + } + if better { + best = r + } + // A tie keeps the earlier rank's candidate, whose global index + // is the earlier one: the pieces are contiguous and ascend. + } + if best < 0 { + return core.Scalar{}, w.fail(base.Errf("spmd: %s shards have no candidate", op)) + } + winner := all[best] + return core.IntScalar(int64(winner.idx)), nil +} + +// argCandidateFor finds the slab's own first extremum, the serial +// walk's rule: the first strictly better element, NaNs skipped. The +// integer dtypes compare exactly in int64, never through float64, +// the reason the core keeps its own int64 walk. +func argCandidateFor(local *core.Array, span Span, op Op) argCandidate { + greater := op == Max + best := -1 + var bv float64 + var biv int64 + intExact := false + walkInt := func(v int64, i int) { + if best < 0 || (greater && v > biv) || (!greater && v < biv) { + best, biv, intExact = i, v, true + } + } + walkFloat := func(v float64, i int) { + if math.IsNaN(v) { + return + } + if best < 0 || (greater && v > bv) || (!greater && v < bv) { + best, bv, intExact = i, v, false + } + } + for i := range local.Len() { + switch local.Dtype() { + case core.Float: + walkFloat(local.FloatAt(i), i) + case core.Float32: + walkFloat(float64(local.RawFloat32s()[i]), i) + case core.Float16: + walkFloat(core.HalfToFloat64(local.RawHalves()[i]), i) + case core.Int: + walkInt(local.RawInts()[i], i) + case core.Int8: + walkInt(int64(local.RawInt8s()[i]), i) + case core.Uint8: + walkInt(int64(local.RawUint8s()[i]), i) + case core.Int16: + walkInt(int64(local.RawInt16s()[i]), i) + case core.Uint16: + walkInt(int64(local.RawUint16s()[i]), i) + case core.Int32: + walkInt(int64(local.RawInt32s()[i]), i) + case core.Uint32: + walkInt(int64(local.RawUint32s()[i]), i) + } + } + if best < 0 { + return argCandidate{} + } + if intExact { + return argCandidate{ok: true, iv: biv, idx: span.Lo + best, isInt: true} + } + return argCandidate{ok: true, value: bv, idx: span.Lo + best} +} + +// AllReduceArgSortShards answers the global permutation that sorts the +// whole array ascending: an Int array of global indices, the +// single-array ArgSort's own answer with its own tie and NaN +// placement, on every rank. The shards' values gather in global order +// and the core sorts them, so the permutation is the single-array +// one by construction. +func (w *World) AllReduceArgSortShards(local *core.Array, span Span) (*core.Array, error) { + ans, err := w.reduceArgSortShards(local, span, 0) + if err != nil { + return nil, err + } + return w.Broadcast(ans, 0) +} + +// ReduceArgSortShards is AllReduceArgSortShards with the answer on the +// root alone; every other rank receives nil. +func (w *World) ReduceArgSortShards(local *core.Array, span Span, root int) (*core.Array, error) { + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + return w.reduceArgSortShards(local, span, root) +} + +func (w *World) reduceArgSortShards(local *core.Array, span Span, root int) (*core.Array, error) { + if err := w.status(); err != nil { + return nil, err + } + if err := w.checkOneDimSpan(local, span); err != nil { + return nil, w.fail(err) + } + switch local.Dtype() { + case core.Int, core.Float32, core.Float16, core.Float: + default: + return nil, w.fail(base.Errf("spmd: ArgSort of dtype %s is not supported; convert with Astype", local.Dtype())) + } + gathered, err := w.gather(local, root) + if err != nil { + return nil, err + } + if w.rank != root { + return nil, nil + } + return core.ArgSort(gathered) +} + +// nanValue is the quiet NaN the no-candidate tests build their +// fixtures from. +var nanValue = math.NaN() diff --git a/spmd/arg_test.go b/spmd/arg_test.go new file mode 100644 index 0000000..a857126 --- /dev/null +++ b/spmd/arg_test.go @@ -0,0 +1,158 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The arg reductions answer the single-array walk's global index, ties +// by the earliest index, NaNs skipped, and the sharded sort answers +// the single-array permutation outright. + +// TestShardedArgMatchesSingleArray pins ArgMax and ArgMin against the +// core's own walk across dtypes, NaNs included, with the tie at the +// earliest index. +func TestShardedArgMatchesSingleArray(t *testing.T) { + for _, size := range []int{1, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537} { + for dt, whole := range fixtureDtypes(t, gn) { + switch dt { + case core.Float, core.Float32, core.Float16, core.Int, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32: + default: + continue + } + for _, op := range []Op{Max, Min} { + var want int + var err error + if op == Max { + want, err = core.ArgMax(whole) + } else { + want, err = core.ArgMin(whole) + } + if err != nil { + t.Fatal(err) + } + err = Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceArgShards(local, span, op) + if err != nil { + return err + } + if got != want { + t.Fatalf("size %d gn %d %s %s: sharded index %d against single-array %d", + w.Size(), gn, dt, op, got, want) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d %s %s: %v", size, gn, dt, op, err) + } + } + } + } + } +} + +// TestShardedArgNoCandidate: an all-NaN array is an error on every +// rank, the single-array walk's own refusal. +func TestShardedArgNoCandidate(t *testing.T) { + nans := make([]float64, 100) + for i := range nans { + nans[i] = math2NaN() + } + whole, err := core.FromFloats(nans, 100) + if err != nil { + t.Fatal(err) + } + err = Launch(3, func(w *World) error { + span := mustPartition(t, 100, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + _, err := w.AllReduceArgShards(local, span, Max) + return err + }) + if err == nil { + t.Fatal("an all-NaN array answered an index") + } +} + +func math2NaN() float64 { return nanValue } + +// TestShardedArgSortMatchesSingleArray: the sharded permutation is +// the single-array ArgSort's own, values, ties, NaN placement and all. +func TestShardedArgSortMatchesSingleArray(t *testing.T) { + for _, size := range []int{1, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537} { + for _, dt := range []core.Dtype{core.Float, core.Float32, core.Float16, core.Int} { + whole := fixtureDtypes(t, gn)[dt] + want, err := core.ArgSort(whole) + if err != nil { + t.Fatal(err) + } + err = Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceArgSortShards(local, span) + if err != nil { + return err + } + if got.Len() != want.Len() { + t.Fatalf("size %d gn %d %s: permutation of %d against %d", + w.Size(), gn, dt, got.Len(), want.Len()) + } + for i := range want.Len() { + if got.RawInts()[i] != want.RawInts()[i] { + t.Fatalf("size %d gn %d %s: permutation differs at %d: %d against %d", + w.Size(), gn, dt, i, got.RawInts()[i], want.RawInts()[i]) + } + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err) + } + } + } + } +} + +// TestShardedArgOverTCP runs the arg and the sort over real +// connections. +func TestShardedArgOverTCP(t *testing.T) { + const gn = 65537 + whole := fixtureDtypes(t, gn)[core.Float] + wantMax, err := core.ArgMax(whole) + if err != nil { + t.Fatal(err) + } + wantSort, err := core.ArgSort(whole) + if err != nil { + t.Fatal(err) + } + runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + gotMax, err := w.AllReduceArgShards(local, span, Max) + if err != nil { + return err + } + if gotMax != wantMax { + t.Fatalf("rank %d: arg %d against %d", w.Rank(), gotMax, wantMax) + } + gotSort, err := w.AllReduceArgSortShards(local, span) + if err != nil { + return err + } + for i := range wantSort.Len() { + if gotSort.RawInts()[i] != wantSort.RawInts()[i] { + t.Fatalf("rank %d: the permutation differs at %d over TCP", w.Rank(), i) + } + } + return nil + }) +} diff --git a/spmd/bench_tcp_test.go b/spmd/bench_tcp_test.go new file mode 100644 index 0000000..7a6d5d9 --- /dev/null +++ b/spmd/bench_tcp_test.go @@ -0,0 +1,177 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "net" + "sync" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The TCP benchmarks run one loopback world of four ranks over real +// framed connections, every rank inside this process: the network path +// is the transport the routed frame pool lives on, and the scheduler +// noise of separate processes stays out. Each benchmark runs twice, as +// alternating sub-benchmarks with the pool on and off through the +// package switch, so both variants of the A/B share one binary and one +// machine state; the verdict comes from the medians across repeated +// rounds. + +// benchmarkTCPBothWays runs one TCP benchmark as two sub-benchmarks, +// the routed frame pool on, then off. +func benchmarkTCPBothWays(b *testing.B, warm func(w *World) (func() error, error)) { + for _, way := range []struct { + name string + pooled bool + }{ + {"pool", true}, + {"base", false}, + } { + b.Run(way.name, func(b *testing.B) { + was := framePoolEnabled + framePoolEnabled = way.pooled + defer func() { framePoolEnabled = was }() + benchmarkTCP(b, warm) + }) + } +} + +// benchmarkTCP assembles the loopback world, runs warm once on every +// rank before the clock starts, then the round it answers b.N times on +// every rank, so one operation is one round of the world. +func benchmarkTCP(b *testing.B, warm func(w *World) (func() error, error)) { + b.Helper() + b.ReportAllocs() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + b.Fatal(err) + } + defer ln.Close() + const size = 4 + var wg sync.WaitGroup + errs := make([]error, size) + wg.Go(func() { + errs[0] = benchTCPRank(b, func() (*World, error) { + return listen(ln, size, Options{Timeout: 2 * time.Minute}) + }, warm) + }) + for r := 1; r < size; r++ { + wg.Go(func() { + errs[r] = benchTCPRank(b, func() (*World, error) { + return Join(ln.Addr().String(), Options{Timeout: 2 * time.Minute}) + }, warm) + }) + } + wg.Wait() + for r, err := range errs { + if err != nil { + b.Fatalf("rank %d: %v", r, err) + } + } +} + +func benchTCPRank(b *testing.B, build func() (*World, error), warm func(w *World) (func() error, error)) error { + w, err := build() + if err != nil { + return err + } + defer w.Close() + round, err := warm(w) + if err != nil { + return err + } + if err := w.Barrier(); err != nil { + return err + } + for range b.N { + if err := round(); err != nil { + return err + } + } + return w.Barrier() +} + +// BenchmarkTCPAllReduceShards is the measurement the pool item names: +// four ranks folding one million float64s to one shared scalar. Its +// frames stay small and every one of them is addressed to rank 0, so +// none of them routes through the hub's outbox: it is the battery's +// control, and the pool should leave it untouched. +func BenchmarkTCPAllReduceShards(b *testing.B) { + const gn = 1 << 20 + whole := fixtureArray(gn) + benchmarkTCPBothWays(b, func(w *World) (func() error, error) { + span := mustPartition(b, gn, w.Size(), w.Rank()) + local, err := dealRows(whole, span) + if err != nil { + return nil, err + } + return func() error { + _, err := w.AllReduceShards(local, span, Sum) + return err + }, nil + }) +} + +// BenchmarkTCPMovement is the movement battery of the TCP tests: +// Scatter, AllGather and a Broadcast whose root is a non-hub rank, so +// its 560 kilobyte frames route through the hub's pump and outbox, the +// path the pool serves. The canonical pieces of 70001 rows are the +// ones that partition gives, empty pieces included, exactly as the +// test the battery mirrors runs them. +func BenchmarkTCPMovement(b *testing.B) { + const gn = 70001 + whole := fixtureArray(gn) + benchmarkTCPBothWays(b, func(w *World) (func() error, error) { + return func() error { + local, _, err := w.Scatter(whole, 0) + if err != nil { + return err + } + got, err := w.AllGather(local) + if err != nil { + return err + } + _, err = w.Broadcast(got, 2) + return err + }, nil + }) +} + +// BenchmarkTCPReduce walks the array reduce, whose chunks travel rank +// to owner through the hub: six routed frames of half a mebibyte a +// round, the heaviest routed traffic in the battery. +func BenchmarkTCPReduce(b *testing.B) { + const per = 1 << 18 + a, err := core.FromFloats(fixture(per), per) + if err != nil { + b.Fatal(err) + } + benchmarkTCPBothWays(b, func(w *World) (func() error, error) { + return func() error { + _, err := w.AllReduce(a, Sum) + return err + }, nil + }) +} + +// BenchmarkTCPExchangeHalos serves the stencil workloads: edge rows +// travel to both neighbours through the hub. The grid sits on the +// canonical partition's own boundaries, four blocks of 65536 rows, so +// every rank holds a real piece, and the halo width puts the edges at +// 128 kilobytes, a size the pool retains: four routed frames a round. +func BenchmarkTCPExchangeHalos(b *testing.B) { + const rows, width, halos = 262144, 8, 2048 + whole := globalRows(rows, width) + benchmarkTCPBothWays(b, func(w *World) (func() error, error) { + span := mustPartition(b, rows, w.Size(), w.Rank()) + local := dealRows2D(b, whole, span.Lo, span.Hi) + return func() error { + _, _, err := w.ExchangeHalos(local, halos) + return err + }, nil + }) +} diff --git a/spmd/bench_test.go b/spmd/bench_test.go new file mode 100644 index 0000000..dcd7b6f --- /dev/null +++ b/spmd/bench_test.go @@ -0,0 +1,118 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "strconv" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The collective benchmarks time one world of four ranks running the +// named collective back to back: every rank loops b.N times, so the +// per-operation number is one round of the collective across the +// world. In-process links put the channel and the arithmetic on the +// clock, not the network. + +func BenchmarkAllReduceShards(b *testing.B) { + const gn = 1 << 20 + whole := fixtureArray(gn) + if err := Launch(4, func(w *World) error { + span := mustPartition(b, gn, w.Size(), w.Rank()) + local, err := dealRows(whole, span) + if err != nil { + return err + } + for i := 0; i < b.N; i++ { + if _, err := w.AllReduceShards(local, span, Sum); err != nil { + return err + } + } + return w.Barrier() + }); err != nil { + b.Fatal(err) + } +} + +func BenchmarkAllReduce(b *testing.B) { + const per = 1 << 18 + a, err := core.FromFloats(fixture(per), per) + if err != nil { + b.Fatal(err) + } + if err := Launch(4, func(w *World) error { + for i := 0; i < b.N; i++ { + if _, err := w.AllReduce(a, Sum); err != nil { + return err + } + } + return w.Barrier() + }); err != nil { + b.Fatal(err) + } +} + +// BenchmarkReduce times the same-shaped-arrays fold on one million +// float64 elements across world sizes 2, 4 and 8 for Sum and Min: the +// fold runs as the chunks land, each arriving chunk joining the +// running accumulator at its rank's turn. +func BenchmarkReduce(b *testing.B) { + const n = 1 << 20 + a, err := core.FromFloats(fixture(n), n) + if err != nil { + b.Fatal(err) + } + for _, size := range []int{2, 4, 8} { + for _, op := range []Op{Sum, Min} { + b.Run(op.String()+"/"+strconv.Itoa(size), func(b *testing.B) { + b.ReportAllocs() + if err := Launch(size, func(w *World) error { + for i := 0; i < b.N; i++ { + if _, err := w.Reduce(a, op, 0); err != nil { + return err + } + } + return w.Barrier() + }); err != nil { + b.Fatal(err) + } + }) + } + } +} + +func BenchmarkBroadcast(b *testing.B) { + a := fixtureArray(1 << 18) + if err := Launch(4, func(w *World) error { + for i := 0; i < b.N; i++ { + if _, err := w.Broadcast(a, 0); err != nil { + return err + } + } + return w.Barrier() + }); err != nil { + b.Fatal(err) + } +} + +func BenchmarkScatterGather(b *testing.B) { + const gn = 1 << 20 + a := fixtureArray(gn) + if err := Launch(4, func(w *World) error { + for i := 0; i < b.N; i++ { + local, span, err := w.Scatter(a, 0) + if err != nil { + return err + } + if _, err := w.Gather(local, 0); err != nil { + return err + } + _ = span + } + return w.Barrier() + }); err != nil { + b.Fatal(err) + } +} diff --git a/spmd/collective.go b/spmd/collective.go new file mode 100644 index 0000000..ad8c5d9 --- /dev/null +++ b/spmd/collective.go @@ -0,0 +1,285 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The movement collectives carry bits and carry nothing else: no +// arithmetic happens in transit, so their answers are deterministic by +// construction. The wire is the only form that travels, and every +// collective's join order is a function of the data, never of the +// order frames happen to arrive in. + +// Broadcast delivers root's array to every rank of the world. The root +// gets the array it passed; every other rank gets an equal copy, bits +// included. A world of one rank hands the array back. +func (w *World) Broadcast(a *core.Array, root int) (*core.Array, error) { + if err := w.status(); err != nil { + return nil, err + } + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + if w.size == 1 { + return a, nil + } + if w.rank == root { + wire, err := encodeArray(nil, a) + if err != nil { + return nil, w.fail(err) + } + for r := range w.size { + if r == root { + continue + } + if err := w.sendTo(r, tagBroadcast, wire); err != nil { + return nil, err + } + } + return a, nil + } + data, err := w.recvFrom(root, tagBroadcast) + if err != nil { + return nil, err + } + got, err := decodeWire(data) + if err != nil { + return nil, w.fail(err) + } + return got, nil +} + +// Scatter deals root's array out along its first dimension: rank r +// receives exactly the canonical partition's piece of the axis, so the +// world's data lands on the boundaries the shard reductions compose +// on. The global shape travels first, so every rank can name its own +// piece before any payload moves; the span that comes back names it. +func (w *World) Scatter(global *core.Array, root int) (*core.Array, Span, error) { + if err := w.status(); err != nil { + return nil, Span{}, err + } + if err := checkRoot(root, w.size); err != nil { + return nil, Span{}, w.fail(err) + } + if global.NDim() == 0 { + return nil, Span{}, w.fail(base.Errf("spmd: Scatter needs a dimension to deal along")) + } + var shape []int + if w.rank == root { + shape = global.Shape() + head, err := encodeHead(nil, global.Dtype(), shape) + if err != nil { + return nil, Span{}, w.fail(err) + } + for r := range w.size { + if r == root { + continue + } + if err := w.sendTo(r, tagScatterHead, head); err != nil { + return nil, Span{}, err + } + } + } else { + data, err := w.recvFrom(root, tagScatterHead) + if err != nil { + return nil, Span{}, err + } + _, shape, err = decodeHead(data) + if err != nil { + return nil, Span{}, w.fail(err) + } + } + if len(shape) == 0 { + return nil, Span{}, w.fail(base.Errf("spmd: the dealt head names no dimension to deal along")) + } + gn := shape[0] + rest, ok := elementCount(shape[1:]) + if !ok { + return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %v overflows the element count", shape)) + } + full, ok := elementCount(shape) + if !ok { + return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %v overflows the element count", shape)) + } + if w.rank == root && full != global.Len() { + return nil, Span{}, w.fail(base.Errf("spmd: Scatter's shape %s names %d elements against the array's %d", + base.ShapeText(shape), full, global.Len())) + } + span, err := Partition(gn, w.size, w.rank) + if err != nil { + return nil, Span{}, w.fail(err) + } + slabShape := make([]int, 0, len(shape)) + slabShape = append(slabShape, span.Len()) + slabShape = append(slabShape, shape[1:]...) + local, err := func() (*core.Array, error) { + if w.rank == root { + wire, err := encodePart(nil, global, slabShape, span.Lo*rest, span.Len()*rest) + if err != nil { + return nil, err + } + for r := range w.size { + if r == root { + continue + } + s, err := Partition(gn, w.size, r) + if err != nil { + return nil, err + } + wire, err := encodePart(nil, global, + append([]int{s.Len()}, shape[1:]...), s.Lo*rest, s.Len()*rest) + if err != nil { + return nil, err + } + if err := w.sendTo(r, tagScatter, wire); err != nil { + return nil, err + } + } + return decodeWire(wire) + } + wire, err := w.recvFrom(root, tagScatter) + if err != nil { + return nil, err + } + return decodeWire(wire) + }() + if err != nil { + return nil, Span{}, w.fail(err) + } + return local, span, nil +} + +// Gather raises the dealt pieces back into the global array on root, +// joining them in rank order, which is the canonical partition's order. +// Ranks other than root answer with nil: they hold their piece, the +// root holds the whole. +func (w *World) Gather(local *core.Array, root int) (*core.Array, error) { + global, err := w.gather(local, root) + if err != nil { + return nil, err + } + if w.rank != root { + return nil, nil + } + return global, nil +} + +// gather is Gather's body; AllGather shares it with a fixed root. +func (w *World) gather(local *core.Array, root int) (*core.Array, error) { + if err := w.status(); err != nil { + return nil, err + } + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + if local.NDim() == 0 { + return nil, w.fail(base.Errf("spmd: Gather needs a dimension to raise along")) + } + wire, err := encodeArray(nil, local) + if err != nil { + return nil, w.fail(err) + } + // A world of one rank is already the whole. + if w.size == 1 { + return local, nil + } + // The pieces travel to the root; the root takes them in rank + // order, which is the partition's order, whatever order they + // arrive in. Its own piece sits at its own rank's place in the + // join. + if w.rank != root { + if err := w.sendTo(root, tagGather, wire); err != nil { + return nil, err + } + return nil, nil + } + pieces := make([][]byte, 0, w.size) + for r := range w.size { + if r == root { + pieces = append(pieces, wire) + continue + } + data, err := w.recvFrom(r, tagGather) + if err != nil { + return nil, err + } + pieces = append(pieces, data) + } + return w.joinPieces(pieces) +} + +// joinPieces builds the global array from the pieces' wire forms: the +// heads must agree in dtype and in every dimension but the first, and +// the payloads concatenate in the order given. +func (w *World) joinPieces(pieces [][]byte) (*core.Array, error) { + var dt core.Dtype + var trailing []int + total := 0 + var payload []byte + for i, piece := range pieces { + pdt, shape, err := decodeHead(piece) + if err != nil { + return nil, w.fail(err) + } + if len(shape) == 0 { + return nil, w.fail(base.Errf("spmd: piece %d has no dimension to raise along", i)) + } + if i == 0 { + dt = pdt + trailing = shape[1:] + } else { + if pdt != dt { + return nil, w.fail(base.Errf("spmd: piece %d is %s against %s", i, pdt, dt)) + } + if len(shape) != len(trailing)+1 { + return nil, w.fail(base.Errf("spmd: piece %d has %d dimensions against %d", i, len(shape), len(trailing)+1)) + } + for d := range trailing { + if shape[d+1] != trailing[d] { + return nil, w.fail(base.Errf("spmd: piece %d disagrees in dimension %d: %d against %d", + i, d+1, shape[d+1], trailing[d])) + } + } + } + if shape[0] > math.MaxInt-total { + return nil, w.fail(base.Errf("spmd: joining the pieces overflows the first dimension at piece %d", i)) + } + total += shape[0] + payload = append(payload, piece[2+8*len(shape):]...) + } + globalShape := make([]int, 0, len(trailing)+1) + globalShape = append(globalShape, total) + globalShape = append(globalShape, trailing...) + head, err := encodeHead(nil, dt, globalShape) + if err != nil { + return nil, w.fail(err) + } + got, err := decodeWire(append(head, payload...)) + if err != nil { + return nil, w.fail(err) + } + return got, nil +} + +// AllGather raises every dealt piece into the whole on every rank: one +// world, one array, identical bits everywhere. +func (w *World) AllGather(local *core.Array) (*core.Array, error) { + whole, err := w.gather(local, 0) + if err != nil { + return nil, err + } + return w.Broadcast(whole, 0) +} + +func checkRoot(root, size int) error { + if root < 0 || root >= size { + return base.Errf("spmd: root %d is outside the world of %d ranks", root, size) + } + return nil +} diff --git a/spmd/collective_test.go b/spmd/collective_test.go new file mode 100644 index 0000000..b28fadd --- /dev/null +++ b/spmd/collective_test.go @@ -0,0 +1,211 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "math" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// fixture builds a deterministic float64 slice whose values carry +// magnitude spread and a fixed pattern, so a shuffled or truncated +// movement shows up in the bits. +func fixture(n int) []float64 { + v := make([]float64, n) + for i := range v { + v[i] = float64((i*7919)%211-105) / 7.0 + if i%97 == 0 { + v[i] = math.Inf(1) + } + if i%89 == 0 { + v[i] = math.Copysign(0, -1) + } + } + return v +} + +func fixtureArray(n int) *core.Array { + a, err := core.FromFloats(fixture(n), n) + if err != nil { + panic(err) + } + return a +} + +// TestBroadcastCarriesTheBits: whatever rank originates the broadcast, +// every rank ends holding the root's exact bits, over the in-process +// world. +func TestBroadcastCarriesTheBits(t *testing.T) { + const n = 1000 + for _, size := range []int{1, 2, 3, 5, 8} { + for _, root := range []int{0, size - 1} { + err := Launch(size, func(w *World) error { + want := fixtureArray(n) + got, err := w.Broadcast(want, root) + if err != nil { + return err + } + if w.Rank() == root { + if got != want { + t.Fatalf("rank %d: the root did not keep its own array", w.Rank()) + } + return nil + } + if !sameBits(want, got) { + t.Fatalf("rank %d: broadcast bits differ from rank %d's", w.Rank(), root) + } + return nil + }) + if err != nil { + t.Fatalf("size %d root %d: %v", size, root, err) + } + } + } +} + +// TestBroadcastCarriesEveryDtype walks the narrow, half and boolean +// element types through the movement path once: the wire is the only +// form that travels, and every dtype has to survive it. +func TestBroadcastCarriesEveryDtype(t *testing.T) { + cases := []*core.Array{ + mk(core.FromBools([]bool{true, false, true}, 3)), + mk(core.HalvesFromArray([]uint16{0x0001, 0x7bff, 0xfc00}, 3)), + mk(core.FromInt8s([]int8{-128, 127, 0}, 3)), + mk(core.FromUint32s([]uint32{4294967295, 0, 7}, 3)), + mk(core.FromComplexes([]complex128{1 + 2i, -0i}, 2)), + } + err := Launch(3, func(w *World) error { + for i, want := range cases { + got, err := w.Broadcast(want, i%w.Size()) + if err != nil { + return err + } + if !sameBits(want, got) { + t.Fatalf("rank %d case %d: bits differ", w.Rank(), i) + } + } + return nil + }) + if err != nil { + t.Fatal(err) + } +} + +// TestScatterGatherRoundTrip deals a global array out and raises it +// back: every rank's piece is the canonical partition's own cut, and +// the gathered whole is the original's exact bits. +func TestScatterGatherRoundTrip(t *testing.T) { + for _, size := range []int{1, 2, 3, 5, 8} { + for _, gn := range []int{0, 1, 100, 65537, 200001} { + for _, root := range []int{0, size - 1} { + err := Launch(size, func(w *World) error { + src, err := core.FromFloats(fixture(gn*3), gn, 3) + if err != nil { + return err + } + want := fixture(gn * 3) + local, span, err := w.Scatter(src, root) + if err != nil { + return err + } + if span != mustPartition(t, gn, w.Size(), w.Rank()) { + t.Fatalf("rank %d: span %v against the partition's %v", + w.Rank(), span, mustPartition(t, gn, w.Size(), w.Rank())) + } + if local.Len() != span.Len()*3 { + t.Fatalf("rank %d: slab of %d elements for a span of %d", + w.Rank(), local.Len(), span.Len()) + } + // Every element the slab carries is the fixture's own + // value at its global index. + for i := 0; i < local.Len(); i++ { + if got := local.FloatAt(i); got != want[span.Lo*3+i] { + t.Fatalf("rank %d element %d: %v", w.Rank(), i, got) + } + } + back, err := w.Gather(local, root) + if err != nil { + return err + } + if w.Rank() == root { + if !sameBits(src, back) { + t.Fatalf("rank %d: the gathered whole differs from the dealt array", w.Rank()) + } + } else if back != nil { + t.Fatalf("rank %d: gather returned a whole to a non-root", w.Rank()) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d root %d: %v", size, gn, root, err) + } + } + } + } +} + +// TestAllGatherRebuildsEverywhere: one deal, one raise, and every rank +// holds the whole. +func TestAllGatherRebuildsEverywhere(t *testing.T) { + const gn = 1000 + err := Launch(4, func(w *World) error { + local, span, err := w.Scatter(fixtureArray(gn), 0) + if err != nil { + return err + } + whole, err := w.AllGather(local) + if err != nil { + return err + } + want := fixtureArray(gn) + if !sameBits(want, whole) { + t.Fatalf("rank %d: the rebuilt whole differs", w.Rank()) + } + if whole.Len() != gn || span.Global != gn { + t.Fatalf("rank %d: whole %d against global %d", w.Rank(), whole.Len(), span.Global) + } + return nil + }) + if err != nil { + t.Fatal(err) + } +} + +// TestMovementOverTCP runs the same movement battery over real +// connections, because the contract says the two transports are one +// machine. +func TestMovementOverTCP(t *testing.T) { + for _, gn := range []int{0, 100, 70001} { + t.Run("", func(t *testing.T) { + runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error { + src := fixtureArray(gn) + local, span, err := w.Scatter(src, 0) + if err != nil { + return err + } + if span != mustPartition(t, gn, w.Size(), w.Rank()) { + t.Fatalf("rank %d: span %v", w.Rank(), span) + } + whole, err := w.AllGather(local) + if err != nil { + return err + } + if !sameBits(src, whole) { + t.Fatalf("rank %d: the whole differs over TCP", w.Rank()) + } + round, err := w.Broadcast(whole, 2) + if err != nil { + return err + } + if !sameBits(whole, round) { + t.Fatalf("rank %d: broadcast over TCP differs", w.Rank()) + } + return nil + }) + }) + } +} diff --git a/spmd/doc.go b/spmd/doc.go new file mode 100644 index 0000000..7be7486 --- /dev/null +++ b/spmd/doc.go @@ -0,0 +1,42 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package spmd runs one program on many ranks: explicit SPMD worlds +// with the collectives of scientific distributed computing, over TCP +// between machines or in process within one. +// +// # The determinism contract +// +// Every collective answers bit-identical results whatever the world +// size, the machine, the worker count or the order messages arrive in. +// The rule that buys this is one: the order of a reduction is a +// function of the data, never of the topology. Family B reductions +// (ReduceShards, AllReduceShards) cut the global length at the fold's +// own block boundaries, fold each block where its elements live, and +// combine the block partials with the same balanced tree the +// single-array fold uses, so a sharded reduction and the single-array +// Sum are one computation and one answer, for one rank or for fifty. +// Family A reductions (Reduce, AllReduce) fold same-shaped arrays in +// rank index order, which the program itself fixes. No collective +// every combines partials in the order they happen to arrive. +// +// # Worlds +// +// Launch builds a world of size ranks in one process, one goroutine +// per rank, over channels: the development surface, the test surface +// and the single-machine surface. Listen and Join build the same world +// over TCP, rank 0 listening and every other rank dialling; ranks are +// assigned in dial order. The sharded reductions' results never +// depend on the assignment; the same-shaped arrays' fold order is the +// rank order the program fixes, so a program that wants bit-reproducible +// family A runs keeps its rank assignment stable. Collectives are bulk synchronous: every rank calls the +// same collective in the same order, as an MPI program does. +// +// Any failure, a deadline included, fails the whole world: the rank +// that saw it and every rank that then touches the world or waits on +// it receive an error, and no collective ever returns a partial +// numeric result. A failed world stays failed. +// +// The package is pure Go on the standard library: no cgo, no +// dependency, no build tag. +package spmd diff --git a/spmd/framepool.go b/spmd/framepool.go new file mode 100644 index 0000000..dd46ee2 --- /dev/null +++ b/spmd/framepool.go @@ -0,0 +1,110 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import "sync" + +// The routed frame pool recycles the payload buffers of the frames the +// hub routes on, the frames whose destination is another rank. Its +// safety rests on one ownership chain and nothing else: the pump takes +// a buffer when it reads a routed frame, the frame hands the buffer to +// the destination link's outbox channel, and the drain, the chain's +// single consumer, writes the frame and returns the buffer at once, +// on a landed write and on a benign skip at a departing peer alike. +// No frame travels a pending stash with the pool's mark on, no +// collective ever sees a pooled buffer, and no count, atom or second +// return point exists: a channel handoff is the ownership transfer. +// A buffer the chain drops on the way, a mislabelled frame or a world +// that ends mid-route, is left to the collector, which makes the loss +// a pool miss and never a double return. This single owner is why the +// first pool attempt's failure mode, a payload retired while a decode +// still read it, cannot arise here. +// +// Two bounds keep the pool from pinning memory. A buffer above the +// retention ceiling is never kept, so one enormous frame cannot pin +// the heap across rounds, and the pool holds at most a fixed number +// of buffers, so a flood of mid-sized ones cannot grow without end. +// Both refusals are pool misses: the buffer goes to the collector, +// which is exactly the behaviour the package had before the pool. + +const ( + // framePoolCeiling is the largest buffer the pool retains, one + // mebibyte. A payload beyond it is a pool miss by rule. + framePoolCeiling = int64(1 << 20) + // framePoolBound is the largest number of buffers the pool holds + // at once; a return that would exceed it is dropped. + framePoolBound = 32 +) + +// framePoolEnabled switches the pool at package level. It exists for +// the A/B measurement, which flips it between alternating +// sub-benchmarks of one binary; the library itself always runs pooled. +var framePoolEnabled = true + +// framePool is the free list behind the pool: payload buffers waiting +// for the next routed frame whose length fits. +type framePool struct { + mu sync.Mutex + free [][]byte +} + +// routedFrames is the process's one pool. The worlds of one process +// share it, which is what makes its bounds process-wide facts. +var routedFrames framePool + +// takeFrameBuffer answers the buffer a routed frame's payload reads +// into, or nil when the frame allocates as usual, outside the pool's +// chain: when the pool is switched off, when the frame stays at the +// hub, or when the length the header announced is beyond the ceiling. +// A qualified frame rides the chain whatever the free list holds: the +// buffer comes from the pool when one fits, and a fresh one joins the +// chain in its place, so the drain has a buffer to return and the +// pool warms with the traffic it serves. The link read calls this +// with the header already parsed, before the payload moves. +func takeFrameBuffer(m message, length int64) []byte { + if !framePoolEnabled || m.dest == 0 || length <= 0 || length > framePoolCeiling { + return nil + } + if buf := routedFrames.take(length); buf != nil { + return buf + } + return make([]byte, length) +} + +// take answers a buffer of exactly length bytes with at least that +// capacity, the smallest retained buffer that fits, or nil when the +// pool holds none. +func (p *framePool) take(length int64) []byte { + if length <= 0 || length > framePoolCeiling { + return nil + } + p.mu.Lock() + defer p.mu.Unlock() + best := -1 + for i, buf := range p.free { + if int64(cap(buf)) >= length && (best < 0 || cap(buf) < cap(p.free[best])) { + best = i + } + } + if best < 0 { + return nil + } + buf := p.free[best] + p.free = append(p.free[:best], p.free[best+1:]...) + return buf[:length] +} + +// retire returns one buffer to the pool under the two bounds: a +// buffer whose capacity exceeds the ceiling, and a return that would +// push the pool past its count bound, are dropped to the collector. +func (p *framePool) retire(data []byte) { + if !framePoolEnabled || cap(data) == 0 || int64(cap(data)) > framePoolCeiling { + return + } + p.mu.Lock() + if len(p.free) < framePoolBound { + p.free = append(p.free, data[:cap(data)]) + } + p.mu.Unlock() +} diff --git a/spmd/framepool_test.go b/spmd/framepool_test.go new file mode 100644 index 0000000..b93f8d5 --- /dev/null +++ b/spmd/framepool_test.go @@ -0,0 +1,212 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "encoding/binary" + "runtime" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// The pool tests cover the two bounds, the exact-length answer, the +// switch, the bound's flatness in the process's heap, and the one +// ownership chain itself: routed frames carrying distinct bits through +// a real hub while the pool recycles every buffer between pump and +// drain. + +// resetPool empties the shared pool so no test inherits another's +// buffers, pins the switch on, and restores both at the end. +func resetPool(t *testing.T) { + t.Helper() + was := framePoolEnabled + t.Cleanup(func() { + framePoolEnabled = was + emptyPool() + }) + emptyPool() + framePoolEnabled = true +} + +func emptyPool() { + routedFrames.mu.Lock() + routedFrames.free = nil + routedFrames.mu.Unlock() +} + +// poolDepth answers how many buffers the pool holds. +func poolDepth() int { + routedFrames.mu.Lock() + defer routedFrames.mu.Unlock() + return len(routedFrames.free) +} + +func TestFramePoolTakeAnswersExactLengths(t *testing.T) { + resetPool(t) + if buf := routedFrames.take(1000); buf != nil { + t.Fatal("an empty pool answered a buffer") + } + routedFrames.retire(make([]byte, 1000)) + buf := routedFrames.take(600) + if buf == nil { + t.Fatal("a retained buffer was not answered") + } + if len(buf) != 600 || cap(buf) < 600 { + t.Fatalf("take answered len %d cap %d for a 600 byte payload", len(buf), cap(buf)) + } + if buf := routedFrames.take(600); buf != nil { + t.Fatal("a taken buffer was answered twice") + } +} + +func TestFramePoolTakesTheSmallestThatFits(t *testing.T) { + resetPool(t) + routedFrames.retire(make([]byte, 5000)) + routedFrames.retire(make([]byte, 900)) + buf := routedFrames.take(800) + if buf == nil { + t.Fatal("a retained buffer was not answered") + } + if cap(buf) != 900 { + t.Fatalf("take answered cap %d when a 900 byte buffer was retained", cap(buf)) + } + buf = routedFrames.take(800) + if buf == nil || cap(buf) != 5000 { + t.Fatalf("the second take answered cap %d, want the 5000 byte buffer", cap(buf)) + } +} + +func TestFramePoolRefusesTheOversized(t *testing.T) { + resetPool(t) + if buf := routedFrames.take(framePoolCeiling + 1); buf != nil { + t.Fatal("a length beyond the ceiling was answered") + } + routedFrames.retire(make([]byte, framePoolCeiling+1)) + if got := poolDepth(); got != 0 { + t.Fatalf("a buffer beyond the ceiling was retained, pool holds %d", got) + } + routedFrames.retire(make([]byte, framePoolCeiling)) + if got := poolDepth(); got != 1 { + t.Fatalf("a buffer at the ceiling was refused, pool holds %d", got) + } + if buf := routedFrames.take(framePoolCeiling); buf == nil { + t.Fatal("a length at the ceiling was not answered") + } +} + +func TestFramePoolHoldsTheBound(t *testing.T) { + resetPool(t) + for range framePoolBound + 5 { + routedFrames.retire(make([]byte, 1000)) + } + if got := poolDepth(); got != framePoolBound { + t.Fatalf("the pool holds %d buffers beyond its bound of %d", got, framePoolBound) + } + if buf := routedFrames.take(1000); buf == nil { + t.Fatal("a bound-full pool answered nothing") + } +} + +func TestFramePoolAnswersNilWhenDisabled(t *testing.T) { + resetPool(t) + framePoolEnabled = false + routedFrames.retire(make([]byte, 1000)) + if got := poolDepth(); got != 0 { + t.Fatalf("a disabled pool retained %d buffers", got) + } + if buf := routedFrames.take(1000); buf != nil { + t.Fatal("a disabled pool answered a buffer") + } +} + +// TestFramePoolStaysFlatOnRepeatedRounds is the bound's proof in the +// heap: the same work run again and again in one process cannot raise +// the heap in use once the pool has warmed, and a rising trend is a +// defect. +func TestFramePoolStaysFlatOnRepeatedRounds(t *testing.T) { + resetPool(t) + sizes := []int64{64 << 10, 256 << 10, framePoolCeiling} + round := func() { + for _, s := range sizes { + buf := routedFrames.take(s) + if buf == nil { + buf = make([]byte, s) + } + for i := range buf { + buf[i] = byte(i) + } + routedFrames.retire(buf) + } + } + for range 50 { + round() + } + runtime.GC() + var before runtime.MemStats + runtime.ReadMemStats(&before) + for range 200 { + round() + } + runtime.GC() + var after runtime.MemStats + runtime.ReadMemStats(&after) + if after.HeapInuse > before.HeapInuse+4<<20 { + t.Fatalf("the heap in use rose from %d to %d bytes across 200 rounds", before.HeapInuse, after.HeapInuse) + } +} + +// routedProbeRounds is how many distinct payloads the ownership test +// pushes through the hub, every one recycled through the pool. +const routedProbeRounds = 300 + +// TestFramePoolCarriesTheBitsOverTCP is the ownership chain exercised: +// a non-hub rank's frames route through the hub's pump, outbox and +// drain, the drain returns every buffer the moment its write lands, +// and the pump reads the next frame into what comes back, so a premature +// return or a shared buffer would scramble the bits the receiver +// checks, most of all under the race detector. +func TestFramePoolCarriesTheBitsOverTCP(t *testing.T) { + resetPool(t) + runTCPWorld(t, 3, Options{Timeout: 30 * time.Second}, func(w *World) error { + if w.Rank() == 2 { + for i := range routedProbeRounds { + buf := make([]byte, 1<<18) + binary.LittleEndian.PutUint64(buf, uint64(i)) + for j := 8; j < len(buf); j += 7 { + buf[j] = byte(i + j) + } + if err := w.sendTo(1, tagHalo, buf); err != nil { + return err + } + } + return nil + } + if w.Rank() != 1 { + return nil + } + for i := range routedProbeRounds { + data, err := w.recvFrom(2, tagHalo) + if err != nil { + return err + } + if len(data) != 1<<18 { + return base.Errf("spmd: probe %d arrived %d bytes long", i, len(data)) + } + if got := binary.LittleEndian.Uint64(data); got != uint64(i) { + return base.Errf("spmd: probe %d arrived with the serial of %d", i, got) + } + for j := 8; j < len(data); j += 7 { + if data[j] != byte(i+j) { + return base.Errf("spmd: probe %d differs at byte %d", i, j) + } + } + } + return nil + }) + if poolDepth() == 0 { + t.Fatal("routed traffic left the pool empty, so no buffer ever rode the chain") + } +} diff --git a/spmd/halo.go b/spmd/halo.go new file mode 100644 index 0000000..6cdce11 --- /dev/null +++ b/spmd/halo.go @@ -0,0 +1,209 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// ExchangeHalos serves the domain-decomposed simulations: a rank whose +// piece is cut along the first axis hands its edge rows to its two +// neighbours and receives theirs. upper carries the rows that precede +// the piece in the global array (the lower edge of rank-1), lower the +// rows that follow it (the upper edge of rank+1); a rank at the +// world's edge receives nil for the side where no neighbour lives, +// and a zero halo width answers the same nil on both sides: nil is +// the answer every zero-length slab takes. The exchange moves bits +// and moves nothing else, so there is +// no order for it to get wrong: the runs in both directions are fixed +// by the rank indices, never by the arrival order. It is +// ExchangeHalosOnGrid on the one-dimensional grid of the whole world, +// cut along axis 0. +func (w *World) ExchangeHalos(local *core.Array, halos int) (*core.Array, *core.Array, error) { + return w.ExchangeHalosOnGrid(local, halos, 0, []int{w.size}) +} + +// ExchangeHalosOnGrid is the halo exchange of a domain-decomposed +// simulation on a process grid. grid lays the world's ranks out +// row-major, grid[a] naming how many positions the grid runs along +// axis a, and its product must be the world's size; rank r sits at +// coordinate (r/prod(grid[axis+1:]))%grid[axis] along axis, so its +// two neighbours there sit one grid step away, at rank-dist and +// rank+dist with dist = prod(grid[axis+1:]), and a rank at the grid's +// edge has no neighbour on that side and receives nil for it. Each +// rank hands the slab of width halos at its piece's edge along axis +// to the neighbour it borders there and receives theirs: upper +// carries the slab that precedes the piece along the axis, lower the +// slab that follows it, both keeping every other dimension of the +// piece's shape. A zero halo width answers nil on both sides, an +// empty piece joins with empty wires and answers no halos, and a +// neighbour with no piece to offer answers nil too, so the protocol +// stays symmetric. The exchange moves bits and moves nothing else, so +// there is no order +// for it to get wrong: the runs in both directions are fixed by the +// rank indices, never by the arrival order. +func (w *World) ExchangeHalosOnGrid(local *core.Array, halos, axis int, grid []int) (*core.Array, *core.Array, error) { + if err := w.status(); err != nil { + return nil, nil, err + } + if halos < 0 { + return nil, nil, w.fail(base.Errf("spmd: a negative halo width %d", halos)) + } + if local.NDim() == 0 { + return nil, nil, w.fail(base.Errf("spmd: the halo exchange needs a dimension to cut the edges from")) + } + if axis < 0 || axis >= local.NDim() { + return nil, nil, w.fail(base.Errf("spmd: halo axis %d is outside the array's %d dimensions", axis, local.NDim())) + } + if axis >= len(grid) { + return nil, nil, w.fail(base.Errf("spmd: halo axis %d is outside the grid's %d axes", axis, len(grid))) + } + extent := 1 + for _, d := range grid { + if d < 0 { + return nil, nil, w.fail(base.Errf("spmd: a grid cannot name a negative extent %d", d)) + } + if extent > 0 && d > math.MaxInt/extent { + return nil, nil, w.fail(base.Errf("spmd: a grid of %s lays out more ranks than the world can hold", base.ShapeText(grid))) + } + extent *= d + } + if extent != w.size { + return nil, nil, w.fail(base.Errf("spmd: a grid of %s lays out %d ranks against the world's %d", base.ShapeText(grid), extent, w.size)) + } + // The rank's place in the row-major grid: the coordinate along the + // cut axis, and the rank distance of one grid step along it. + dist := 1 + for _, d := range grid[axis+1:] { + dist *= d + } + coord := (w.rank / dist) % grid[axis] + hasLower, hasUpper := coord > 0, coord+1 < grid[axis] + shape := local.Shape() + if local.Len() == 0 { + // An empty piece owns no slab, so it carries no edges: it + // still joins the exchange with empty wires, so the + // neighbours' protocol stays symmetric, and answers no halos. + emptyShape := make([]int, len(shape)) + copy(emptyShape, shape) + emptyShape[axis] = 0 + empty, err := encodeHead(nil, local.Dtype(), emptyShape) + if err != nil { + return nil, nil, w.fail(err) + } + if hasUpper { + if err := w.sendTo(w.rank+dist, tagHalo, empty); err != nil { + return nil, nil, err + } + } + if hasLower { + if err := w.sendTo(w.rank-dist, tagHalo, empty); err != nil { + return nil, nil, err + } + } + if hasUpper { + if _, err := w.recvFrom(w.rank+dist, tagHalo); err != nil { + return nil, nil, err + } + } + if hasLower { + if _, err := w.recvFrom(w.rank-dist, tagHalo); err != nil { + return nil, nil, err + } + } + return nil, nil, nil + } + span := shape[axis] + if halos > span { + return nil, nil, w.fail(base.Errf("spmd: a halo width of %d rows exceeds the piece's %d", halos, span)) + } + lead, rest := 1, 1 + for _, d := range shape[:axis] { + lead *= d + } + for _, d := range shape[axis+1:] { + rest *= d + } + edgeShape := make([]int, len(shape)) + copy(edgeShape, shape) + edgeShape[axis] = halos + // The edge slab is one run of halos*rest elements per leading + // position, contiguous in the row-major layout; the wire carries + // the runs joined under the slab's own shape, whose extent along + // the axis is the halo width and whose other extents are the + // piece's own. + edge := func(high bool) ([]byte, error) { + wire, err := encodeHead(nil, local.Dtype(), edgeShape) + if err != nil { + return nil, err + } + count := halos * rest + for p := range lead { + first := p*span*rest + (span-halos)*rest + if !high { + first = p * span * rest + } + part, err := encodePart(nil, local, []int{count}, first, count) + if err != nil { + return nil, err + } + wire = append(wire, part[partHeadLen:]...) + } + return wire, nil + } + // Two phases: the lower halos travel to the upper neighbour + // first, the upper halos to the lower one second. The sends ride + // the links' buffers and the networked links drain through the + // hub, so no rank waits on a rank that is waiting on it. + if hasUpper { + highEdge, err := edge(true) + if err != nil { + return nil, nil, w.fail(err) + } + if err := w.sendTo(w.rank+dist, tagHalo, highEdge); err != nil { + return nil, nil, err + } + } + var lower *core.Array + if hasUpper { + data, err := w.recvFrom(w.rank+dist, tagHalo) + if err != nil { + return nil, nil, err + } + lower, err = decodeWire(data) + if err != nil { + return nil, nil, w.fail(err) + } + if lower.Len() == 0 { + lower = nil // the neighbour owns no slab + } + } + if hasLower { + lowEdge, err := edge(false) + if err != nil { + return nil, nil, w.fail(err) + } + if err := w.sendTo(w.rank-dist, tagHalo, lowEdge); err != nil { + return nil, nil, err + } + } + var upper *core.Array + if hasLower { + data, err := w.recvFrom(w.rank-dist, tagHalo) + if err != nil { + return nil, nil, err + } + upper, err = decodeWire(data) + if err != nil { + return nil, nil, w.fail(err) + } + if upper.Len() == 0 { + upper = nil // the neighbour owns no slab + } + } + return upper, lower, nil +} diff --git a/spmd/halo_grid_test.go b/spmd/halo_grid_test.go new file mode 100644 index 0000000..b5ddb7e --- /dev/null +++ b/spmd/halo_grid_test.go @@ -0,0 +1,372 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "slices" + "strings" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The grid halo tests pin ExchangeHalosOnGrid against the whole array: +// the rank's row-major place in the grid, its neighbours along the cut +// axis and the slabs that travel between them, with the serial stencil +// as the exacter judge of the geometry. + +// gridFixture builds the [rows, cols] fixture the grid tests deal out. +func gridFixture(rows, cols int) *core.Array { + vals := make([]float64, rows*cols) + for i := range vals { + vals[i] = float64(i*37%(rows*cols)) * 0.5 + } + return mk(core.FromFloats(vals, rows, cols)) +} + +// tileSpan cuts tile idx's run of an extent dealt into parts tiles, +// the deal that leaves the remainders on the last tiles. +func tileSpan(extent, parts, idx int) (int, int) { + return idx * extent / parts, (idx + 1) * extent / parts +} + +// gridCoords maps a rank into a row-major grid, coordinate a being +// (rank/prod(grid[a+1:]))%grid[a]. +func gridCoords(rank int, grid []int) []int { + coords := make([]int, len(grid)) + rest := 1 + for a := len(grid) - 1; a >= 0; a-- { + coords[a] = (rank / rest) % grid[a] + rest *= grid[a] + } + return coords +} + +// gridRank is gridCoords' inverse: the rank a row-major grid puts on +// the named coordinates. +func gridRank(coords []int, grid []int) int { + rank := 0 + for a := range grid { + rank = rank*grid[a] + coords[a] + } + return rank +} + +// dealTile cuts rows [rlo, rhi) by cols [clo, chi) off a 2-D whole, +// an empty array for an empty range. +func dealTile(t *testing.T, whole *core.Array, rlo, rhi, clo, chi int) *core.Array { + t.Helper() + width := whole.Shape()[1] + wire, err := encodeHead(nil, whole.Dtype(), []int{rhi - rlo, chi - clo}) + if err != nil { + t.Fatal(err) + } + for r := rlo; r < rhi; r++ { + part, err := encodePart(nil, whole, []int{chi - clo}, r*width+clo, chi-clo) + if err != nil { + t.Fatal(err) + } + wire = append(wire, part[2+8:]...) + } + a, err := decodeWire(wire) + if err != nil { + t.Fatal(err) + } + return a +} + +// checkGridExchange runs one rank's grid exchange and judges its halos +// against the neighbouring tiles' own edge slabs, worked out from the +// row-major grid arithmetic beside the exchange: a halo keeps the +// piece's other dimensions, so the expected slab is the neighbour's +// tile narrowed to the halo width along the cut axis. Empty pieces +// answer nothing, and empty neighbours and grid edges hand nil back. +func checkGridExchange(t *testing.T, w *World, whole *core.Array, halos, axis int, grid []int) { + t.Helper() + two := grid + if len(two) == 1 { + two = []int{two[0], 1} + } + coords := gridCoords(w.Rank(), two) + rlo, rhi, clo, chi := rankTile(whole, two, w.Rank()) + local := dealTile(t, whole, rlo, rhi, clo, chi) + upper, lower, err := w.ExchangeHalosOnGrid(local, halos, axis, grid) + if err != nil { + t.Fatalf("rank %d: %v", w.Rank(), err) + } + if local.Len() == 0 { + if upper != nil || lower != nil { + t.Fatalf("rank %d: an empty piece answered halos", w.Rank()) + } + return + } + sides := []struct { + name string + got *core.Array + off int + }{ + {"upper", upper, -1}, + {"lower", lower, 1}, + } + for _, side := range sides { + got := side.got + if got != nil && got.Len() == 0 { + got = nil // a zero-width slab answers nothing + } + c := coords[axis] + side.off + if c < 0 || c >= two[axis] { + if got != nil { + t.Fatalf("rank %d: a %s halo arrived where no neighbour lives", w.Rank(), side.name) + } + continue + } + sideCoords := slices.Clone(coords) + sideCoords[axis] = c + nrlo, nrhi, nclo, nchi := rankTile(whole, two, gridRank(sideCoords, two)) + if nrhi == nrlo || nchi == nclo { + // The neighbour owns no elements, so it carried no slab. + if got != nil { + t.Fatalf("rank %d: a %s halo arrived from an empty neighbour", w.Rank(), side.name) + } + continue + } + // The neighbour's slab: its tile's edge along the cut axis, + // kept whole on every other dimension. + cutLo, cutHi := nclo, nchi + if axis == 0 { + cutLo, cutHi = nrlo, nrhi + } + sLo, sHi := cutLo, cutLo+halos + if side.off < 0 { + sLo, sHi = cutHi-halos, cutHi + } + var want *core.Array + if axis == 0 { + want = dealTile(t, whole, sLo, sHi, nclo, nchi) + } else { + want = dealTile(t, whole, nrlo, nrhi, sLo, sHi) + } + if want.Len() == 0 { + if got != nil { + t.Fatalf("rank %d: a %s halo arrived for a zero-width slab", w.Rank(), side.name) + } + continue + } + if got == nil { + t.Fatalf("rank %d: the %s halo is nil with a living neighbour", w.Rank(), side.name) + } + if !slices.Equal(got.Shape(), want.Shape()) || !sameBits(want, got) { + t.Fatalf("rank %d: the %s halo differs from the neighbouring tile's slab along axis %d", + w.Rank(), side.name, axis) + } + } +} + +// rankTile is dealTile's own cut for one rank of a 2-D grid. +func rankTile(whole *core.Array, grid []int, rank int) (rlo, rhi, clo, chi int) { + coords := gridCoords(rank, grid) + rlo, rhi = tileSpan(whole.Shape()[0], grid[0], coords[0]) + clo, chi = tileSpan(whole.Shape()[1], grid[1], coords[1]) + return rlo, rhi, clo, chi +} + +// TestExchangeHalosOnGridMatchesTheWhole walks the whole matrix the +// API claims: world sizes 1 to 8, halo widths 0, 1 and 3, every +// two-factor grid of each size beside the degenerate one-dimensional +// one, cut along every axis the grid names. The expected slabs come +// from the rank's grid coordinates worked out beside the exchange, so +// a wrong row-major mapping, a wrong neighbour distance or a wrong +// slab cannot pass: on the grid [2, 3] the axis 1 neighbours are +// rank+-1 and the axis 0 neighbours rank+-3, and the checks know it +// independently. +func TestExchangeHalosOnGridMatchesTheWhole(t *testing.T) { + // The extents divide out to tiles of at least three elements on + // every factor grid up to eight positions, so even the widest + // pinned halo fits the piece it travels from. + const rows, cols = 24, 24 + whole := gridFixture(rows, cols) + for _, size := range []int{1, 2, 3, 5, 8} { + grids := [][]int{{size}} + for a := 1; a <= size; a++ { + if size%a == 0 { + grids = append(grids, []int{a, size / a}) + } + } + for _, grid := range grids { + for _, halos := range []int{0, 1, 3} { + for axis := range len(grid) { + err := Launch(size, func(w *World) error { + checkGridExchange(t, w, whole, halos, axis, grid) + return nil + }) + if err != nil { + t.Fatalf("size %d grid %v halos %d axis %d: %v", size, grid, halos, axis, err) + } + } + } + } + } +} + +// TestExchangeHalosOnGridEmptyPieces: the pieces the tiling leaves +// empty join the exchange symmetrically, answering no halos, and the +// neighbours see nil from the empty side, in both grid orientations. +func TestExchangeHalosOnGridEmptyPieces(t *testing.T) { + whole := gridFixture(1, 2) + for _, grid := range [][]int{{2, 2}, {3, 2}, {2, 3}} { + for axis := range 2 { + err := Launch(grid[0]*grid[1], func(w *World) error { + checkGridExchange(t, w, whole, 1, axis, grid) + return nil + }) + if err != nil { + t.Fatalf("grid %v axis %d: %v", grid, axis, err) + } + } + } +} + +// TestExchangeHalosOnGridStencilMatchesSerial runs the five-point +// stencil over a field tiled 2 by 3, exchanging halos along both grid +// axes, and compares every tile element for element against the same +// step computed on the whole field. The comparison is the matrix +// claim made flesh: on the row-major grid [2, 3] the axis 1 +// neighbours are rank+-1, the axis 0 neighbours rank+-3, and the +// distributed stencil may only agree bit for bit when the slabs +// deliver exactly the neighbours the serial walk sees. +func TestExchangeHalosOnGridStencilMatchesSerial(t *testing.T) { + const rows, cols = 6, 12 + whole := gridFixture(rows, cols) + serial := make([]float64, rows*cols) + for r := range rows { + for c := range cols { + up, down := 0.0, 0.0 + if r > 0 { + up = whole.FloatAt((r-1)*cols + c) + } + if r < rows-1 { + down = whole.FloatAt((r+1)*cols + c) + } + left, right := 0.0, 0.0 + if c > 0 { + left = whole.FloatAt(r*cols + c - 1) + } + if c < cols-1 { + right = whole.FloatAt(r*cols + c + 1) + } + serial[r*cols+c] = (up + left + whole.FloatAt(r*cols+c) + right + down) / 5 + } + } + const halos = 1 + err := Launch(6, func(w *World) error { + rlo, rhi, clo, chi := rankTile(whole, []int{2, 3}, w.Rank()) + local := dealTile(t, whole, rlo, rhi, clo, chi) + upper, lower, err := w.ExchangeHalosOnGrid(local, halos, 0, []int{2, 3}) + if err != nil { + return err + } + leftHalo, rightHalo, err := w.ExchangeHalosOnGrid(local, halos, 1, []int{2, 3}) + if err != nil { + return err + } + tileRows, tileCols := rhi-rlo, chi-clo + for ri := range tileRows { + for cj := range tileCols { + gi, gj := rlo+ri, clo+cj + centre := local.FloatAt(ri*tileCols + cj) + up := 0.0 + if ri == 0 { + if gi > 0 { + up = upper.FloatAt(cj) // the halo's last row adjoins the tile + } + } else { + up = local.FloatAt((ri-1)*tileCols + cj) + } + down := 0.0 + if ri == tileRows-1 { + if gi < rows-1 { + down = lower.FloatAt(cj) + } + } else { + down = local.FloatAt((ri+1)*tileCols + cj) + } + left := 0.0 + if cj == 0 { + if gj > 0 { + left = leftHalo.FloatAt(ri) // the halo's last column adjoins the tile + } + } else { + left = local.FloatAt(ri*tileCols + cj - 1) + } + right := 0.0 + if cj == tileCols-1 { + if gj < cols-1 { + right = rightHalo.FloatAt(ri) + } + } else { + right = local.FloatAt(ri*tileCols + cj + 1) + } + got := (up + left + centre + right + down) / 5 + if got != serial[gi*cols+gj] { + t.Fatalf("rank %d global (%d, %d): %v against the serial %v", + w.Rank(), gi, gj, got, serial[gi*cols+gj]) + } + } + } + return w.Barrier() + }) + if err != nil { + t.Fatal(err) + } +} + +// TestExchangeHalosOnGridOverTCP runs the grid exchange over real +// connections, the axis 1 slabs included. +func TestExchangeHalosOnGridOverTCP(t *testing.T) { + whole := gridFixture(24, 18) + runTCPWorld(t, 6, Options{Timeout: 30 * time.Second}, func(w *World) error { + checkGridExchange(t, w, whole, 2, 1, []int{2, 3}) + checkGridExchange(t, w, whole, 2, 0, []int{2, 3}) + return nil + }) +} + +// TestExchangeHalosOnGridRefusesTheHostile: every invalid argument is +// a named error before any frame moves. Every rank makes the same +// invalid call, so an exchange that ever started would deadlock the +// world instead of answering, which is what pins the ordering too. +func TestExchangeHalosOnGridRefusesTheHostile(t *testing.T) { + local := mk(core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)) + cases := []struct { + name string + local *core.Array + halos, axis int + grid []int + want string + }{ + {"negative halos", local, -1, 0, []int{2}, "negative halo width"}, + {"axis past the array", local, 1, 2, []int{2, 2}, "outside the array"}, + {"negative axis", local, 1, -1, []int{2, 2}, "outside the array"}, + {"axis past the grid", local, 1, 1, []int{2}, "outside the grid"}, + {"grid short of the world", local, 1, 0, []int{1, 3}, "against the world's"}, + {"negative grid extent", local, 1, 0, []int{-2}, "negative extent"}, + {"halos wider than the axis", local, 3, 0, []int{2}, "exceeds the piece's"}, + } + for _, tc := range cases { + err := Launch(2, func(w *World) error { + _, _, err := w.ExchangeHalosOnGrid(tc.local, tc.halos, tc.axis, tc.grid) + if err == nil { + t.Fatalf("%s: the hostile call was accepted", tc.name) + } + if !strings.Contains(err.Error(), tc.want) { + t.Fatalf("%s: the error %q does not name the fault", tc.name, err) + } + return nil + }) + if err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + } +} diff --git a/spmd/halo_test.go b/spmd/halo_test.go new file mode 100644 index 0000000..fc66f64 --- /dev/null +++ b/spmd/halo_test.go @@ -0,0 +1,180 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// globalRows builds the [rows, width] fixture the halo tests deal out. +func globalRows(rows, width int) *core.Array { + vals := make([]float64, rows*width) + for i := range vals { + vals[i] = float64(i*31%(rows*width)) * 0.25 + } + a, err := core.FromFloats(vals, rows, width) + if err != nil { + panic(err) + } + return a +} + +// dealRows2D cuts rows [lo, hi) off the global array, an empty array +// for an empty range. +func dealRows2D(tb testing.TB, whole *core.Array, lo, hi int) *core.Array { + tb.Helper() + rest := 1 + for _, d := range whole.Shape()[1:] { + rest *= d + } + wire, err := encodePart(nil, whole, append([]int{hi - lo}, whole.Shape()[1:]...), lo*rest, (hi-lo)*rest) + if err != nil { + tb.Fatal(err) + } + a, err := decodeWire(wire) + if err != nil { + tb.Fatal(err) + } + return a +} + +// checkHalos asserts one rank's halos against the whole's neighbouring +// rows, treating empty neighbours and empty pieces as no data. +func checkHalos(t *testing.T, w *World, whole *core.Array, rows, halos int, upper, lower *core.Array) { + t.Helper() + span := mustPartition(t, rows, w.Size(), w.Rank()) + neighbourRows := func(rank int) int { + if rank < 0 || rank >= w.Size() { + return 0 + } + return mustPartition(t, rows, w.Size(), rank).Len() + } + if span.Len() == 0 { + if upper != nil || lower != nil { + t.Fatalf("rank %d: an empty piece answered halos", w.Rank()) + } + return + } + if upper != nil && upper.Len() == 0 { + upper = nil + } + if lower != nil && lower.Len() == 0 { + lower = nil + } + if w.Rank() == 0 || neighbourRows(w.Rank()-1) == 0 { + if upper != nil { + t.Fatalf("rank %d: an upper halo arrived where no rows live", w.Rank()) + } + } else { + if upper == nil { + t.Fatalf("rank %d: the upper halo is nil with a living neighbour", w.Rank()) + } + if want := dealRows2D(t, whole, span.Lo-halos, span.Lo); !sameBits(want, upper) { + t.Fatalf("rank %d: the upper halo differs from rows [%d, %d)", w.Rank(), span.Lo-halos, span.Lo) + } + } + if w.Rank()+1 == w.Size() || neighbourRows(w.Rank()+1) == 0 { + if lower != nil { + t.Fatalf("rank %d: a lower halo arrived where no rows live", w.Rank()) + } + } else { + if lower == nil { + t.Fatalf("rank %d: the lower halo is nil with a living neighbour", w.Rank()) + } + if want := dealRows2D(t, whole, span.Hi, span.Hi+halos); !sameBits(want, lower) { + t.Fatalf("rank %d: the lower halo differs from rows [%d, %d)", w.Rank(), span.Hi, span.Hi+halos) + } + } +} + +// TestExchangeHalosMatchTheWhole: every rank's upper and lower halos +// are exactly the whole's neighbouring rows, and edges without a +// living neighbour answer nil. +func TestExchangeHalosMatchTheWhole(t *testing.T) { + for _, size := range []int{1, 2, 3, 5, 8} { + for _, halos := range []int{0, 1, 3} { + const rows, width = 40, 5 + whole := globalRows(rows, width) + err := Launch(size, func(w *World) error { + span := mustPartition(t, rows, w.Size(), w.Rank()) + local := dealRows2D(t, whole, span.Lo, span.Hi) + upper, lower, err := w.ExchangeHalos(local, halos) + if err != nil { + return err + } + checkHalos(t, w, whole, rows, halos, upper, lower) + return nil + }) + if err != nil { + t.Fatalf("size %d halos %d: %v", size, halos, err) + } + } + } +} + +// TestExchangeHalosStencilMatchesSerial runs a three-point smoothing +// step over a sharded series and compares it, element for element, +// with the same step computed on the whole array: the halos must make +// the distributed stencil see exactly the neighbours the serial one +// sees. +func TestExchangeHalosStencilMatchesSerial(t *testing.T) { + const n = 1000 + whole := fixtureArray(n) + serial := make([]float64, n) + for i := 1; i < n-1; i++ { + serial[i] = (whole.FloatAt(i-1) + whole.FloatAt(i) + whole.FloatAt(i+1)) / 3 + } + err := Launch(4, func(w *World) error { + span := mustPartition(t, n, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + upper, lower, err := w.ExchangeHalos(local, 1) + if err != nil { + return err + } + for i := range span.Len() { + g := span.Lo + i + if g == 0 || g == n-1 { + continue + } + var left, right float64 + if i == 0 { + left = upper.FloatAt(0) + } else { + left = local.FloatAt(i - 1) + } + if i == span.Len()-1 { + right = lower.FloatAt(0) + } else { + right = local.FloatAt(i + 1) + } + if got := (left + local.FloatAt(i) + right) / 3; got != serial[g] { + t.Fatalf("rank %d global %d: %v against the serial %v", w.Rank(), g, got, serial[g]) + } + } + return w.Barrier() + }) + if err != nil { + t.Fatal(err) + } +} + +// TestExchangeHalosOverTCP runs the neighbour exchange over real +// connections. +func TestExchangeHalosOverTCP(t *testing.T) { + const rows, width = 200, 3 + whole := globalRows(rows, width) + runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error { + span := mustPartition(t, rows, w.Size(), w.Rank()) + local := dealRows2D(t, whole, span.Lo, span.Hi) + upper, lower, err := w.ExchangeHalos(local, 2) + if err != nil { + return err + } + checkHalos(t, w, whole, rows, 2, upper, lower) + return nil + }) +} diff --git a/spmd/partition.go b/spmd/partition.go new file mode 100644 index 0000000..80ab950 --- /dev/null +++ b/spmd/partition.go @@ -0,0 +1,79 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Span names one rank's contiguous piece of a global axis of data: the +// global length, the piece's first element and one past its last. +type Span struct { + // Global is the axis length every rank's piece is a cut of. + Global int + // Lo is the piece's first global element index; Hi is one past its + // last. The piece is Lo..Hi, possibly empty. + Lo, Hi int +} + +// Len returns the number of elements the piece carries along the axis. +func (s Span) Len() int { return s.Hi - s.Lo } + +// Partition cuts a global axis of globalN elements into the world's +// contiguous pieces: rank rank's piece is [Lo, Hi). Every piece starts +// and ends on a boundary of the canonical fold partition, which is what +// lets the shards' reductions compose into the single-array fold's +// exact bits, so a program that cuts its data any other way gives up +// that contract. With more ranks than blocks, the pieces beyond the blocks are empty, +// wherever the boundary falls, so rank 0 may hold nothing at all. +// +// A negative globalN, a size below 1 or a rank outside [0, size) is +// refused with an error naming the input. Valid arguments never fail. +func Partition(globalN, size, rank int) (Span, error) { + if globalN < 0 { + return Span{}, base.Errf("spmd: Partition of a negative global length %d", globalN) + } + if size < 1 || rank < 0 || rank >= size { + return Span{}, base.Errf("spmd: Partition of rank %d in a world of %d", rank, size) + } + parts := core.FoldParts(globalN) + return Span{ + Global: globalN, + Lo: core.FoldBoundary(globalN, rank*parts/size), + Hi: core.FoldBoundary(globalN, (rank+1)*parts/size), + }, nil +} + +// checkSpan verifies that a caller's span is the canonical partition's +// own cut for this rank and that the local slab leads with exactly the +// span's run of the axis. Any other cut is refused by name: the +// bit-identity contract lives on the canonical boundaries. +func (w *World) checkSpan(span Span, local *core.Array) error { + want, err := Partition(span.Global, w.size, w.rank) + if err != nil { + return err + } + if span != want { + return base.Errf("spmd: rank %d holds [%d, %d) of %d, but the canonical partition puts this rank on [%d, %d); cut the data with Partition", + w.rank, span.Lo, span.Hi, span.Global, want.Lo, want.Hi) + } + if local.NDim() == 0 || local.Shape()[0] != span.Len() { + lead := 0 + if local.NDim() > 0 { + lead = local.Shape()[0] + } + return base.Errf("spmd: rank %d's slab leads with %d elements for a span of %d", w.rank, lead, span.Len()) + } + return nil +} + +// checkOneDimSpan is checkSpan for the vector reductions: the shards +// carry one-dimensional arrays, whose length is the span's run. +func (w *World) checkOneDimSpan(local *core.Array, span Span) error { + if local.NDim() != 1 { + return base.Errf("spmd: the sharded product, norm and dot carry 1-D arrays; got %d dimensions", local.NDim()) + } + return w.checkSpan(span, local) +} diff --git a/spmd/partition_test.go b/spmd/partition_test.go new file mode 100644 index 0000000..a888654 --- /dev/null +++ b/spmd/partition_test.go @@ -0,0 +1,44 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "strings" + "testing" +) + +// mustPartition is Partition for the tests' fixed valid arguments: the +// piece, or a failure the tests can never meet. Benchmarks pass their +// own tb. +func mustPartition(tb testing.TB, globalN, size, rank int) Span { + tb.Helper() + span, err := Partition(globalN, size, rank) + if err != nil { + tb.Fatal(err) + } + return span +} + +// TestPartitionRefusesInvalidArguments: the inputs Partition refuses +// come back as errors naming the input, never as panics. +func TestPartitionRefusesInvalidArguments(t *testing.T) { + for _, tc := range []struct { + name string + globalN, size, rank int + want string + }{ + {"negative global length", -1, 4, 0, "Partition of a negative global length -1"}, + {"empty world", 100, 0, 0, "Partition of rank 0 in a world of 0"}, + {"negative rank", 100, 4, -1, "Partition of rank -1 in a world of 4"}, + {"rank past the world", 100, 4, 4, "Partition of rank 4 in a world of 4"}, + } { + span, err := Partition(tc.globalN, tc.size, tc.rank) + if err == nil { + t.Fatalf("%s: Partition answered %v", tc.name, span) + } + if !strings.Contains(err.Error(), tc.want) { + t.Fatalf("%s: the error does not name the input: %v", tc.name, err) + } + } +} diff --git a/spmd/proc_test.go b/spmd/proc_test.go new file mode 100644 index 0000000..76e0697 --- /dev/null +++ b/spmd/proc_test.go @@ -0,0 +1,201 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "fmt" + "math" + "net" + "os" + "os/exec" + "strconv" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The loopback tests run the package as a cluster: N real processes on +// one machine, rank 0 inside the test process and the rest re-executed +// out of this same test binary. It is the same code a multi-machine +// run executes, which is what makes the multi-machine contract +// testable without one. + +func TestMain(m *testing.M) { + if os.Getenv("TENSOR_SPMD_WORKER") != "" { + os.Exit(spmdWorker(os.Getenv("TENSOR_SPMD_ADDR"))) + } + os.Exit(m.Run()) +} + +// spmdWorker is the program a re-executed worker process runs: join +// the world, reduce the shards, compare the bits, leave cleanly. The +// fixture is deterministic, so no data crosses with the address. +func spmdWorker(addr string) int { + rank, _ := strconv.Atoi(os.Getenv("TENSOR_SPMD_RANK")) + size, _ := strconv.Atoi(os.Getenv("TENSOR_SPMD_SIZE")) + w, err := Join(addr, Options{Timeout: 2 * time.Minute}) + if err != nil { + fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err) + return 1 + } + const gn = 131073 + whole := fixtureArray(gn) + want := core.Sum(whole) + span, err := Partition(gn, size, w.Rank()) + if err != nil { + fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err) + return 1 + } + local, err := dealRows(whole, span) + if err != nil { + fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err) + return 1 + } + got, err := w.AllReduceShards(local, span, Sum) + if err != nil { + fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err) + return 1 + } + if scalarBits(got) != scalarBits(want) { + fmt.Fprintf(os.Stderr, "worker %d: sharded %s against single-array %s\n", + rank, scalarBits(got), scalarBits(want)) + return 1 + } + if err := w.Barrier(); err != nil { + fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err) + return 1 + } + if err := w.Close(); err != nil { + fmt.Fprintf(os.Stderr, "worker %d: %v\n", rank, err) + return 1 + } + return 0 +} + +// dealRows cuts a shard's rows off a 1-D whole, bits exactly. +func dealRows(whole *core.Array, span Span) (*core.Array, error) { + wire, err := encodePart(nil, whole, []int{span.Len()}, span.Lo, span.Len()) + if err != nil { + return nil, err + } + return decodeWire(wire) +} + +func TestMultiprocessLoopback(t *testing.T) { + const size = 4 + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + rankErr := make(chan error, 1) + go func() { + w, err := listen(ln, size, Options{Timeout: 2 * time.Minute}) + if err != nil { + rankErr <- err + return + } + defer w.Close() + const gn = 131073 + whole := fixtureArray(gn) + want := core.Sum(whole) + span, err := Partition(gn, size, w.Rank()) + if err != nil { + rankErr <- err + return + } + local, err := dealRows(whole, span) + if err != nil { + rankErr <- err + return + } + got, err := w.AllReduceShards(local, span, Sum) + if err != nil { + rankErr <- err + return + } + if scalarBits(got) != scalarBits(want) { + rankErr <- fmt.Errorf("rank 0: sharded %s against single-array %s", + scalarBits(got), scalarBits(want)) + return + } + if err := w.Barrier(); err != nil { + rankErr <- err + return + } + rankErr <- nil + }() + workers := make([]*exec.Cmd, size-1) + for r := 1; r < size; r++ { + cmd := exec.Command(os.Args[0], "-test.run=^$") + cmd.Env = append(os.Environ(), + "TENSOR_SPMD_WORKER=1", + "TENSOR_SPMD_ADDR="+ln.Addr().String(), + "TENSOR_SPMD_RANK="+strconv.Itoa(r), + "TENSOR_SPMD_SIZE="+strconv.Itoa(size)) + workers[r-1] = cmd + } + for r, cmd := range workers { + if err := cmd.Start(); err != nil { + t.Fatalf("worker %d never started: %v", r+1, err) + } + } + for r, cmd := range workers { + if err := cmd.Wait(); err != nil { + t.Fatalf("worker %d failed: %v", r+1, err) + } + } + if err := <-rankErr; err != nil { + t.Fatalf("rank 0: %v", err) + } +} + +// TestRepeatedRunsAnswerIdenticalBits is the arrival-order claim, +// exercised: one world runs the sharded reduction and the movement +// collectives interleaved many times, and six worlds run the whole +// battery again, so goroutine scheduling arrives at every order it +// can find and the bits may not move once. +func TestRepeatedRunsAnswerIdenticalBits(t *testing.T) { + const gn = 200001 + whole := fixtureArray(gn) + want := core.Sum(whole) + wantMax, err := core.Max(whole) + if err != nil { + t.Fatal(err) + } + for attempt := range 6 { + err := Launch(5, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local, err := dealRows(whole, span) + if err != nil { + return err + } + for i := range 8 { + got, err := w.AllReduceShards(local, span, Sum) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("attempt %d run %d: sharded %s against single-array %s", + attempt, i, scalarBits(got), scalarBits(want)) + } + gotMax, err := w.AllReduceShards(local, span, Max) + if err != nil { + return err + } + if math.Float64bits(gotMax.Float()) != math.Float64bits(wantMax.Float()) { + t.Fatalf("attempt %d run %d: sharded max moved", attempt, i) + } + if _, err := w.AllGather(local); err != nil { + return err + } + } + return nil + }) + if err != nil { + t.Fatalf("attempt %d: %v", attempt, err) + } + } +} diff --git a/spmd/reduce.go b/spmd/reduce.go new file mode 100644 index 0000000..bc70b4c --- /dev/null +++ b/spmd/reduce.go @@ -0,0 +1,1291 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "encoding/binary" + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Op names a reduction's arithmetic. +type Op uint8 + +const ( + // Sum adds: shard partials through the balanced tree, rank arrays + // in rank order. + Sum Op = iota + // Min takes the smallest; NaN never wins. + Min + // Max takes the largest; NaN never wins. + Max + // Any reports whether any element is true; Bool arrays only. + Any + // All reports whether every element is true; Bool arrays only. + All + // Prod multiplies: shard partials through the balanced product + // tree, rank arrays in rank order. + Prod +) + +func (o Op) String() string { + switch o { + case Sum: + return "sum" + case Min: + return "min" + case Max: + return "max" + case Any: + return "any" + case All: + return "all" + case Prod: + return "prod" + } + return "?" +} + +// The payload kinds a shard's contribution can travel as. The kind is +// part of the wire, so the root combines exactly what the ranks +// folded, never a guess from the dtype. +const ( + kindNone byte = 0 // a rank whose piece is empty + kindSumF64 byte = 1 // one float64 fold per block + kindSumC128 byte = 2 // one complex fold per block + kindSumI64 byte = 3 // one exact int64 sum per block + kindExtF64 byte = 4 // one float64 extremum and candidate flag per block + kindExtI64 byte = 5 // one exact int64 extremum per block + kindFlag byte = 6 // one 0/1 flag for the whole piece + kindProdF64 byte = 7 // one float64 product per block + kindProdF32 byte = 8 // one native float32 product per block + kindProdHalf byte = 9 // one exact half product per block, in float64 + kindProdI64 byte = 10 // one exact int64 product per block + kindArg byte = 11 // one extremum candidate: value, index and flags +) + +// intPayload is the integer element type a shard folds exactly. +type intPayload interface { + int64 | int8 | uint8 | int16 | uint16 | int32 | uint32 +} + +// AllReduceShards answers the reduction of one global array whose +// canonical pieces the ranks hold, every rank receiving the answer. +// Each rank folds its own blocks with the single block fold, the block +// values gather at rank 0, and core's own tree and extremum rules +// combine them: the answer is the single-array reduction's exact bits, +// whatever the world's size, whatever the order the frames arrive in. +// Any and All answer the integer-class scalar, 1 or 0. +func (w *World) AllReduceShards(local *core.Array, span Span, op Op) (core.Scalar, error) { + ans, err := w.reduceShards(local, span, op, 0) + if err != nil { + return core.Scalar{}, err + } + return w.broadcastScalar(ans, 0) +} + +// ReduceShards is AllReduceShards with the answer on the root alone; +// every other rank receives nil. +func (w *World) ReduceShards(local *core.Array, span Span, op Op, root int) (*core.Scalar, error) { + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + ans, err := w.reduceShards(local, span, op, root) + if err != nil { + return nil, err + } + if w.rank != root { + return nil, nil + } + return &ans, nil +} + +// blockValues is one rank's contribution to a sharded reduction: the +// index of its first block, how many blocks it covers, the payload +// kind, and the values. +type blockValues struct { + first int + kind byte + f []float64 + f32 []float32 + c []complex128 + i []int64 + oks []bool +} + +// blocks reports how many partition blocks the piece covers. +func (bv blockValues) blocks() int { + switch bv.kind { + case kindSumF64, kindSumI64: + return len(bv.f) + len(bv.i) + case kindSumC128: + return len(bv.c) + case kindExtF64: + return len(bv.oks) + case kindExtI64: + return len(bv.i) + case kindProdF64: + return len(bv.f) + case kindProdF32: + return len(bv.f32) + case kindProdHalf, kindProdI64: + return len(bv.f) + len(bv.i) + } + return 0 +} + +// spanBlocks names the partition blocks the span covers: [first, last). +// A span that the partition gave out always sits on block boundaries. +func spanBlocks(span Span) (first, last, parts int) { + parts = core.FoldParts(span.Global) + first, last = -1, -1 + for c := 0; c <= parts; c++ { + b := core.FoldBoundary(span.Global, c) + if b == span.Lo && first < 0 { + first = c + } + if b == span.Hi { + last = c + } + } + return first, last, parts +} + +// reduceShards runs the sharded reduction; the answer exists at the +// root. +func (w *World) reduceShards(local *core.Array, span Span, op Op, root int) (core.Scalar, error) { + if err := w.status(); err != nil { + return core.Scalar{}, err + } + if err := w.checkSpan(span, local); err != nil { + return core.Scalar{}, w.fail(err) + } + if span.Global == 0 { + // An empty global axis answers without an exchange: the + // single-array reduction answers the same constants, so the + // contract holds at the degenerate length too. The dtype + // validates first, exactly as a non-empty array would. + switch op { + case Prod: + if dt := local.Dtype(); dt == core.Bool || dt == core.Int8 || dt == core.Uint8 || dt == core.Int16 || dt == core.Uint16 || dt == core.Int32 || dt == core.Uint32 || dt == core.Complex { + return core.Scalar{}, w.fail(base.Errf("spmd: prod of dtype %s is not supported; convert with Astype", dt)) + } + case Min, Max: + if dt := local.Dtype(); dt == core.Bool || dt == core.Complex { + return core.Scalar{}, w.fail(base.Errf("spmd: %s of dtype %s has no ordering", op, dt)) + } + case Any, All: + if dt := local.Dtype(); dt != core.Bool { + return core.Scalar{}, w.fail(base.Errf("spmd: %s needs Bool elements, got %s", op, dt)) + } + } + switch op { + case Sum: + switch local.Dtype() { + case core.Complex: + return core.ComplexScalar(0), nil + case core.Int, core.Int8, core.Uint8, core.Int16, core.Uint16, + core.Int32, core.Uint32, core.Bool: + return core.IntScalar(0), nil + default: + return core.FloatScalar(0), nil + } + case Any: + return core.IntScalar(0), nil + case All: + return core.IntScalar(1), nil + case Prod: + // The empty product is the multiplicative identity, the + // answer the single-array product keeps. + if local.Dtype() == core.Int { + return core.IntScalar(1), nil + } + return core.FloatScalar(1), nil + default: + return core.Scalar{}, w.fail(base.Errf("spmd: %s of an empty array has no answer", op)) + } + } + bv, err := w.foldBlocks(local, span, op) + if err != nil { + return core.Scalar{}, w.fail(err) + } + if w.rank != root { + if err := w.sendTo(root, tagShardsValues, encodeBlockValues(bv)); err != nil { + return core.Scalar{}, err + } + return core.Scalar{}, nil + } + all := make([]blockValues, w.size) + all[w.rank] = bv + for r := range w.size { + if r == root { + continue + } + data, err := w.recvFrom(r, tagShardsValues) + if err != nil { + return core.Scalar{}, err + } + if all[r], err = decodeBlockValues(data); err != nil { + return core.Scalar{}, w.fail(err) + } + } + return w.combineShards(all, span, op) +} + +// foldBlocks computes the rank's own blocks' values with the same +// per-block arithmetic the single-array fold uses, dtype by dtype. +func (w *World) foldBlocks(local *core.Array, span Span, op Op) (blockValues, error) { + first, last, _ := spanBlocks(span) + if first < 0 || last < 0 || first > last { + return blockValues{}, base.Errf("spmd: rank %d's span [%d, %d) is not on the partition's block boundaries", + w.rank, span.Lo, span.Hi) + } + bv := blockValues{first: first} + if first == last { + return bv, nil + } + n := span.Global + count := last - first + blockF64 := func(c int) []float64 { + return sliceBlock(local.RawFloats()[:local.Len()], n, first, c) + } + switch op { + case Sum: + switch local.Dtype() { + case core.Float: + bv.kind = kindSumF64 + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldRange(blockF64(c)) + } + case core.Float32: + bv.kind = kindSumF64 + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldRangeF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c)) + } + case core.Float16: + bv.kind = kindSumF64 + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldRangeF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c)) + } + case core.Complex: + bv.kind = kindSumC128 + bv.c = make([]complex128, count) + for c := first; c < last; c++ { + bv.c[c-first] = core.FoldRangeC128(sliceBlock(local.RawComplexes()[:local.Len()], n, first, c)) + } + case core.Int: + bv.kind = kindSumI64 + bv.i = exactSums(local.RawInts()[:local.Len()], n, first, last) + case core.Int8: + bv.kind = kindSumI64 + bv.i = exactSums(local.RawInt8s()[:local.Len()], n, first, last) + case core.Uint8: + bv.kind = kindSumI64 + bv.i = exactSums(local.RawUint8s()[:local.Len()], n, first, last) + case core.Int16: + bv.kind = kindSumI64 + bv.i = exactSums(local.RawInt16s()[:local.Len()], n, first, last) + case core.Uint16: + bv.kind = kindSumI64 + bv.i = exactSums(local.RawUint16s()[:local.Len()], n, first, last) + case core.Int32: + bv.kind = kindSumI64 + bv.i = exactSums(local.RawInt32s()[:local.Len()], n, first, last) + case core.Uint32: + bv.kind = kindSumI64 + bv.i = exactSums(local.RawUint32s()[:local.Len()], n, first, last) + case core.Bool: + bv.kind = kindSumI64 + src := local.RawBools()[:local.Len()] + bv.i = make([]int64, count) + for c := first; c < last; c++ { + var s int64 + for _, v := range sliceBlock(src, n, first, c) { + if v { + s++ + } + } + bv.i[c-first] = s + } + default: + return blockValues{}, base.Errf("spmd: sum of dtype %s is not a reduction", local.Dtype()) + } + case Min, Max: + greater := op == Max + switch local.Dtype() { + case core.Float: + bv.kind = kindExtF64 + bv.f = make([]float64, count) + bv.oks = make([]bool, count) + for c := first; c < last; c++ { + bv.f[c-first], bv.oks[c-first] = core.ExtremeRange(blockF64(c), greater) + } + case core.Float32: + bv.kind = kindExtF64 + bv.f = make([]float64, count) + bv.oks = make([]bool, count) + for c := first; c < last; c++ { + bv.f[c-first], bv.oks[c-first] = core.ExtremeRangeF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c), greater) + } + case core.Float16: + bv.kind = kindExtF64 + bv.f = make([]float64, count) + bv.oks = make([]bool, count) + for c := first; c < last; c++ { + bv.f[c-first], bv.oks[c-first] = core.ExtremeRangeF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c), greater) + } + case core.Int: + bv.kind = kindExtI64 + bv.i = exactExtremes(local.RawInts()[:local.Len()], n, first, last, greater) + case core.Int8: + bv.kind = kindExtI64 + bv.i = exactExtremes(local.RawInt8s()[:local.Len()], n, first, last, greater) + case core.Uint8: + bv.kind = kindExtI64 + bv.i = exactExtremes(local.RawUint8s()[:local.Len()], n, first, last, greater) + case core.Int16: + bv.kind = kindExtI64 + bv.i = exactExtremes(local.RawInt16s()[:local.Len()], n, first, last, greater) + case core.Uint16: + bv.kind = kindExtI64 + bv.i = exactExtremes(local.RawUint16s()[:local.Len()], n, first, last, greater) + case core.Int32: + bv.kind = kindExtI64 + bv.i = exactExtremes(local.RawInt32s()[:local.Len()], n, first, last, greater) + case core.Uint32: + bv.kind = kindExtI64 + bv.i = exactExtremes(local.RawUint32s()[:local.Len()], n, first, last, greater) + default: + return blockValues{}, base.Errf("spmd: %s of dtype %s has no ordering", op, local.Dtype()) + } + case Prod: + switch local.Dtype() { + case core.Float: + bv.kind = kindProdF64 + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldProd(sliceBlock(local.RawFloats()[:local.Len()], n, first, c)) + } + case core.Float32: + bv.kind = kindProdF32 + bv.f32 = make([]float32, count) + for c := first; c < last; c++ { + bv.f32[c-first] = core.FoldProdF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c)) + } + case core.Float16: + bv.kind = kindProdHalf + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldProdF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c)) + } + case core.Int: + bv.kind = kindProdI64 + bv.i = make([]int64, count) + for c := first; c < last; c++ { + bv.i[c-first] = core.FoldProdI64(sliceBlock(local.RawInts()[:local.Len()], n, first, c)) + } + default: + return blockValues{}, base.Errf("spmd: prod of dtype %s is not supported; convert with Astype", local.Dtype()) + } + case Any, All: + if local.Dtype() != core.Bool { + return blockValues{}, base.Errf("spmd: %s needs Bool elements, got %s", op, local.Dtype()) + } + bv.kind = kindFlag + flag := op == All + for _, v := range local.RawBools()[:local.Len()] { + if op == Any && v { + flag = true + break + } + if op == All && !v { + flag = false + break + } + } + bv.i = []int64{bToF(flag)} + default: + return blockValues{}, base.Errf("spmd: unknown reduction op %s", op) + } + return bv, nil +} + +// sliceBlock is block c of the canonical partition of n elements, seen +// inside a slab whose first block is first. +func sliceBlock[T any](src []T, n, first, c int) []T { + baseIdx := core.FoldBoundary(n, first) + return src[core.FoldBoundary(n, c)-baseIdx : core.FoldBoundary(n, c+1)-baseIdx] +} + +// exactSums folds each of the piece's blocks into an exact int64 sum. +func exactSums[T intPayload](src []T, n, first, last int) []int64 { + out := make([]int64, last-first) + for c := first; c < last; c++ { + var s int64 + for _, v := range sliceBlock(src, n, first, c) { + s += int64(v) + } + out[c-first] = s + } + return out +} + +// exactExtremes folds each of the piece's blocks into an exact int64 +// extremum: the values compare exactly, so the combine needs no flags. +func exactExtremes[T intPayload](src []T, n, first, last int, greater bool) []int64 { + out := make([]int64, last-first) + for c := first; c < last; c++ { + blk := sliceBlock(src, n, first, c) + best := int64(blk[0]) + for _, v := range blk[1:] { + if w := int64(v); (greater && w > best) || (!greater && w < best) { + best = w + } + } + out[c-first] = best + } + return out +} + +// payloadKind names the kind every non-empty piece carries; the pieces +// agree because the program runs one op on one dtype, and the root, whose +// own piece may be empty, reads it from the first rank that folded +// anything. +func payloadKind(all []blockValues) (byte, error) { + kind := kindNone + for _, bv := range all { + if bv.kind == kindNone { + continue + } + if kind == kindNone { + kind = bv.kind + } else if bv.kind != kind { + return 0, base.Errf("spmd: the shards disagree on the payload kind: %d against %d", bv.kind, kind) + } + } + return kind, nil +} + +// combineShards lays the ranks' pieces over the whole partition and +// combines them with the same functions the single-array reduction +// combines its chunks with. +func (w *World) combineShards(all []blockValues, span Span, op Op) (core.Scalar, error) { + first, last, parts := spanBlocks(span) + if first < 0 || last < 0 { + return core.Scalar{}, base.Errf("spmd: the root's span is not on the partition's boundaries") + } + kind, kindErr := payloadKind(all) + if kindErr != nil { + return core.Scalar{}, kindErr + } + if op != Any && op != All { + covered := 0 + for _, bv := range all { + covered += bv.blocks() + } + if covered != parts { + return core.Scalar{}, base.Errf("spmd: the shards' %d blocks do not cover the partition's %d", covered, parts) + } + // The root joins in rank order, which is the partition's + // order only when every piece starts at its own rank's first + // block; a piece naming any other start would scramble the + // tree's addends into a wrong number returned as success. + for r, bv := range all { + if bv.kind == kindNone || bv.kind == kindFlag { + continue + } + if bv.first != r*parts/w.size { + return core.Scalar{}, base.Errf("spmd: rank %d's piece starts at block %d, the partition puts it at %d", r, bv.first, r*parts/w.size) + } + } + } + switch op { + case Sum: + switch kind { + case kindSumF64: + vals := make([]float64, 0, parts) + for _, bv := range all { + vals = append(vals, bv.f...) + } + return core.FloatScalar(core.TreeSum(vals)), nil + case kindSumC128: + vals := make([]complex128, 0, parts) + for _, bv := range all { + vals = append(vals, bv.c...) + } + return core.ComplexScalar(core.TreeSum(vals)), nil + case kindSumI64: + vals := make([]int64, 0, parts) + for _, bv := range all { + vals = append(vals, bv.i...) + } + return core.IntScalar(core.TreeSum(vals)), nil + default: + return core.Scalar{}, base.Errf("spmd: sum shards disagree on the payload kind") + } + case Min, Max: + greater := op == Max + switch kind { + case kindExtF64: + vals := make([]float64, 0, parts) + oks := make([]bool, 0, parts) + for _, bv := range all { + vals = append(vals, bv.f...) + oks = append(oks, bv.oks...) + } + return core.FloatScalar(core.CombineExtrema(vals, oks, greater)), nil + case kindExtI64: + var best int64 + have := false + for _, bv := range all { + for _, v := range bv.i { + if !have || (greater && v > best) || (!greater && v < best) { + best, have = v, true + } + } + } + return core.IntScalar(best), nil + default: + return core.Scalar{}, base.Errf("spmd: %s shards disagree on the payload kind", op) + } + case Prod: + switch kind { + case kindProdF64: + vals := make([]float64, 0, parts) + for _, bv := range all { + vals = append(vals, bv.f...) + } + return core.FloatScalar(core.TreeProd(vals)), nil + case kindProdF32: + vals := make([]float32, 0, parts) + for _, bv := range all { + vals = append(vals, bv.f32...) + } + return core.FloatScalar(float64(core.TreeProd(vals))), nil + case kindProdHalf: + vals := make([]float64, 0, parts) + for _, bv := range all { + vals = append(vals, bv.f...) + } + return core.FloatScalar(core.TreeProdHalf(vals)), nil + case kindProdI64: + vals := make([]int64, 0, parts) + for _, bv := range all { + vals = append(vals, bv.i...) + } + return core.IntScalar(core.TreeProd(vals)), nil + default: + return core.Scalar{}, base.Errf("spmd: prod shards disagree on the payload kind") + } + case Any, All: + flag := op == All + for _, bv := range all { + if bv.kind != kindFlag || len(bv.i) == 0 { + continue + } + v := bv.i[0] != 0 + if op == Any && v { + flag = true + } + if op == All && !v { + flag = false + } + } + return core.IntScalar(bToF(flag)), nil + } + return core.Scalar{}, base.Errf("spmd: unknown reduction op %s", op) +} + +// bToF carries a boolean in the int64 slot the integer-class scalar +// answers through. +func bToF(b bool) int64 { + if b { + return 1 + } + return 0 +} + +// The shard contribution's wire form: first block index, payload kind, +// block count, then the values. Every field is checked against every +// other on the way back in. + +func encodeBlockValues(bv blockValues) []byte { + buf := binary.LittleEndian.AppendUint64(nil, uint64(bv.first)) + buf = append(buf, bv.kind) + buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.f))) + buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.f32))) + buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.c))) + buf = binary.LittleEndian.AppendUint64(buf, uint64(len(bv.i))) + for _, v := range bv.f { + buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(v)) + } + for _, v := range bv.c { + buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(real(v))) + buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(imag(v))) + } + for _, v := range bv.i { + buf = binary.LittleEndian.AppendUint64(buf, uint64(v)) + } + for _, v := range bv.f32 { + buf = binary.LittleEndian.AppendUint32(buf, math.Float32bits(v)) + } + for _, ok := range bv.oks { + b := byte(0) + if ok { + b = 1 + } + buf = append(buf, b) + } + return buf +} + +func decodeBlockValues(data []byte) (blockValues, error) { + const head = 8 + 1 + 8*4 + if len(data) < head { + return blockValues{}, base.Errf("spmd: a shard contribution of %d bytes is shorter than its head", len(data)) + } + // Every count is bounded as its unsigned wire value, before any + // signed conversion: a top-bit count would come back negative and + // sneak past an upper-bound check. + firstU := binary.LittleEndian.Uint64(data) + if firstU > math.MaxInt { + return blockValues{}, base.Errf("spmd: a shard contribution names a first block of %d", firstU) + } + bv := blockValues{first: int(firstU)} + bv.kind = data[8] + nf, n32, nc, ni, err := decodeCounts(data[9:]) + if err != nil { + return blockValues{}, err + } + data = data[head:] + want := (nf+ni)*8 + n32*4 + nc*16 + if bv.kind == kindExtF64 { + want += nf + } + if bv.kind == kindArg { + want += 2 + } + if len(data) != want { + return blockValues{}, base.Errf("spmd: a shard contribution carries %d value bytes for %d floats, %d float32s, %d complexes and %d integers", + len(data), nf, n32, nc, ni) + } + bv.f = make([]float64, nf) + for i := range bv.f { + bv.f[i] = math.Float64frombits(binary.LittleEndian.Uint64(data[i*8:])) + } + data = data[nf*8:] + bv.c = make([]complex128, nc) + for i := range bv.c { + re := math.Float64frombits(binary.LittleEndian.Uint64(data[i*16:])) + im := math.Float64frombits(binary.LittleEndian.Uint64(data[i*16+8:])) + bv.c[i] = complex(re, im) + } + data = data[nc*16:] + bv.i = make([]int64, ni) + for i := range bv.i { + bv.i[i] = int64(binary.LittleEndian.Uint64(data[i*8:])) + } + data = data[ni*8:] + bv.f32 = make([]float32, n32) + for i := range bv.f32 { + bv.f32[i] = math.Float32frombits(binary.LittleEndian.Uint32(data[i*4:])) + } + data = data[n32*4:] + if bv.kind == kindExtF64 { + bv.oks = make([]bool, nf) + for i := range bv.oks { + switch data[i] { + case 0: + case 1: + bv.oks[i] = true + default: + return blockValues{}, base.Errf("spmd: a shard's candidate flag %d is neither 0 nor 1", data[i]) + } + } + } + if bv.kind == kindArg { + bv.oks = make([]bool, 2) + for i := range bv.oks { + switch data[i] { + case 0: + case 1: + bv.oks[i] = true + default: + return blockValues{}, base.Errf("spmd: a shard's candidate flag %d is neither 0 nor 1", data[i]) + } + } + } + return bv, nil +} + +// decodeCounts reads the four block counts as unsigned wire values +// and bounds each before any signed conversion: a top-bit count would +// come back negative and sneak past an upper-bound check. +func decodeCounts(data []byte) (nf, n32, nc, ni int, err error) { + raw := [4]uint64{ + binary.LittleEndian.Uint64(data[0:8]), + binary.LittleEndian.Uint64(data[8:16]), + binary.LittleEndian.Uint64(data[16:24]), + binary.LittleEndian.Uint64(data[24:32]), + } + for k, u := range raw { + if u > maxBlocks { + return 0, 0, 0, 0, base.Errf("spmd: a shard contribution names an impossible count %d in slot %d", u, k) + } + } + return int(raw[0]), int(raw[1]), int(raw[2]), int(raw[3]), nil +} + +// maxBlocks bounds the counts a contribution may name before anything +// is allocated: the partition itself never exceeds this. +const maxBlocks = 1 << 12 + +// The scalar answer travels as one fixed frame: a kind byte, then the +// int64, float64 and complex slots, each little-endian. Only the slot +// the kind names is meaningful; every slot always travels, so the form +// has no variable length to get wrong. +const ( + scalarKindInt byte = 1 + scalarKindFloat byte = 2 + scalarKindComplex byte = 3 +) + +const scalarFrameLen = 1 + 8 + 8 + 16 + +func encodeScalar(s core.Scalar) []byte { + buf := make([]byte, scalarFrameLen) + switch { + case s.IsComplex(): + buf[0] = scalarKindComplex + binary.LittleEndian.PutUint64(buf[17:], math.Float64bits(real(s.Complex()))) + binary.LittleEndian.PutUint64(buf[25:], math.Float64bits(imag(s.Complex()))) + case s.IsFloat(): + buf[0] = scalarKindFloat + binary.LittleEndian.PutUint64(buf[9:], math.Float64bits(s.Float())) + default: + buf[0] = scalarKindInt + binary.LittleEndian.PutUint64(buf[1:], uint64(s.Int())) + } + return buf +} + +func sameShape(a, b []int) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +func decodeScalar(data []byte) (core.Scalar, error) { + if len(data) != scalarFrameLen { + return core.Scalar{}, base.Errf("spmd: a scalar answer of %d bytes is not the fixed %d", len(data), scalarFrameLen) + } + switch data[0] { + case scalarKindInt: + return core.IntScalar(int64(binary.LittleEndian.Uint64(data[1:]))), nil + case scalarKindFloat: + return core.FloatScalar(math.Float64frombits(binary.LittleEndian.Uint64(data[9:]))), nil + case scalarKindComplex: + re := math.Float64frombits(binary.LittleEndian.Uint64(data[17:])) + im := math.Float64frombits(binary.LittleEndian.Uint64(data[25:])) + return core.ComplexScalar(complex(re, im)), nil + } + return core.Scalar{}, base.Errf("spmd: a scalar answer of kind %d is none of the library's", data[0]) +} + +// broadcastScalar carries the root's answer to every rank. +func (w *World) broadcastScalar(s core.Scalar, root int) (core.Scalar, error) { + if w.size == 1 { + return s, nil + } + if w.rank == root { + wire := encodeScalar(s) + for r := range w.size { + if r == root { + continue + } + if err := w.sendTo(r, tagShardsWhole, wire); err != nil { + return core.Scalar{}, err + } + } + return s, nil + } + data, err := w.recvFrom(root, tagShardsWhole) + if err != nil { + return core.Scalar{}, err + } + got, err := decodeScalar(data) + if err != nil { + return core.Scalar{}, w.fail(err) + } + return got, nil +} + +// Reduce folds the ranks' same-shaped arrays together and answers on +// the root alone: element e of the answer is the fold of the ranks' +// elements e in rank index order, left to right. The fold's order is +// the program's own rank order, never the arrival order, so the same +// program and data answer the same bits on one machine and across a +// cluster. Every other rank receives nil. +func (w *World) Reduce(a *core.Array, op Op, root int) (*core.Array, error) { + if err := w.status(); err != nil { + return nil, err + } + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + if err := checkArrayOp(a, op); err != nil { + return nil, w.fail(err) + } + if w.size == 1 { + return a, nil + } + k := w.size + n := a.Len() + bounds := func(c int) (int, int) { return c * n / k, (c + 1) * n / k } + // Every rank hands its chunk c to the rank that owns c. + for c := range k { + if c == w.rank { + continue + } + lo, hi := bounds(c) + wire, err := encodePart(nil, a, []int{hi - lo}, lo, hi-lo) + if err != nil { + return nil, w.fail(err) + } + if err := w.sendTo(c, tagReduce, wire); err != nil { + return nil, err + } + } + // The fold runs as the chunks land: the rank's own chunk seeds the + // accumulator, rank 0's chunk replaces it when it arrives, and + // every later chunk joins the running accumulator at its rank's + // turn. The chain of operands is chunk 0 left to chunk k-1 right, + // exactly the rank order it always was, so the bits cannot move, + // while folding one chunk overlaps the transfer of the next and + // no rank piles the whole exchange up before it computes. + lo, hi := bounds(w.rank) + ownWire, err := encodePart(nil, a, []int{hi - lo}, lo, hi-lo) + if err != nil { + return nil, w.fail(err) + } + own, err := decodeWire(ownWire) + if err != nil { + return nil, w.fail(err) + } + acc := own + for i := range k { + chunk := own + if i != w.rank { + data, rerr := w.recvFrom(i, tagReduce) + if rerr != nil { + return nil, rerr + } + chunk, rerr = decodeWire(data) + if rerr != nil { + return nil, w.fail(rerr) + } + if chunk.Dtype() != own.Dtype() { + return nil, w.fail(base.Errf("spmd: rank %d's %s disagrees with rank %d's %s", + i, chunk.Dtype(), w.rank, own.Dtype())) + } + } + if i == 0 { + acc = chunk // the chain starts at rank 0's chunk + continue + } + var ferr error + switch op { + case Sum: + acc, ferr = core.Add(acc, chunk) + case Min: + acc, ferr = core.Minimum(acc, chunk) + case Max: + acc, ferr = core.Maximum(acc, chunk) + case Any: + acc, ferr = core.Or(acc, chunk) + case All: + acc, ferr = core.And(acc, chunk) + case Prod: + acc, ferr = core.Mul(acc, chunk) + } + if ferr != nil { + return nil, w.fail(ferr) + } + } + // The folded chunk travels back to the root, which joins the + // world's chunks in rank order and reshapes to the input's shape. + flat, err := w.gather(acc, root) + if err != nil { + return nil, err + } + if w.rank != root { + return nil, nil + } + return core.Reshape(flat, a.Shape()...) +} + +// AllReduce is Reduce with the answer on every rank: one world, one +// answer, identical bits everywhere. +func (w *World) AllReduce(a *core.Array, op Op) (*core.Array, error) { + whole, err := w.Reduce(a, op, 0) + if err != nil { + return nil, err + } + return w.Broadcast(whole, 0) +} + +// checkArrayOp refuses the combinations that carry no arithmetic +// before any frame moves: Bool has no sum and no ordering, complex has +// no ordering, and Any with All are Bool's alone. +func checkArrayOp(a *core.Array, op Op) error { + dt := a.Dtype() + switch op { + case Sum: + if dt == core.Bool { + return base.Errf("spmd: sum of Bool arrays has no arithmetic; Any and All are Bool's reductions") + } + case Prod: + if dt == core.Bool { + return base.Errf("spmd: prod of Bool arrays has no arithmetic; Any and All are Bool's reductions") + } + case Min, Max: + if dt == core.Bool || dt == core.Complex { + return base.Errf("spmd: %s of dtype %s has no ordering", op, dt) + } + case Any, All: + if dt != core.Bool { + return base.Errf("spmd: %s needs Bool elements, got %s", op, dt) + } + default: + return base.Errf("spmd: unknown reduction op %s", op) + } + return nil +} + +// The sharded product, norm and dot extend the reduction families to +// the vector operations the iterative solvers live on. All three carry +// one-dimensional arrays: the canonical partition cuts the axis, the +// shards fold their own blocks with the single block kernels, and the +// combine is the same tree the single-array answers use. + +// AllReduceNormShards answers the Lp norm of one global array whose +// canonical pieces the ranks hold, every rank receiving the answer. +// The power sums fold the canonical blocks with the norm fold's own +// per-element arithmetic and combine through the same balanced tree, +// so the answer is the single-array norm's exact bits at any world +// size. p must be finite and positive; the infinity norm is a maximum +// and belongs to Max, not here. +func (w *World) AllReduceNormShards(local *core.Array, span Span, p float64) (core.Scalar, error) { + ans, err := w.reduceNormShards(local, span, p, 0) + if err != nil { + return core.Scalar{}, err + } + return w.broadcastScalar(ans, 0) +} + +// ReduceNormShards is AllReduceNormShards with the answer on the root +// alone; every other rank receives nil. +func (w *World) ReduceNormShards(local *core.Array, span Span, p float64, root int) (*core.Scalar, error) { + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + ans, err := w.reduceNormShards(local, span, p, root) + if err != nil { + return nil, err + } + if w.rank != root { + return nil, nil + } + return &ans, nil +} + +func (w *World) reduceNormShards(local *core.Array, span Span, p float64, root int) (core.Scalar, error) { + if err := w.status(); err != nil { + return core.Scalar{}, err + } + if math.IsNaN(p) || p <= 0 || math.IsInf(p, 1) { + return core.Scalar{}, w.fail(base.Errf("spmd: Norm shards need a finite positive p, got %v", p)) + } + if err := w.checkOneDimSpan(local, span); err != nil { + return core.Scalar{}, w.fail(err) + } + if dt := local.Dtype(); dt != core.Float && dt != core.Float32 && dt != core.Float16 && dt != core.Int { + return core.Scalar{}, w.fail(base.Errf("spmd: norm of dtype %s is not supported; convert with Astype", dt)) + } + if span.Global == 0 { + return core.FloatScalar(0), nil + } + sums, err := w.foldNormBlocks(local, span, p) + if err != nil { + return core.Scalar{}, err + } + if w.rank != root { + if err := w.sendTo(root, tagShardsValues, encodeBlockValues(blockValues{first: blockFirst(span), kind: kindSumF64, f: sums})); err != nil { + return core.Scalar{}, err + } + return core.Scalar{}, nil + } + all := make([][]float64, w.size) + all[w.rank] = sums + for r := range w.size { + if r == root { + continue + } + data, err := w.recvFrom(r, tagShardsValues) + if err != nil { + return core.Scalar{}, err + } + bv, err := decodeBlockValues(data) + if err != nil { + return core.Scalar{}, w.fail(err) + } + if bv.kind != kindSumF64 { + return core.Scalar{}, w.fail(base.Errf("spmd: norm shards disagree on the payload kind")) + } + all[r] = bv.f + } + total := make([]float64, 0, core.FoldParts(span.Global)) + for _, part := range all { + total = append(total, part...) + } + return core.FloatScalar(core.NormRoot(core.TreeSum(total), p)), nil +} + +// foldNormBlocks computes the rank's own blocks' power sums with the +// norm fold's per-element arithmetic, dtype by dtype. +func (w *World) foldNormBlocks(local *core.Array, span Span, p float64) ([]float64, error) { + first, last, _ := spanBlocks(span) + if first < 0 || last < 0 || first > last { + return nil, base.Errf("spmd: rank %d's span [%d, %d) is not on the partition's block boundaries", + w.rank, span.Lo, span.Hi) + } + n := span.Global + count := last - first + sums := make([]float64, count) + for c := first; c < last; c++ { + switch local.Dtype() { + case core.Float: + sums[c-first] = core.FoldNormPower(sliceBlock(local.RawFloats()[:local.Len()], n, first, c), p) + case core.Float32: + sums[c-first] = core.FoldNormPowerF32(sliceBlock(local.RawFloat32s()[:local.Len()], n, first, c), p) + case core.Float16: + sums[c-first] = core.FoldNormPowerF16(sliceBlock(local.RawHalves()[:local.Len()], n, first, c), p) + case core.Int: + sums[c-first] = core.FoldNormPowerI64(sliceBlock(local.RawInts()[:local.Len()], n, first, c), p) + default: + return nil, base.Errf("spmd: norm of dtype %s is not supported; convert with Astype", local.Dtype()) + } + } + return sums, nil +} + +// blockFirst names the partition block a span starts on. +func blockFirst(span Span) int { + first, _, _ := spanBlocks(span) + return first +} + +// AllReduceDotShards answers the dot product of two global arrays the +// ranks hold as the same canonical pieces, every rank receiving the +// answer. The shards fold their own blocks with the single block dot +// kernel and the block values combine through the same balanced tree, +// so the answer is the single-array Dot's exact bits at any world +// size. The two arrays must carry the same dtype. +func (w *World) AllReduceDotShards(x, y *core.Array, span Span) (core.Scalar, error) { + ans, err := w.reduceDotShards(x, y, span, 0) + if err != nil { + return core.Scalar{}, err + } + return w.broadcastScalar(ans, 0) +} + +// ReduceDotShards is AllReduceDotShards with the answer on the root +// alone; every other rank receives nil. +func (w *World) ReduceDotShards(x, y *core.Array, span Span, root int) (*core.Scalar, error) { + if err := checkRoot(root, w.size); err != nil { + return nil, w.fail(err) + } + ans, err := w.reduceDotShards(x, y, span, root) + if err != nil { + return nil, err + } + if w.rank != root { + return nil, nil + } + return &ans, nil +} + +func (w *World) reduceDotShards(x, y *core.Array, span Span, root int) (core.Scalar, error) { + if err := w.status(); err != nil { + return core.Scalar{}, err + } + if err := w.checkOneDimSpan(x, span); err != nil { + return core.Scalar{}, w.fail(err) + } + if y.NDim() != 1 || y.Len() != span.Len() { + return core.Scalar{}, w.fail(base.Errf("spmd: the second shard leads with %d elements against the span's %d", y.Len(), span.Len())) + } + if x.Dtype() != y.Dtype() { + return core.Scalar{}, w.fail(base.Errf("spmd: Dot shards carry %s against %s; convert with Astype", x.Dtype(), y.Dtype())) + } + switch x.Dtype() { + case core.Int: + case core.Float, core.Float32, core.Float16, core.Complex: + default: + return core.Scalar{}, w.fail(base.Errf("spmd: dot of dtype %s is not supported; convert with Astype", x.Dtype())) + } + if span.Global == 0 { + if x.Dtype() == core.Int { + return core.IntScalar(0), nil + } + if x.Dtype() == core.Complex { + return core.ComplexScalar(0), nil + } + return core.FloatScalar(0), nil + } + bv, err := w.foldDotBlocks(x, y, span) + if err != nil { + return core.Scalar{}, w.fail(err) + } + if w.rank != root { + if err := w.sendTo(root, tagShardsValues, encodeBlockValues(bv)); err != nil { + return core.Scalar{}, err + } + return core.Scalar{}, nil + } + all := make([]blockValues, w.size) + all[w.rank] = bv + for r := range w.size { + if r == root { + continue + } + data, err := w.recvFrom(r, tagShardsValues) + if err != nil { + return core.Scalar{}, err + } + if all[r], err = decodeBlockValues(data); err != nil { + return core.Scalar{}, w.fail(err) + } + } + return w.combineDot(all, span) +} + +// foldDotBlocks computes the rank's own blocks' dot values with the +// single block dot kernel, dtype by dtype. +func (w *World) foldDotBlocks(x, y *core.Array, span Span) (blockValues, error) { + first, last, _ := spanBlocks(span) + if first < 0 || last < 0 || first > last { + return blockValues{}, base.Errf("spmd: rank %d's span [%d, %d) is not on the partition's block boundaries", + w.rank, span.Lo, span.Hi) + } + n := span.Global + count := last - first + bv := blockValues{first: first} + if count == 0 { + return bv, nil + } + switch x.Dtype() { + case core.Float: + bv.kind = kindSumF64 + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldDot( + sliceBlock(x.RawFloats()[:x.Len()], n, first, c), + sliceBlock(y.RawFloats()[:y.Len()], n, first, c)) + } + case core.Float32: + bv.kind = kindSumF64 + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldDotF32( + sliceBlock(x.RawFloat32s()[:x.Len()], n, first, c), + sliceBlock(y.RawFloat32s()[:y.Len()], n, first, c)) + } + case core.Float16: + bv.kind = kindSumF64 + bv.f = make([]float64, count) + for c := first; c < last; c++ { + bv.f[c-first] = core.FoldDotF16( + sliceBlock(x.RawHalves()[:x.Len()], n, first, c), + sliceBlock(y.RawHalves()[:y.Len()], n, first, c)) + } + case core.Complex: + bv.kind = kindSumC128 + bv.c = make([]complex128, count) + for c := first; c < last; c++ { + bv.c[c-first] = core.FoldDotC128( + sliceBlock(x.RawComplexes()[:x.Len()], n, first, c), + sliceBlock(y.RawComplexes()[:y.Len()], n, first, c)) + } + case core.Int: + bv.kind = kindSumI64 + bv.i = make([]int64, count) + for c := first; c < last; c++ { + bv.i[c-first] = core.FoldDotI64( + sliceBlock(x.RawInts()[:x.Len()], n, first, c), + sliceBlock(y.RawInts()[:y.Len()], n, first, c)) + } + default: + return blockValues{}, base.Errf("spmd: dot of dtype %s is not supported; convert with Astype", x.Dtype()) + } + return bv, nil +} + +// combineDot folds the shards' dot values through the balanced tree. +func (w *World) combineDot(all []blockValues, span Span) (core.Scalar, error) { + first, _, parts := spanBlocks(span) + if first < 0 { + return core.Scalar{}, base.Errf("spmd: the root's span is not on the partition's boundaries") + } + covered := 0 + for r, bv := range all { + covered += bv.blocks() + if bv.kind == kindNone { + continue + } + if bv.first != r*parts/w.size { + return core.Scalar{}, base.Errf("spmd: rank %d's piece starts at block %d, the partition puts it at %d", r, bv.first, r*parts/w.size) + } + } + if covered != parts { + return core.Scalar{}, base.Errf("spmd: the shards' %d blocks do not cover the partition's %d", covered, parts) + } + kind, kindErr := payloadKind(all) + if kindErr != nil { + return core.Scalar{}, kindErr + } + switch kind { + case kindSumI64: + vals := make([]int64, 0, parts) + for _, bv := range all { + vals = append(vals, bv.i...) + } + return core.IntScalar(core.TreeSum(vals)), nil + case kindSumC128: + vals := make([]complex128, 0, parts) + for _, bv := range all { + vals = append(vals, bv.c...) + } + return core.ComplexScalar(core.TreeSum(vals)), nil + case kindSumF64: + vals := make([]float64, 0, parts) + for _, bv := range all { + vals = append(vals, bv.f...) + } + return core.FloatScalar(core.TreeSum(vals)), nil + default: + return core.Scalar{}, base.Errf("spmd: dot shards disagree on the payload kind") + } +} diff --git a/spmd/reduce_test.go b/spmd/reduce_test.go new file mode 100644 index 0000000..ecc7990 --- /dev/null +++ b/spmd/reduce_test.go @@ -0,0 +1,886 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "encoding/binary" + "math" + "strings" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The sharded reduction tests judge the whole contract in one claim: +// whatever the world's size, the sharded answer carries the +// single-array reduction's exact bits. The single-array answer is +// computed from the same fixture beside the test, so the comparison is +// the contract and not a self-check. + +// sparseFixture spreads magnitudes and drops NaNs and an infinity at +// fixed spots, so the folds meet cancellation and the extremum rules +// meet their NaN edges. +func sparseFixture(n int) []float64 { + v := make([]float64, n) + for i := range v { + v[i] = float64((i*6559)%2001-1000) * math.Pow(10, float64(i%7)-3) + if i%401 == 3 { + v[i] = math.NaN() + } + if i%401 == 200 { + v[i] = math.Inf(-1) + } + } + return v +} + +// fixtureDtypes builds the same data in every element type the shard +// reductions carry. +func fixtureDtypes(t *testing.T, gn int) map[core.Dtype]*core.Array { + t.Helper() + build := func(a *core.Array, err error) *core.Array { + if err != nil { + t.Fatal(err) + } + return a + } + vals := sparseFixture(gn) + ints := make([]int64, gn) + bools := make([]bool, gn) + for i := range ints { + ints[i] = int64((i*13)%97) - 48 + bools[i] = i%3 == 0 + } + halves := make([]uint16, gn) + f32s := make([]float32, gn) + for i := range halves { + halves[i] = core.HalfFromFloat64(vals[i]) + f32s[i] = float32(vals[i]) + } + complexes := make([]complex128, gn) + for i := range complexes { + complexes[i] = complex(vals[i], -vals[i]/2) + } + return map[core.Dtype]*core.Array{ + core.Float: build(core.FromFloats(vals, gn)), + core.Float32: build(core.FromFloat32s(f32s, gn)), + core.Float16: build(core.HalvesFromArray(halves, gn)), + core.Complex: build(core.FromComplexes(complexes, gn)), + core.Int: build(core.FromInts(ints, gn)), + core.Int8: build(core.FromInt8s(narrow8s(ints), gn)), + core.Uint8: build(core.FromUint8s(narrowu8s(ints), gn)), + core.Int16: build(core.FromInt16s(narrow16s(ints), gn)), + core.Uint16: build(core.FromUint16s(narrowu16s(ints), gn)), + core.Int32: build(core.FromInt32s(narrow32s(ints), gn)), + core.Uint32: build(core.FromUint32s(narrowu32s(ints), gn)), + core.Bool: build(core.FromBools(bools, gn)), + } +} + +func narrow8s(v []int64) []int8 { + out := make([]int8, len(v)) + for i := range v { + out[i] = int8(v[i]) + } + return out +} + +func narrowu8s(v []int64) []uint8 { + out := make([]uint8, len(v)) + for i := range v { + out[i] = uint8(v[i]) + } + return out +} + +func narrow16s(v []int64) []int16 { + out := make([]int16, len(v)) + for i := range v { + out[i] = int16(v[i]) + } + return out +} + +func narrowu16s(v []int64) []uint16 { + out := make([]uint16, len(v)) + for i := range v { + out[i] = uint16(v[i]) + } + return out +} + +func narrow32s(v []int64) []int32 { + out := make([]int32, len(v)) + for i := range v { + out[i] = int32(v[i]) + } + return out +} + +func narrowu32s(v []int64) []uint32 { + out := make([]uint32, len(v)) + for i := range v { + out[i] = uint32(v[i]) + } + return out +} + +// narrowSliceFor deals any dtype's rows off the whole array, keeping +// the bits exactly: what Scatter deals, built without the collectives +// under test. +func narrowSliceFor(t *testing.T, whole *core.Array, span Span) *core.Array { + t.Helper() + wire, err := encodePart(nil, whole, append([]int{span.Len()}, whole.Shape()[1:]...), span.Lo, span.Len()) + if err != nil { + t.Fatal(err) + } + a, err := decodeWire(wire) + if err != nil { + t.Fatal(err) + } + return a +} + +// scalarBits renders a scalar's exact bits, NaN payloads and signed +// zeros included, so the comparison is bit for bit. +func scalarBits(s core.Scalar) string { + switch { + case s.IsComplex(): + c := s.Complex() + return "c" + fBits(real(c)) + "/" + fBits(imag(c)) + case s.IsFloat(): + return "f" + fBits(s.Float()) + default: + return "i" + iFormat(s.Int()) + } +} + +func fBits(f float64) string { return iFormat(int64(math.Float64bits(f))) } + +func iFormat(i int64) string { + if i < 0 { + return "-" + uFormat(-uint64(i)) + } + return uFormat(uint64(i)) +} + +func uFormat(u uint64) string { + if u == 0 { + return "0" + } + var b []byte + for u > 0 { + b = append([]byte{byte('0' + u%10)}, b...) + u /= 10 + } + return string(b) +} + +// TestShardedSumIsTheSingleArraySum is the headline claim of the whole +// package: shard the data, reduce the shards, and the bits are the +// single array's bits, for every world size, every element type and +// every length that touches the partition's edges. +func TestShardedSumIsTheSingleArraySum(t *testing.T) { + for _, size := range []int{1, 2, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537, 131073, 327681} { + for dt, whole := range fixtureDtypes(t, gn) { + want := core.Sum(whole) + err := Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceShards(local, span, Sum) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("size %d gn %d %s: sharded %s against single-array %s", + w.Size(), gn, dt, scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err) + } + } + } + } +} + +// TestShardedExtremaAreTheSingleArrayExtrema pins Min and Max against +// the single-array walk, NaN rules and the all-NaN fallback included. +func TestShardedExtremaAreTheSingleArrayExtrema(t *testing.T) { + for _, size := range []int{1, 2, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537, 131073} { + for _, dt := range []core.Dtype{core.Float, core.Float32, core.Float16, core.Int, core.Int8} { + whole := fixtureDtypes(t, gn)[dt] + for _, op := range []Op{Min, Max} { + want, err := (func() (core.Scalar, error) { + if op == Min { + return core.Min(whole) + } + return core.Max(whole) + })() + if err != nil { + t.Fatal(err) + } + err = Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceShards(local, span, op) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("size %d gn %d %s %s: sharded %s against single-array %s", + w.Size(), gn, dt, op, scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d %s %s: %v", size, gn, dt, op, err) + } + } + } + } + } + // The all-NaN array: every block reports no candidate, so the answer + // is the last block's fallback, the single-array walk's own rule. + const gn = 131073 // two blocks + allNaN := make([]float64, gn) + for i := range allNaN { + allNaN[i] = math.NaN() + } + nans := mk(core.FromFloats(allNaN, gn)) + for _, op := range []Op{Min, Max} { + want, err := (func() (core.Scalar, error) { + if op == Min { + return core.Min(nans) + } + return core.Max(nans) + })() + if err != nil { + t.Fatal(err) + } + err = Launch(3, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, nans, span) + got, err := w.AllReduceShards(local, span, op) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("all-NaN %s: sharded %s against single-array %s", + op, scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatalf("all-NaN %s: %v", op, err) + } + } +} + +// TestShardedBoolReductions pins Any and All against the single-array +// answers. +func TestShardedBoolReductions(t *testing.T) { + for _, size := range []int{1, 2, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537} { + whole := fixtureDtypes(t, gn)[core.Bool] + anyWant, _ := core.Any(whole) + allWant, _ := core.All(whole) + err := Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + anyGot, err := w.AllReduceShards(local, span, Any) + if err != nil { + return err + } + allGot, err := w.AllReduceShards(local, span, All) + if err != nil { + return err + } + if anyGot.Int() != bToF(anyWant) || allGot.Int() != bToF(allWant) { + t.Fatalf("size %d gn %d: any %d all %d against %v/%v", + w.Size(), gn, anyGot.Int(), allGot.Int(), anyWant, allWant) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d: %v", size, gn, err) + } + } + } +} + +// TestReduceShardsAnswersTheRootAlone pins the root-directed shape of +// the collective: the root holds the answer, nobody else holds +// anything. +func TestReduceShardsAnswersTheRootAlone(t *testing.T) { + const gn = 131074 + whole := fixtureDtypes(t, gn)[core.Float] + want := core.Sum(whole) + err := Launch(3, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.ReduceShards(local, span, Sum, 2) + if err != nil { + return err + } + if w.Rank() == 2 { + if got == nil || scalarBits(*got) != scalarBits(want) { + t.Fatalf("the root's sharded sum against the single-array %s", scalarBits(want)) + } + return nil + } + if got != nil { + t.Fatalf("rank %d received an answer it should not have", w.Rank()) + } + return nil + }) + if err != nil { + t.Fatal(err) + } +} + +// TestShardsRefuseAForeignSpan: a span the partition did not cut is an +// error naming the canonical boundaries, never a number. +func TestShardsRefuseAForeignSpan(t *testing.T) { + const gn = 65537 + whole := fixtureDtypes(t, gn)[core.Float] + err := Launch(3, func(w *World) error { + span := mustPartition(t, gn, 3, w.Rank()) + local := narrowSliceFor(t, whole, mustPartition(t, gn, 3, w.Rank())) + if w.Rank() == 1 { + // A plausible but foreign cut: one element over. + span.Lo++ + } + _, err := w.AllReduceShards(local, span, Sum) + return err + }) + if err == nil || !strings.Contains(err.Error(), "canonical partition") { + t.Fatalf("a foreign span did not name the canonical boundaries: %v", err) + } +} + +// TestShardedSumOverTCP runs the headline claim over real connections. +func TestShardedSumOverTCP(t *testing.T) { + const gn = 131073 + whole := fixtureDtypes(t, gn)[core.Float] + want := core.Sum(whole) + runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceShards(local, span, Sum) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("rank %d: sharded %s against single-array %s", w.Rank(), scalarBits(got), scalarBits(want)) + } + return nil + }) +} + +// TestAllReduceFoldsInRankOrder pins family A: elementwise across the +// ranks' same-shaped arrays, folded in rank index order, which the +// expected answer rebuilds with the core's own pairwise ops. +func TestAllReduceFoldsInRankOrder(t *testing.T) { + for _, size := range []int{1, 2, 3, 5} { + rankArrays := make([]*core.Array, size) + for r := range rankArrays { + vals := make([]float64, 100) + for i := range vals { + vals[i] = float64(r+1) * float64((i*17)%31-15) / 3.0 + } + rankArrays[r] = mk(core.FromFloats(vals, 100)) + } + want := rankArrays[0] + for _, a := range rankArrays[1:] { + next, err := core.Add(want, a) + if err != nil { + t.Fatal(err) + } + want = next + } + wantMin := rankArrays[0] + for _, a := range rankArrays[1:] { + next, err := core.Minimum(wantMin, a) + if err != nil { + t.Fatal(err) + } + wantMin = next + } + err := Launch(size, func(w *World) error { + got, err := w.AllReduce(rankArrays[w.Rank()], Sum) + if err != nil { + return err + } + if !sameBits(want, got) { + t.Fatalf("size %d: the AllReduce sum differs from the rank-order fold", size) + } + got, err = w.AllReduce(rankArrays[w.Rank()], Min) + if err != nil { + return err + } + if !sameBits(wantMin, got) { + t.Fatalf("size %d: the AllReduce min differs from the rank-order fold", size) + } + return nil + }) + if err != nil { + t.Fatalf("size %d: %v", size, err) + } + } +} + +// TestShardContributionRefusesTheHostile: the counts the wire names +// are bounded as unsigned values before anything is allocated, so a +// top-bit count is refused rather than turned into a negative length, +// and the candidate flags of an extremum contribution are read back as +// the 0/1 bytes the encoder writes, an honest run and an empty one +// included. +func TestShardContributionRefusesTheHostile(t *testing.T) { + head := make([]byte, 8+1+8*4) + head[8] = kindSumF64 + binary.LittleEndian.PutUint64(head[9:], uint64(1)<<63) + if _, err := decodeBlockValues(head); err == nil { + t.Fatal("a top-bit float count decoded") + } + honest := binary.LittleEndian.AppendUint64(nil, 0) + honest = append(honest, kindSumF64) + honest = binary.LittleEndian.AppendUint64(honest, 1) + honest = binary.LittleEndian.AppendUint64(honest, 0) + honest = binary.LittleEndian.AppendUint64(honest, 0) + honest = binary.LittleEndian.AppendUint64(honest, 0) + honest = binary.LittleEndian.AppendUint64(honest, math.Float64bits(1)) + bv, err := decodeBlockValues(honest) + if err != nil || len(bv.f) != 1 || bv.f[0] != 1 { + t.Fatalf("an honest contribution was refused: %v", err) + } + extremum := func(flags ...byte) []byte { + buf := binary.LittleEndian.AppendUint64(nil, 0) + buf = append(buf, kindExtF64) + buf = binary.LittleEndian.AppendUint64(buf, uint64(len(flags))) + buf = binary.LittleEndian.AppendUint64(buf, 0) + buf = binary.LittleEndian.AppendUint64(buf, 0) + buf = binary.LittleEndian.AppendUint64(buf, 0) + for range flags { + buf = binary.LittleEndian.AppendUint64(buf, math.Float64bits(1)) + } + return append(buf, flags...) + } + if _, err := decodeBlockValues(extremum()); err != nil { + t.Fatalf("an empty flag run was refused: %v", err) + } + bv, err = decodeBlockValues(extremum(1, 0)) + if err != nil { + t.Fatalf("honest candidate flags were refused: %v", err) + } + if !bv.oks[0] || bv.oks[1] { + t.Fatalf("the candidate flags came back as %v", bv.oks) + } + if _, err := decodeBlockValues(extremum(1, 2)); err == nil { + t.Fatal("a candidate flag byte of 2 decoded") + } +} + +// TestShardedSumAtZeroLength: the degenerate global axis answers the +// single-array reduction's zero at any world size, no exchange needed. +func TestShardedSumAtZeroLength(t *testing.T) { + for _, size := range []int{1, 2, 3, 5} { + empty := mk(core.FromFloats(nil, 0)) + want := core.Sum(empty) + err := Launch(size, func(w *World) error { + span := mustPartition(t, 0, w.Size(), w.Rank()) + got, err := w.AllReduceShards(empty, span, Sum) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("size %d: sharded %s against single-array %s", + w.Size(), scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatalf("size %d: %v", size, err) + } + } +} + +// TestShardedExtremaReachTheRemotePieces pins the candidate flags' +// travel: the extremum living strictly on a non-root rank, without a +// tie to mask it, must win the combine whoever holds it. The saturating +// fixture above cannot tell this, because its extrema repeat in every +// window and the tie rule hands the answer to the root anyway. +func TestShardedExtremaReachTheRemotePieces(t *testing.T) { + // Ascending data: the maximum lives on the last rank alone, and the + // minimum on the root. The root holds its own block here. + const gn = 131073 + asc := make([]float64, gn) + for i := range asc { + asc[i] = float64(i) + } + whole := mk(core.FromFloats(asc, gn)) + for _, op := range []Op{Min, Max} { + want, err := (func() (core.Scalar, error) { + if op == Min { + return core.Min(whole) + } + return core.Max(whole) + })() + if err != nil { + t.Fatal(err) + } + err = Launch(2, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceShards(local, span, op) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("%s: sharded %s against single-array %s", + op, scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatalf("%s: %v", op, err) + } + } + // A world whose root holds nothing at all (four ranks over three + // blocks) and whose maximum lives only in the middle block: with + // every root-side candidate absent, the combine must still take the + // middle block's value, never the last block's. + const gn2 = 131073 + pyr := make([]float64, gn2) + for i := range pyr { + d := i - gn2/2 + if d < 0 { + d = -d + } + pyr[i] = -float64(d) + } + peak := mk(core.FromFloats(pyr, gn2)) + want, err := core.Max(peak) + if err != nil { + t.Fatal(err) + } + err = Launch(4, func(w *World) error { + span := mustPartition(t, gn2, w.Size(), w.Rank()) + local := narrowSliceFor(t, peak, span) + got, err := w.AllReduceShards(local, span, Max) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("peak: sharded %s against single-array %s", + scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatal(err) + } + // A tie in value with different bits: the earlier block's signed zero + // must win, whichever sign each block holds. + const gn3 = 131073 // two blocks + for _, tc := range []struct { + first, second float64 + }{ + {0, math.Copysign(0, -1)}, // tie keeps the first block's +0 + {math.Copysign(0, -1), 0}, // tie keeps the first block's -0 + } { + vals := make([]float64, gn3) + for i := range vals { + vals[i] = tc.first + if i >= gn3/2 { + vals[i] = tc.second + } + } + zeros := mk(core.FromFloats(vals, gn3)) + want, err := core.Max(zeros) + if err != nil { + t.Fatal(err) + } + err = Launch(2, func(w *World) error { + span := mustPartition(t, gn3, w.Size(), w.Rank()) + local := narrowSliceFor(t, zeros, span) + got, err := w.AllReduceShards(local, span, Max) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("tie: sharded max %s against single-array %s", + scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatal(err) + } + } +} + +// TestShardedExtremaOverTCP runs the extrema's candidate flags over +// real connections, the pattern of TestShardedSumOverTCP. +func TestShardedExtremaOverTCP(t *testing.T) { + const gn = 131073 + pyr := make([]float64, gn) + for i := range pyr { + d := i - gn/2 + if d < 0 { + d = -d + } + pyr[i] = -float64(d) + } + peak := mk(core.FromFloats(pyr, gn)) + want, err := core.Max(peak) + if err != nil { + t.Fatal(err) + } + runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, peak, span) + got, err := w.AllReduceShards(local, span, Max) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("rank %d: sharded %s against single-array %s", + w.Rank(), scalarBits(got), scalarBits(want)) + } + return nil + }) +} + +// TestShardedEmptyAxisMatchesTheDtypeRules pins the degenerate length's +// dtype rules against the non-empty path's: a Bool sum counts its trues +// and answers zero, and Any and All keep refusing the non-Bool dtypes. +func TestShardedEmptyAxisMatchesTheDtypeRules(t *testing.T) { + emptyBool := mk(core.FromBools(nil, 0)) + if got := core.Sum(emptyBool); got.Int() != 0 { + t.Fatalf("the single-array Bool sum of an empty array: %v", got) + } + err := Launch(3, func(w *World) error { + span := mustPartition(t, 0, w.Size(), w.Rank()) + got, err := w.AllReduceShards(emptyBool, span, Sum) + if err != nil { + return err + } + if got.Int() != 0 { + t.Fatalf("the empty Bool sum answered %v", got) + } + if any, err := w.AllReduceShards(emptyBool, span, Any); err != nil || any.Int() != 0 { + t.Fatalf("the empty Bool Any answered %v %v", any, err) + } + if all, err := w.AllReduceShards(emptyBool, span, All); err != nil || all.Int() != 1 { + t.Fatalf("the empty Bool All answered %v %v", all, err) + } + emptyFloat := mk(core.FromFloats(nil, 0)) + if _, err := w.AllReduceShards(emptyFloat, span, Any); err == nil { + t.Fatal("the empty Any of a Float shard was accepted") + } + if _, err := w.AllReduceShards(emptyFloat, span, All); err == nil { + t.Fatal("the empty All of a Float shard was accepted") + } + return nil + }) + if err != nil { + t.Fatal(err) + } +} + +// prodFixture keeps every factor a relative hair away from one, so a +// long product stays finite and meaningful against the single-array +// answer. +func prodFixture(n int) []float64 { + v := make([]float64, n) + for i := range v { + v[i] = 1 + float64(int64(i%2001)-1000)/1e6 + } + return v +} + +// TestShardedProdIsTheSingleArrayProd: the sharded product carries the +// single-array product's exact bits, the dtype's own per-step rounding +// included. +func TestShardedProdIsTheSingleArrayProd(t *testing.T) { + for _, size := range []int{1, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537, 131073} { + build := func(t *testing.T) map[core.Dtype]*core.Array { + build1 := func(a *core.Array, err error) *core.Array { + if err != nil { + t.Fatal(err) + } + return a + } + vals := prodFixture(gn) + ints := make([]int64, gn) + for i := range ints { + ints[i] = int64(i%5) - 2 + } + halves := make([]uint16, gn) + f32s := make([]float32, gn) + for i := range halves { + halves[i] = core.HalfFromFloat64(vals[i]) + f32s[i] = float32(vals[i]) + } + return map[core.Dtype]*core.Array{ + core.Float: build1(core.FromFloats(vals, gn)), + core.Float32: build1(core.FromFloat32s(f32s, gn)), + core.Float16: build1(core.HalvesFromArray(halves, gn)), + core.Int: build1(core.FromInts(ints, gn)), + } + } + for dt, whole := range build(t) { + want, err := core.Prod(whole, 0, false) + if err != nil { + t.Fatal(err) + } + var wantBits string + switch dt { + case core.Int: + wantBits = scalarBits(core.IntScalar(want.RawInts()[0])) + case core.Float16: + wantBits = scalarBits(core.FloatScalar(core.HalfToFloat64(want.RawHalves()[0]))) + case core.Float32: + wantBits = scalarBits(core.FloatScalar(float64(want.RawFloat32s()[0]))) + default: + wantBits = scalarBits(core.FloatScalar(want.FloatAt(0))) + } + err = Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceShards(local, span, Prod) + if err != nil { + return err + } + if scalarBits(got) != wantBits { + t.Fatalf("size %d gn %d %s: sharded %s against single-array %s", + w.Size(), gn, dt, scalarBits(got), wantBits) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d %s: %v", size, gn, dt, err) + } + } + } + } +} + +// TestShardedNormIsTheSingleArrayNorm pins the power sums and their +// closing against the single-array norm, finite exponents only. +func TestShardedNormIsTheSingleArrayNorm(t *testing.T) { + for _, size := range []int{1, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537, 131073} { + for _, p := range []float64{1, 2, 3.5} { + whole := sparseFixtureOf(gn, core.Float) + want, err := core.Norm(whole, p, 0, false) + if err != nil { + t.Fatal(err) + } + err = Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + local := narrowSliceFor(t, whole, span) + got, err := w.AllReduceNormShards(local, span, p) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(core.FloatScalar(want.FloatAt(0))) { + t.Fatalf("size %d gn %d p %v: sharded %s against single-array %s", + w.Size(), gn, p, scalarBits(got), scalarBits(core.FloatScalar(want.FloatAt(0)))) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d p %v: %v", size, gn, p, err) + } + } + } + } +} + +// sparseFixtureOf is sparseFixture landing in one dtype. +func sparseFixtureOf(gn int, dt core.Dtype) *core.Array { + vals := sparseFixture(gn) + var a *core.Array + var err error + switch dt { + case core.Float32: + f32s := make([]float32, gn) + for i := range f32s { + f32s[i] = float32(vals[i]) + } + a, err = core.FromFloat32s(f32s, gn) + default: + a, err = core.FromFloats(vals, gn) + } + if err != nil { + panic(err) + } + return a +} + +// TestShardedDotIsTheSingleArrayDot: two equally sharded arrays answer +// the single-array Dot's exact bits. +func TestShardedDotIsTheSingleArrayDot(t *testing.T) { + for _, size := range []int{1, 3, 5, 8} { + for _, gn := range []int{1, 100, 65537, 131073} { + x := fixtureDtypes(t, gn)[core.Float] + yWhole := sparseFixtureOf(gn, core.Float) + want, err := core.Dot(x, yWhole) + if err != nil { + t.Fatal(err) + } + err = Launch(size, func(w *World) error { + span := mustPartition(t, gn, w.Size(), w.Rank()) + lx := narrowSliceFor(t, x, span) + ly := narrowSliceFor(t, yWhole, span) + got, err := w.AllReduceDotShards(lx, ly, span) + if err != nil { + return err + } + if scalarBits(got) != scalarBits(want) { + t.Fatalf("size %d gn %d: sharded %s against single-array %s", + w.Size(), gn, scalarBits(got), scalarBits(want)) + } + return nil + }) + if err != nil { + t.Fatalf("size %d gn %d: %v", size, gn, err) + } + } + } +} + +// TestShardedVectorRefusals: the infinity norm belongs to Max, and the +// vector reductions carry 1-D arrays alone. +func TestShardedVectorRefusals(t *testing.T) { + err := Launch(2, func(w *World) error { + span := mustPartition(t, 100, w.Size(), w.Rank()) + local := narrowSliceFor(t, sparseFixtureOf(100, core.Float), span) + if _, err := w.AllReduceNormShards(local, span, math.Inf(1)); err == nil { + t.Fatal("the infinity norm was accepted") + } + two, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) + if err != nil { + return err + } + span2 := mustPartition(t, 2, w.Size(), w.Rank()) + if _, err := w.AllReduceNormShards(two, span2, 2); err == nil { + t.Fatal("a two-dimensional shard was accepted") + } + if _, err := w.AllReduceShards(two, span2, Prod); err == nil { + t.Fatal("a two-dimensional prod shard was accepted") + } + return nil + }) + if err != nil { + t.Fatal(err) + } +} diff --git a/spmd/transport.go b/spmd/transport.go new file mode 100644 index 0000000..05305dd --- /dev/null +++ b/spmd/transport.go @@ -0,0 +1,237 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "bufio" + "bytes" + "encoding/binary" + "io" + "net" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// A frame is one message on a link: a 24-byte header, then the payload. +// +// offset 0: tag u8 +// offset 1: version u8 (frameVersion, a guard against a confused peer) +// offset 2: reserved u16 +// offset 4: from u32, little-endian, the sending rank +// offset 8: dest u32, little-endian, the receiving rank +// offset 12: length u64, little-endian, payload bytes that follow +// +// The header is the only place a length arrives from the wire, and no +// read ever allocates a payload longer than the world's message +// ceiling: a peer announcing more is an error before a byte of payload +// is read. From and dest name the endpoints because a networked world +// is a star over rank 0: a frame from rank 2 to rank 5 rides rank 2's +// connection in, and rank 5's connection out. +const ( + frameHeaderLen = 24 + frameVersion = 1 +) + +// The tags the collectives speak. The handshake speaks its own words on +// the fresh connection, before any frame. +const ( + tagBarrier uint8 = 1 + tagBarrierAck uint8 = 2 + tagBroadcast uint8 = 3 + tagScatterHead uint8 = 4 + tagScatter uint8 = 5 + tagGather uint8 = 6 + tagShardsValues uint8 = 8 + tagShardsWhole uint8 = 9 + tagReduce uint8 = 10 + tagHalo uint8 = 11 +) + +// message is one frame in the world's own terms. +type message struct { + tag uint8 + from int + dest int + data []byte + // pooled marks a payload whose buffer the hub's pump took from the + // routed frame pool, so the drain, the frame's single consumer, + // returns it there after the write. Every other frame leaves it + // false and its buffer belongs to whoever holds the frame. + pooled bool +} + +// peer is one rank's link. An in-process world joins the two ranks' +// channels directly: the sender writes into the receiver's inbox. A +// networked hub pumps each of its links with one reader and one writer; +// a networked rank that is not the hub drives its single connection +// itself and lets the hub route by dest. +type peer struct { + rank int + // inbox carries frames from this rank that this world reads. In + // process it is the direct channel from the peer; at the hub a + // reader feeds it; a non-hub rank leaves it nil and reads its one + // connection itself. + inbox chan message + // outbox carries frames to this rank's connection. In process it + // is the direct channel to the peer; at the hub a writer drains + // it; a non-hub rank leaves it nil and writes its one connection + // itself. + outbox chan message + // gone closes when the peer's world has left; only an in-process + // link has one, and only a sender waits on it. + gone chan struct{} + // conn is this world's end of the peer's networked link. + conn *tcpLink + // pending holds frames that arrived before the collective asked + // for their rank and tag, so no answer ever depends on the + // arrival order. Owned by the one goroutine that drives the + // world's collectives. + pending []message +} + +// tcpLink is a framed TCP connection to one peer rank. +type tcpLink struct { + conn net.Conn + rd *bufio.Reader + wr *bufio.Writer +} + +func newTCPLink(conn net.Conn) *tcpLink { + return &tcpLink{conn: conn, rd: bufio.NewReader(conn), wr: bufio.NewWriter(conn)} +} + +func (l *tcpLink) readFrame(max int64, deadline time.Time) (message, error) { + return l.readFrameInto(max, deadline, nil) +} + +// readFrameInto reads one frame as readFrame does. Take, when not +// nil, is asked with the parsed header and the payload length the +// header announced for the buffer the payload reads into; a nil +// answer allocates the payload as usual. A buffer take supplied +// leaves the read with the message's pooled mark on, which commits it +// to the single ownership chain the routed frame pool lives on: the +// pump that reads the frame hands it through one outbox channel to +// the one drain, which returns the buffer after the write. +func (l *tcpLink) readFrameInto(max int64, deadline time.Time, take func(message, int64) []byte) (message, error) { + if err := l.conn.SetReadDeadline(deadline); err != nil { + return message{}, err + } + var head [frameHeaderLen]byte + if _, err := io.ReadFull(l.rd, head[:]); err != nil { + return message{}, err + } + if head[1] != frameVersion { + return message{}, base.Errf("spmd: frame version %d from the link is not %d", head[1], frameVersion) + } + m := message{ + tag: head[0], + from: int(binary.LittleEndian.Uint32(head[4:])), + dest: int(binary.LittleEndian.Uint32(head[8:])), + } + length := int64(binary.LittleEndian.Uint64(head[12:])) + if length < 0 || length > max { + return message{}, base.Errf("spmd: the link announces a %d byte payload beyond the %d byte ceiling", length, max) + } + if take != nil { + if buf := take(m, length); buf != nil { + m.data, m.pooled = buf, true + } + } + if m.data == nil { + m.data = make([]byte, length) + } + if _, err := io.ReadFull(l.rd, m.data); err != nil { + return message{}, err + } + return m, nil +} + +func (l *tcpLink) writeFrame(m message, deadline time.Time) error { + if err := l.conn.SetWriteDeadline(deadline); err != nil { + return err + } + var head [frameHeaderLen]byte + head[0] = m.tag + head[1] = frameVersion + binary.LittleEndian.PutUint32(head[4:], uint32(m.from)) + binary.LittleEndian.PutUint32(head[8:], uint32(m.dest)) + binary.LittleEndian.PutUint64(head[12:], uint64(len(m.data))) + if _, err := l.wr.Write(head[:]); err != nil { + return err + } + if _, err := l.wr.Write(m.data); err != nil { + return err + } + return l.wr.Flush() +} + +// The handshake words on a fresh TCP connection: the joining rank +// sends hello, the listening rank answers with the world's size and +// the joining rank's place in it. +const handshakeLen = 16 + +var handshakeMagic = [4]byte{'T', 'S', 'P', 'M'} + +// sendHello is the joining side's word. +func sendHello(conn net.Conn) error { + var hello [handshakeLen]byte + copy(hello[0:4], handshakeMagic[:]) + hello[4] = frameVersion + _, err := conn.Write(hello[:]) + return err +} + +// readHello is the listening side's read of it. +func readHello(conn net.Conn, deadline time.Time) error { + if err := conn.SetReadDeadline(deadline); err != nil { + return err + } + var hello [handshakeLen]byte + if _, err := io.ReadFull(conn, hello[:]); err != nil { + return err + } + if !bytes.Equal(hello[0:4], handshakeMagic[:]) { + return base.Errf("spmd: the joining connection did not say the spmd magic") + } + if hello[4] != frameVersion { + return base.Errf("spmd: joining protocol version %d is not %d", hello[4], frameVersion) + } + return nil +} + +// sendWelcome is the listening side's answer: the world size and the +// rank the joining connection carries. +func sendWelcome(conn net.Conn, size, rank int) error { + var welcome [handshakeLen]byte + copy(welcome[0:4], handshakeMagic[:]) + welcome[4] = frameVersion + binary.LittleEndian.PutUint32(welcome[8:], uint32(size)) + binary.LittleEndian.PutUint32(welcome[12:], uint32(rank)) + _, err := conn.Write(welcome[:]) + return err +} + +// readWelcome is the joining side's read of it. +func readWelcome(conn net.Conn, deadline time.Time) (size, rank int, err error) { + if err := conn.SetReadDeadline(deadline); err != nil { + return 0, 0, err + } + var welcome [handshakeLen]byte + if _, err := io.ReadFull(conn, welcome[:]); err != nil { + return 0, 0, err + } + if !bytes.Equal(welcome[0:4], handshakeMagic[:]) { + return 0, 0, base.Errf("spmd: the listener did not answer with the spmd magic") + } + if welcome[4] != frameVersion { + return 0, 0, base.Errf("spmd: listener protocol version %d is not %d", welcome[4], frameVersion) + } + size = int(binary.LittleEndian.Uint32(welcome[8:])) + rank = int(binary.LittleEndian.Uint32(welcome[12:])) + if size < 1 || rank < 0 || rank >= size { + return 0, 0, base.Errf("spmd: the listener answered with an impossible size %d and rank %d", size, rank) + } + return size, rank, nil +} diff --git a/spmd/wire.go b/spmd/wire.go new file mode 100644 index 0000000..eee110d --- /dev/null +++ b/spmd/wire.go @@ -0,0 +1,315 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "encoding/binary" + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The wire form of an array is a self-describing frame payload: one +// dtype byte, one dimension-count byte, the extents as little-endian +// int64, then the elements as little-endian fixed-width raw bits in +// row-major order. Every multi-byte field is written little-endian by +// explicit conversion, never by copying the payload's memory, so the +// form is the same on every architecture the library builds for. +// Floating-point payloads keep their exact bit patterns, NaN payloads +// and signed zeros included. + +// dtypeWidth is the wire width of one element per dtype. +func dtypeWidth(dt core.Dtype) int { + switch dt { + case core.Int8, core.Uint8, core.Bool: + return 1 + case core.Int16, core.Uint16, core.Float16: + return 2 + case core.Int32, core.Uint32, core.Float32: + return 4 + case core.Int, core.Float: + return 8 + case core.Complex: + return 16 + } + return 0 +} + +// encodeArray appends the wire form of a to dst and returns the grown +// slice. +func encodeArray(dst []byte, a *core.Array) ([]byte, error) { + return encodePart(dst, a, a.Shape(), 0, a.Len()) +} + +// encodeHead appends just the head of a wire form: dtype, dimension +// count and extents. A rank that holds the head can name its piece of +// the data before any payload flows. +func encodeHead(dst []byte, dt core.Dtype, shape []int) ([]byte, error) { + if dtypeWidth(dt) == 0 { + return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt) + } + if len(shape) > 255 { + return nil, base.Errf("spmd: an array of %d dimensions exceeds the wire form's 255", len(shape)) + } + dst = append(dst, byte(dt), byte(len(shape))) + for _, d := range shape { + if d < 0 { + return nil, base.Errf("spmd: a wire form cannot name a negative extent %d", d) + } + dst = binary.LittleEndian.AppendUint64(dst, uint64(d)) + } + return dst, nil +} + +// decodeHead reads a wire form's head: the dtype and the shape. It +// validates only the head's own length; the payload against the shape +// is decodeArray's check. +func decodeHead(wire []byte) (core.Dtype, []int, error) { + if len(wire) < 2 { + return 0, nil, base.Errf("spmd: a wire form needs at least two bytes, got %d", len(wire)) + } + dt := core.Dtype(wire[0]) + if dtypeWidth(dt) == 0 { + return 0, nil, base.Errf("spmd: dtype %s cannot travel the wire", dt) + } + ndim := int(wire[1]) + if len(wire) < 2+8*ndim { + return 0, nil, base.Errf("spmd: a wire form for %d dimensions is short by %d bytes", ndim, 2+8*ndim-len(wire)) + } + shape := make([]int, ndim) + for i := range shape { + v := binary.LittleEndian.Uint64(wire[2+8*i:]) + if v > math.MaxInt { + return 0, nil, base.Errf("spmd: an extent of %d does not fit this machine's int", v) + } + shape[i] = int(v) + } + return dt, shape, nil +} + +// decodeWire builds an array from its complete wire form, head and +// payload together. Every extent is bounded before anything is +// allocated: an overflowing shape or a payload that does not carry +// exactly the named elements is an error, never a partial answer. +func decodeWire(wire []byte) (*core.Array, error) { + dt, shape, err := decodeHead(wire) + if err != nil { + return nil, err + } + return decodeArray(dt, shape, wire[2+8*len(shape):]) +} + +// partHeadLen is the byte length of a one-dimensional part's wire +// head, the shape encodePart writes for `[]int{count}`: one dtype +// byte, one dimension-count byte, and the single uint64 extent. +const partHeadLen = 2 + 8 + +// encodePart appends the wire form of a contiguous element range of a, +// presented under the given shape: the head names the shape, the +// payload carries elements [first, first+count) of a. Only the array's +// own elements take part: a rebased view's payload may run past its +// element count, so every payload walk stops at Len. +func encodePart(dst []byte, a *core.Array, shape []int, first, count int) ([]byte, error) { + dt := a.Dtype() + width := dtypeWidth(dt) + if width == 0 { + return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt) + } + if len(shape) > 255 { + return nil, base.Errf("spmd: an array of %d dimensions exceeds the wire form's 255", len(shape)) + } + if first < 0 || count < 0 || first+count > a.Len() { + return nil, base.Errf("spmd: element range [%d, %d) is outside the array's %d elements", first, first+count, a.Len()) + } + for _, d := range shape { + if d < 0 { + return nil, base.Errf("spmd: a wire form cannot name a negative extent %d", d) + } + } + dst = append(dst, byte(dt), byte(len(shape))) + for _, d := range shape { + dst = binary.LittleEndian.AppendUint64(dst, uint64(d)) + } + start := len(dst) + switch dt { + case core.Int: + for _, v := range a.RawInts()[first : first+count] { + dst = binary.LittleEndian.AppendUint64(dst, uint64(v)) + } + case core.Float: + for _, v := range a.RawFloats()[first : first+count] { + dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(v)) + } + case core.Float32: + for _, v := range a.RawFloat32s()[first : first+count] { + dst = binary.LittleEndian.AppendUint32(dst, math.Float32bits(v)) + } + case core.Float16: + for _, v := range a.RawHalves()[first : first+count] { + dst = binary.LittleEndian.AppendUint16(dst, v) + } + case core.Complex: + for _, v := range a.RawComplexes()[first : first+count] { + dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(real(v))) + dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(imag(v))) + } + case core.Bool: + for _, v := range a.RawBools()[first : first+count] { + b := byte(0) + if v { + b = 1 + } + dst = append(dst, b) + } + case core.Int8: + for _, v := range a.RawInt8s()[first : first+count] { + dst = append(dst, byte(v)) + } + case core.Uint8: + for _, v := range a.RawUint8s()[first : first+count] { + dst = append(dst, v) + } + case core.Int16: + for _, v := range a.RawInt16s()[first : first+count] { + dst = binary.LittleEndian.AppendUint16(dst, uint16(v)) + } + case core.Uint16: + for _, v := range a.RawUint16s()[first : first+count] { + dst = binary.LittleEndian.AppendUint16(dst, v) + } + case core.Int32: + for _, v := range a.RawInt32s()[first : first+count] { + dst = binary.LittleEndian.AppendUint32(dst, uint32(v)) + } + case core.Uint32: + for _, v := range a.RawUint32s()[first : first+count] { + dst = binary.LittleEndian.AppendUint32(dst, v) + } + } + if got := len(dst) - start; got != count*width { + return nil, base.Errf("spmd: array of dtype %s encoded %d bytes for %d elements", dt, got, count) + } + return dst, nil +} + +// decodeArray builds an array from the wire payload of a dtype and +// shape the caller has read. The payload must carry exactly the +// elements the shape names at the dtype's width; the caller has +// already bounded payload by the world's message ceiling. +func decodeArray(dt core.Dtype, shape []int, payload []byte) (*core.Array, error) { + width := dtypeWidth(dt) + if width == 0 { + return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt) + } + n, ok := elementCount(shape) + if !ok { + return nil, base.Errf("spmd: shape %v overflows the element count", shape) + } + if int64(len(payload)) != int64(n)*int64(width) { + return nil, base.Errf("spmd: wire payload of %d bytes does not carry %d elements of %s", len(payload), n, dt) + } + switch dt { + case core.Int: + vals := make([]int64, n) + for i := range vals { + vals[i] = int64(binary.LittleEndian.Uint64(payload[i*8:])) + } + return core.FromInts(vals, shape...) + case core.Float: + vals := make([]float64, n) + for i := range vals { + vals[i] = math.Float64frombits(binary.LittleEndian.Uint64(payload[i*8:])) + } + return core.FromFloats(vals, shape...) + case core.Float32: + vals := make([]float32, n) + for i := range vals { + vals[i] = math.Float32frombits(binary.LittleEndian.Uint32(payload[i*4:])) + } + return core.FromFloat32s(vals, shape...) + case core.Float16: + vals := make([]uint16, n) + for i := range vals { + vals[i] = binary.LittleEndian.Uint16(payload[i*2:]) + } + return core.HalvesFromArray(vals, shape...) + case core.Complex: + vals := make([]complex128, n) + for i := range vals { + re := math.Float64frombits(binary.LittleEndian.Uint64(payload[i*16:])) + im := math.Float64frombits(binary.LittleEndian.Uint64(payload[i*16+8:])) + vals[i] = complex(re, im) + } + return core.FromComplexes(vals, shape...) + case core.Bool: + vals := make([]bool, n) + for i := range vals { + switch payload[i] { + case 0: + case 1: + vals[i] = true + default: + return nil, base.Errf("spmd: bool wire byte %d at index %d is neither 0 nor 1", payload[i], i) + } + } + return core.FromBools(vals, shape...) + case core.Int8: + vals := make([]int8, n) + for i := range vals { + vals[i] = int8(payload[i]) + } + return core.FromInt8s(vals, shape...) + case core.Uint8: + vals := make([]uint8, n) + for i := range vals { + vals[i] = payload[i] + } + return core.FromUint8s(vals, shape...) + case core.Int16: + vals := make([]int16, n) + for i := range vals { + vals[i] = int16(binary.LittleEndian.Uint16(payload[i*2:])) + } + return core.FromInt16s(vals, shape...) + case core.Uint16: + vals := make([]uint16, n) + for i := range vals { + vals[i] = binary.LittleEndian.Uint16(payload[i*2:]) + } + return core.FromUint16s(vals, shape...) + case core.Int32: + vals := make([]int32, n) + for i := range vals { + vals[i] = int32(binary.LittleEndian.Uint32(payload[i*4:])) + } + return core.FromInt32s(vals, shape...) + case core.Uint32: + vals := make([]uint32, n) + for i := range vals { + vals[i] = binary.LittleEndian.Uint32(payload[i*4:]) + } + return core.FromUint32s(vals, shape...) + } + return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt) +} + +// elementCount is the product of the shape, reported as false when any +// extent is negative or the product overflows an int. +func elementCount(shape []int) (int, bool) { + n := 1 + for _, d := range shape { + if d < 0 { + return 0, false + } + if d == 0 { + return 0, true + } + if n > math.MaxInt/d { + return 0, false + } + n *= d + } + return n, true +} diff --git a/spmd/wire_test.go b/spmd/wire_test.go new file mode 100644 index 0000000..21ff908 --- /dev/null +++ b/spmd/wire_test.go @@ -0,0 +1,175 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "bytes" + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The wire form is judged by one rule: the array that comes out carries +// the exact bits of the array that went in, shape and dtype included. +func TestWireRoundTrip(t *testing.T) { + nanLow := math.Float64frombits(0x7ff8000000000001) // NaN, low payload bit set + nanHigh := math.Float64frombits(0xfff8deadbeef0000) + negZero := math.Copysign(0, -1) + cases := []*core.Array{ + mk(core.FromFloats([]float64{1, -2.5, 0, negZero, math.Inf(1), math.Inf(-1), nanLow, nanHigh, math.MaxFloat64, math.SmallestNonzeroFloat64}, 10)), + mk(core.FromFloat32s([]float32{1.5, -0.25, float32(negZero), float32(math.Inf(1)), float32(nanLow)}, 5)), + mk(core.FromComplexes([]complex128{1 + 2i, complex(negZero, nanHigh), complex(0, negZero)}, 3)), + mk(core.FromInts([]int64{math.MaxInt64, math.MinInt64, -1, 0, 42}, 5)), + mk(core.FromBools([]bool{true, false, true, true}, 4)), + mk(core.FromInt8s([]int8{math.MinInt8, math.MaxInt8, -1, 0}, 4)), + mk(core.FromUint8s([]uint8{0, 255, 128}, 3)), + mk(core.FromInt16s([]int16{math.MinInt16, math.MaxInt16, -1}, 3)), + mk(core.FromUint16s([]uint16{0, 65535, 32768}, 3)), + mk(core.FromInt32s([]int32{math.MinInt32, math.MaxInt32, -1}, 3)), + mk(core.FromUint32s([]uint32{0, 4294967295, 2147483648}, 3)), + // Float16 raw halves: denormals, infinities, NaN payloads, the + // whole span the payload can hold. + mk(core.HalvesFromArray([]uint16{0x0001, 0x03ff, 0x7bff, 0x7c00, 0xfc00, 0x7e00, 0x7eaa}, 7)), + // Empty and multi-dimensional shapes. + mk(core.FromFloats(nil, 0)), + mk(core.FromInts([]int64{1, 2, 3, 4, 5, 6}, 2, 3)), + mk(core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}, 2, 1, 2, 3)), + } + for i, a := range cases { + wire, err := encodeArray(nil, a) + if err != nil { + t.Fatalf("case %d (%s): %v", i, a.Dtype(), err) + } + b, err := decodeWire(wire) + if err != nil { + t.Fatalf("case %d (%s): %v", i, a.Dtype(), err) + } + if b.Dtype() != a.Dtype() { + t.Fatalf("case %d: dtype %s came back as %s", i, a.Dtype(), b.Dtype()) + } + if !sameShape(a.Shape(), b.Shape()) { + t.Fatalf("case %d: shape %v came back as %v", i, a.Shape(), b.Shape()) + } + if !sameBits(a, b) { + t.Fatalf("case %d (%s): the round trip moved bits", i, a.Dtype()) + } + } +} + +// TestWireRefusesTheHostile is the hostile-reader rule: a form that +// does not carry exactly what its head names is an error, never a short +// read and never a partial answer. +func TestWireRefusesTheHostile(t *testing.T) { + a := mk(core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)) + wire, err := encodeArray(nil, a) + if err != nil { + t.Fatal(err) + } + payload := wire[2+8*2:] + if _, err := decodeWire(wire[:len(wire)-1]); err == nil { + t.Fatal("a short payload decoded") + } + if _, err := decodeWire(append(bytes.Clone(wire), 0)); err == nil { + t.Fatal("a long payload decoded") + } + if _, err := decodeWire(wire[:1]); err == nil { + t.Fatal("a headless form decoded") + } + if _, err := decodeWire(wire[:6]); err == nil { + t.Fatal("a truncated head decoded") + } + big := []int{math.MaxInt32, math.MaxInt32} + if _, err := decodeArray(core.Float, big, nil); err == nil { + t.Fatal("an overflowing shape decoded") + } + if _, err := decodeArray(core.Float, []int{2, -2}, payload); err == nil { + t.Fatal("a negative extent decoded") + } + badBool := mk(core.FromBools([]bool{true, false}, 2)) + bw, err := encodeArray(nil, badBool) + if err != nil { + t.Fatal(err) + } + bp := bytes.Clone(bw) + bp[2+8] = 2 + if _, err := decodeWire(bp); err == nil { + t.Fatal("a bool byte of 2 decoded") + } +} + +// mk builds a fixture or panics; the fixtures are package-level +// literals, so a panic lands in the test that declared them. +func mk(a *core.Array, err error) *core.Array { + if err != nil { + panic(err) + } + return a +} + +// sameBits compares two arrays' raw payloads bit for bit, dtype by +// dtype. NaN payloads and signed zeros are part of the contract. +func sameBits(a, b *core.Array) bool { + if a.Len() != b.Len() || a.Dtype() != b.Dtype() { + return false + } + switch a.Dtype() { + case core.Int: + return slicesEqual(a.RawInts()[:a.Len()], b.RawInts()[:b.Len()]) + case core.Float: + x, y := a.RawFloats()[:a.Len()], b.RawFloats()[:b.Len()] + for i := range x { + if math.Float64bits(x[i]) != math.Float64bits(y[i]) { + return false + } + } + return true + case core.Float32: + x, y := a.RawFloat32s()[:a.Len()], b.RawFloat32s()[:b.Len()] + for i := range x { + if math.Float32bits(x[i]) != math.Float32bits(y[i]) { + return false + } + } + return true + case core.Float16: + return slicesEqual(a.RawHalves()[:a.Len()], b.RawHalves()[:b.Len()]) + case core.Complex: + x, y := a.RawComplexes()[:a.Len()], b.RawComplexes()[:b.Len()] + for i := range x { + if math.Float64bits(real(x[i])) != math.Float64bits(real(y[i])) || + math.Float64bits(imag(x[i])) != math.Float64bits(imag(y[i])) { + return false + } + } + return true + case core.Bool: + return slicesEqual(a.RawBools()[:a.Len()], b.RawBools()[:b.Len()]) + case core.Int8: + return slicesEqual(a.RawInt8s()[:a.Len()], b.RawInt8s()[:b.Len()]) + case core.Uint8: + return slicesEqual(a.RawUint8s()[:a.Len()], b.RawUint8s()[:b.Len()]) + case core.Int16: + return slicesEqual(a.RawInt16s()[:a.Len()], b.RawInt16s()[:b.Len()]) + case core.Uint16: + return slicesEqual(a.RawUint16s()[:a.Len()], b.RawUint16s()[:b.Len()]) + case core.Int32: + return slicesEqual(a.RawInt32s()[:a.Len()], b.RawInt32s()[:b.Len()]) + case core.Uint32: + return slicesEqual(a.RawUint32s()[:a.Len()], b.RawUint32s()[:b.Len()]) + } + return false +} + +func slicesEqual[T comparable](a, b []T) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} diff --git a/spmd/world.go b/spmd/world.go new file mode 100644 index 0000000..3ff9f61 --- /dev/null +++ b/spmd/world.go @@ -0,0 +1,576 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "errors" + "io" + "net" + "sync" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// Default bounds of a networked world: one collective's wait, one +// frame's payload, and the number of frames a link may queue. All +// three exist so that a stuck or hostile peer is an error, never a +// hang and never an allocation. +const ( + defaultTimeout = 10 * time.Minute + defaultMaxMessage = int64(16) << 30 + linkQueue = 64 + // pendingCap bounds how many unmatched frames one recvFrom may + // hold: a peer flooding frames nobody asked for is an error, + // never an unbounded allocator. + pendingCap = 4096 +) + +// Options bounds a networked world. An in-process world from Launch +// takes none of them: its links are channels, and progress is the +// program's own business, as it is in MPI. +type Options struct { + // Timeout bounds one collective's wait on the network: dialing, + // the handshake and every send and receive carry it as a + // deadline, refreshed each time a frame moves. Zero means the + // default of ten minutes; a negative value means no deadline at + // all. + Timeout time.Duration + // MaxMessage is the largest frame payload the world accepts, in + // bytes. A peer announcing more is refused before any allocation. + // Zero means the default of 16 GiB. + MaxMessage int64 +} + +func (o Options) timeout() time.Duration { + switch { + case o.Timeout > 0: + return o.Timeout + case o.Timeout < 0: + return 0 // no deadline + default: + return defaultTimeout + } +} + +func (o Options) maxMessage() int64 { + if o.MaxMessage > 0 { + return o.MaxMessage + } + return defaultMaxMessage +} + +// World is one rank's end of an SPMD world: its place in it, the links +// to the other ranks and the collectives. A World is driven by one +// goroutine: like an MPI rank, it never runs two collectives at once. +type World struct { + rank int + size int + timeout time.Duration + maxMessage int64 + networked bool + + peers []*peer // peers[r] is the link to rank r; nil for this rank + closers []io.Closer + + // The hub's readers and writers. The writers are waited on before + // an orderly close, so the last collective's frames are on the + // wire before the connections go away; the readers are waited on + // after, once those connections have broken their blocking reads. + drainWg sync.WaitGroup + pumpWg sync.WaitGroup + + done chan struct{} // closed when the world failed or left + closeDone sync.Once + failErr error +} + +// Rank returns this rank's index, from 0 to Size-1. +func (w *World) Rank() int { return w.rank } + +// Size returns the number of ranks in the world. +func (w *World) Size() int { return w.size } + +// Close ends a networked world: whatever the collectives queued is +// written to the wire, then the connections close and the other ranks +// see the departure as their next receive failing. On an in-process +// world it is a no-op, because Launch tears the world down. +func (w *World) Close() error { + w.leave() + w.drainWg.Wait() + var first error + for _, c := range w.closers { + if err := c.Close(); err != nil && first == nil { + first = err + } + } + w.pumpWg.Wait() + return first +} + +// leave closes done without recording a failure: the ordinary exit of +// the rank's program. +func (w *World) leave() { + w.closeDone.Do(func() { close(w.done) }) +} + +// fail records the world's first failure, closes done so that every +// wait on this world wakes, and closes the networked links. It returns +// the failure an outside caller should see. +func (w *World) fail(err error) error { + w.closeDone.Do(func() { + w.failErr = err + close(w.done) + for _, c := range w.closers { + c.Close() + } + }) + return w.status() +} + +// status is the entry check every public operation makes: a world that +// failed, or a rank whose program already returned, answers with an +// error and never with data. +func (w *World) status() error { + select { + case <-w.done: + if w.failErr != nil { + return base.Errf("spmd: rank %d world is in a failed state: %v", w.rank, w.failErr) + } + return base.Errf("spmd: rank %d world has left", w.rank) + default: + return nil + } +} + +// deadline is the absolute time one collective may run to on a +// networked world, refreshed each time a frame moves; the zero time +// means no deadline, which is what an in-process world always gets. +func (w *World) deadline() time.Time { + if !w.networked || w.timeout <= 0 { + return time.Time{} + } + return time.Now().Add(w.timeout) +} + +// depart closes the peer's gone channel: the peer closed its side of +// the link, so no further frame will ever come from it. +func (p *peer) depart() { + select { + case <-p.gone: + default: + close(p.gone) + } +} + +// hasLeft reports whether the peer closed its side of the link. +func (p *peer) hasLeft() bool { + if p.gone == nil { + return false + } + select { + case <-p.gone: + return true + default: + return false + } +} + +// sendTo carries one frame to rank r. On an in-process world it walks +// the direct channel; on a networked world it walks this rank's own +// link to the hub, because rank 0 routes every frame by its +// destination. +func (w *World) sendTo(r int, tag uint8, payload []byte) error { + if err := w.status(); err != nil { + return err + } + m := message{tag: tag, from: w.rank, dest: r, data: payload} + p := w.peers[r] + if p != nil && p.outbox != nil { + select { + case p.outbox <- m: + // On a buffered networked link a queued frame is not yet + // a delivered frame, so a peer that left between the two + // dropped it. An in-process handoff is the delivery + // itself: the receiver taking the frame and then leaving + // is its own healthy business. + if w.networked && p.hasLeft() { + return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r)) + } + return nil + case <-p.gone: + return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r)) + case <-w.done: + return w.status() + } + } + // A networked rank that is not the hub owns one connection, and + // every word it sends rides it; the hub reads the destination. + if err := w.peers[0].conn.writeFrame(m, w.deadline()); err != nil { + return w.fail(base.Errf("spmd: rank %d sending to rank %d: %v", w.rank, r, err)) + } + return nil +} + +// recvFrom returns the payload of the next frame rank r sent under the +// wanted tag, holding frames that arrived earlier for other ranks or +// tags until the collective asks for them: no answer ever depends on +// the arrival order. Any failure fails the world. +func (w *World) recvFrom(r int, tag uint8) ([]byte, error) { + if err := w.status(); err != nil { + return nil, err + } + p := w.peers[r] + if w.networked && w.rank != 0 { + // One stream carries every rank's words to a rank that is not + // the hub, so one pending stash serves them all. + p = w.peers[0] + } + for i, m := range p.pending { + if m.from == r && m.tag == tag { + p.pending = append(p.pending[:i], p.pending[i+1:]...) + return m.data, nil + } + } + for { + var m message + var err error + switch { + case p.inbox != nil: + select { + case m = <-p.inbox: + case <-p.gone: + // The peer left: whatever it queued before leaving is + // still in the buffer and still counts; only an empty + // buffer means the frames will never come. + for { + select { + case m = <-p.inbox: + if m.from == r && m.tag == tag { + return m.data, nil + } + p.pending = append(p.pending, m) + continue + default: + } + return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: rank %d left the world", w.rank, r, r)) + } + case <-w.done: + return nil, w.status() + } + default: + // This rank drives its single connection; the hub has + // already routed whatever was not for it. + m, err = p.conn.readFrame(w.maxMessage, w.deadline()) + if err != nil { + return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: %v", w.rank, r, w.readErr(err))) + } + if m.dest != w.rank { + return nil, w.fail(base.Errf("spmd: rank %d got a frame addressed to rank %d", w.rank, m.dest)) + } + } + if m.from == r && m.tag == tag { + return m.data, nil + } + p.pending = append(p.pending, m) + if len(p.pending) >= pendingCap { + return nil, w.fail(base.Errf("spmd: rank %d holds %d unmatched frames against rank %d", w.rank, len(p.pending), r)) + } + } +} + +// readErr names a networked read failure for what it is: a peer that +// closed or dropped its connection. +func (w *World) readErr(err error) error { + if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) { + return errors.New("the peer closed its connection") + } + return err +} + +// pump reads one networked hub link for the world's lifetime: frames +// for rank 0 join the link's inbox, frames for anybody else join that +// rank's outbox unchanged. The hub is a post office, never an +// interpreter: what a frame carries is the collectives' business. A +// frame bound for another rank reads into a buffer from the routed +// frame pool, whose ownership rides the message through the outbox +// channel to the link's one drain. A peer that closes its connection +// has left the world, which is its own business too; only a broken or +// unreadable link fails this world. +func (w *World) pump(p *peer) { + for { + m, err := p.conn.readFrameInto(w.maxMessage, w.deadline(), takeFrameBuffer) + if err != nil { + if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) || errors.Is(err, net.ErrClosed) { + p.depart() + return + } + w.fail(base.Errf("spmd: rank 0 reading from rank %d: %v", p.rank, w.readErr(err))) + return + } + if m.from != p.rank || m.dest < 0 || m.dest >= w.size { + w.fail(base.Errf("spmd: rank 0 got a mislabelled frame from rank %d", p.rank)) + return + } + var out chan message + if m.dest == 0 { + out = p.inbox + } else { + out = w.peers[m.dest].outbox + } + select { + case out <- m: + case <-w.done: + return + } + } +} + +// drain writes one networked hub link for the world's lifetime: it is +// the only goroutine that ever writes the connection, so the frames +// the collective logic and the routed traffic send share one ordered +// stream without a lock. It is also the single consumer of the +// routed frames the pump queues, which makes it the one place a +// pooled payload is returned: writeFrame is the payload's last +// reader, so the buffer goes back the moment the write returns, +// whatever the answer was. When the world ends it writes out whatever +// the last collectives queued before it leaves, so an orderly close +// never drops a frame that was sent. A write against a link whose peer +// already left is not failed, because the departure was the peer's own +// clean act; any other write error is, named at once rather than left +// to surface later as a hang. +func (w *World) drain(p *peer) { + write := func(m message) bool { + err := p.conn.writeFrame(m, w.deadline()) + if m.pooled { + routedFrames.retire(m.data) + } + if err == nil { + return true + } + if p.hasLeft() || w.ended() { + return false + } + w.fail(base.Errf("spmd: rank 0 writing to rank %d: %v", p.rank, err)) + return false + } + for { + select { + case m := <-p.outbox: + if !write(m) { + return + } + case <-w.done: + for { + select { + case m := <-p.outbox: + if !write(m) { + return + } + default: + return + } + } + } + } +} + +// ended reports whether the world's done channel has closed. +func (w *World) ended() bool { + select { + case <-w.done: + return true + default: + return false + } +} + +// startPumps launches the hub's readers and writers, one of each per +// link. They live as long as the world does. +func (w *World) startPumps() { + for r := 1; r < w.size; r++ { + p := w.peers[r] + w.pumpWg.Go(func() { w.pump(p) }) + w.drainWg.Go(func() { w.drain(p) }) + } +} + +// Launch runs the same function on size ranks of one process, one +// goroutine per rank, over in-process links: the same collectives, the +// same answers and the same rules as a networked world, which makes it +// the development and test surface of the package. The errors of the +// ranks that failed come back joined in rank order, so the report is +// deterministic and the rank that caused the trouble is in it; a rank +// whose function panics fails its world, and the panic comes back as +// that rank's error. +func Launch(size int, fn func(w *World) error) error { + if size < 1 { + return base.Errf("spmd: a world needs at least one rank, got %d", size) + } + worlds := make([]*World, size) + for r := range worlds { + worlds[r] = &World{ + rank: r, + size: size, + done: make(chan struct{}), + peers: make([]*peer, size), + } + } + // One channel per direction of every pair: the sender's outbox is + // the receiver's inbox, so a frame crosses without a middleman and + // a receive wakes the moment its rank's world leaves. + for r := range worlds { + for q := r + 1; q < size; q++ { + rToQ := make(chan message, linkQueue) + qToR := make(chan message, linkQueue) + worlds[r].peers[q] = &peer{rank: q, inbox: qToR, outbox: rToQ, gone: worlds[q].done} + worlds[q].peers[r] = &peer{rank: r, inbox: rToQ, outbox: qToR, gone: worlds[r].done} + } + } + errs := make([]error, size) + var wg sync.WaitGroup + for r := range worlds { + wg.Go(func() { + defer func() { + if p := recover(); p != nil { + errs[r] = base.Errf("spmd: rank %d panicked: %v", r, p) + worlds[r].fail(errs[r]) + } + worlds[r].leave() + }() + errs[r] = fn(worlds[r]) + }) + } + wg.Wait() + return errors.Join(errs...) +} + +// Listen builds the rank 0 end of a networked world: it listens on the +// address until every other rank has joined, assigning ranks in dial +// order. Connections that do not say the spmd handshake are closed +// and skipped, so they cannot take a rank's place. +func Listen(addr string, size int, opts Options) (*World, error) { + ln, err := net.Listen("tcp", addr) + if err != nil { + return nil, base.Errf("spmd: listening on %s: %v", addr, err) + } + return listen(ln, size, opts) +} + +// listen assembles the rank 0 world on a ready listener; the tests use +// it to hand over a listener whose address they already know. +func listen(ln net.Listener, size int, opts Options) (*World, error) { + if size < 1 { + ln.Close() + return nil, base.Errf("spmd: a world needs at least one rank, got %d", size) + } + timeout := opts.timeout() + w := &World{ + rank: 0, + size: size, + timeout: timeout, + maxMessage: opts.maxMessage(), + networked: true, + done: make(chan struct{}), + peers: make([]*peer, size), + } + w.closers = append(w.closers, ln) + if tcp, ok := ln.(*net.TCPListener); ok && timeout > 0 { + tcp.SetDeadline(time.Now().Add(timeout)) + } + deadline := w.deadline() + for joined := 1; joined < size; joined++ { + conn, err := ln.Accept() + if err != nil { + return nil, w.fail(base.Errf("spmd: rank 0 accepting rank %d on %s: %v", joined, ln.Addr(), err)) + } + if err := readHello(conn, deadline); err != nil { + // A connection that does not say the handshake is not one + // of ours; it takes no rank's place. + conn.Close() + joined-- + continue + } + if err := sendWelcome(conn, size, joined); err != nil { + conn.Close() + return nil, w.fail(base.Errf("spmd: rank 0 welcoming rank %d: %v", joined, err)) + } + w.peers[joined] = &peer{ + rank: joined, + conn: newTCPLink(conn), + inbox: make(chan message, linkQueue), + outbox: make(chan message, linkQueue), + gone: make(chan struct{}), + } + w.closers = append(w.closers, conn) + } + // Every rank has its link; nobody else joins this world. The + // listener was closers[0], and its job is done. + ln.Close() + w.closers = w.closers[1:] + w.startPumps() + return w, nil +} + +// Join builds the other ranks' end of a networked world: it dials the +// listening rank 0, which answers with the world's size and this +// connection's rank. The world is a star over rank 0's listener, so +// one address is the whole world's knowledge; rank 0 routes every +// frame to its destination. +func Join(addr string, opts Options) (*World, error) { + timeout := opts.timeout() + d := net.Dialer{Timeout: timeout} + conn, err := d.Dial("tcp", addr) + if err != nil { + return nil, base.Errf("spmd: rank dialling %s: %v", addr, err) + } + w := &World{ + timeout: timeout, + maxMessage: opts.maxMessage(), + networked: true, + done: make(chan struct{}), + closers: []io.Closer{conn}, + } + if err := sendHello(conn); err != nil { + return nil, w.fail(base.Errf("spmd: rank saying hello to %s: %v", addr, err)) + } + size, rank, err := readWelcome(conn, w.deadline()) + if err != nil { + return nil, w.fail(base.Errf("spmd: rank reading %s's welcome: %v", addr, err)) + } + if rank == 0 { + return nil, w.fail(base.Errf("spmd: the listener at %s answered a joining connection with rank 0", addr)) + } + w.rank, w.size = rank, size + w.peers = make([]*peer, size) + w.peers[0] = &peer{rank: 0, conn: newTCPLink(conn)} + return w, nil +} + +// Barrier blocks until every rank of the world has reached it. It +// carries no data and no arithmetic, so there is nothing in it to be +// anything but deterministic. +func (w *World) Barrier() error { + if err := w.status(); err != nil { + return err + } + if w.rank == 0 { + for r := 1; r < w.size; r++ { + if _, err := w.recvFrom(r, tagBarrier); err != nil { + return err + } + } + for r := 1; r < w.size; r++ { + if err := w.sendTo(r, tagBarrierAck, nil); err != nil { + return err + } + } + return nil + } + if err := w.sendTo(0, tagBarrier, nil); err != nil { + return err + } + _, err := w.recvFrom(0, tagBarrierAck) + return err +} diff --git a/spmd/world_test.go b/spmd/world_test.go new file mode 100644 index 0000000..47bd229 --- /dev/null +++ b/spmd/world_test.go @@ -0,0 +1,416 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package spmd + +import ( + "encoding/binary" + "errors" + "io" + "net" + "strings" + "sync" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +// The world tests cover both transports with the same battery: an +// in-process world over channels and a loopback TCP world over real +// connections, because the contract says they are the same machine. + +func TestLaunchOneRank(t *testing.T) { + err := Launch(1, func(w *World) error { + if w.Rank() != 0 || w.Size() != 1 { + t.Fatalf("rank %d of %d", w.Rank(), w.Size()) + } + return w.Barrier() + }) + if err != nil { + t.Fatal(err) + } +} + +func TestLaunchBarrier(t *testing.T) { + for _, size := range []int{2, 3, 5, 8} { + t.Run("", func(t *testing.T) { + reached := make([]int, size) + err := Launch(size, func(w *World) error { + if err := w.Barrier(); err != nil { + return err + } + reached[w.Rank()] = 1 + return w.Barrier() + }) + if err != nil { + t.Fatal(err) + } + for r, got := range reached { + if got != 1 { + t.Fatalf("rank %d never reported", r) + } + } + }) + } +} + +// TestLaunchFailsTogether is the fail-fast rule: the rank that sees the +// error fails its world, every other rank's next collective answers +// with an error, and Launch returns the first rank's error in rank +// order. +func TestLaunchFailsTogether(t *testing.T) { + boom := errors.New("boom") + saw := make([]error, 3) + err := Launch(3, func(w *World) error { + if w.Rank() == 1 { + return boom + } + saw[w.Rank()] = w.Barrier() + return saw[w.Rank()] + }) + if err == nil { + t.Fatal("a world with a failing rank returned nil") + } + for r, got := range saw { + if r == 1 { + continue + } + if got == nil { + t.Fatalf("rank %d's barrier survived a failed world", r) + } + } +} + +func TestLaunchPanicIsAnError(t *testing.T) { + err := Launch(2, func(w *World) error { + if w.Rank() == 1 { + panic("rank one fell over") + } + return w.Barrier() + }) + if err == nil || !strings.Contains(err.Error(), "rank 1 panicked") { + t.Fatalf("a panicking rank came back as %v", err) + } +} + +// runTCPWorld assembles a loopback TCP world of size ranks, running fn +// on every rank in its own goroutine, and fails the test if any rank +// errors. The listener the tests hand over lets rank 0 know its +// address before it starts. A final barrier runs after fn on every +// rank, the orderly end of the program: no rank tears its world down +// while another still expects words from it. +func runTCPWorld(t *testing.T, size int, opts Options, fn func(w *World) error) { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + var wg sync.WaitGroup + errs := make([]error, size) + wg.Go(func() { + w, err := listen(ln, size, opts) + if err != nil { + errs[0] = err + return + } + defer w.Close() + if err := fn(w); err != nil { + errs[0] = err + return + } + errs[0] = w.Barrier() + }) + for r := 1; r < size; r++ { + wg.Go(func() { + w, err := Join(ln.Addr().String(), opts) + if err != nil { + errs[r] = err + return + } + defer w.Close() + if err := fn(w); err != nil { + errs[r] = err + return + } + errs[r] = w.Barrier() + }) + } + wg.Wait() + for r, err := range errs { + if err != nil { + t.Fatalf("rank %d: %v", r, err) + } + } +} + +func TestTCPBarrier(t *testing.T) { + runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error { + if err := w.Barrier(); err != nil { + return err + } + return w.Barrier() + }) +} + +// TestTCPRankOrderIsDialOrder pins the one place ranks come from: the +// order the peers dial in, which the collectives' results never depend +// on. +func TestTCPRankOrderIsDialOrder(t *testing.T) { + runTCPWorld(t, 3, Options{Timeout: 30 * time.Second}, func(w *World) error { + return w.Barrier() + }) +} + +// TestTCPStrayConnectionTakesNoRank dials the listener with a +// connection that says nothing the handshake would recognise; the world +// must still assemble on the real ranks. +func TestTCPStrayConnectionTakesNoRank(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + // The stray speaks first, then goes silent. + stray, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + t.Fatal(err) + } + defer stray.Close() + if _, err := stray.Write([]byte("not the magic at all, but long enough to fill the read")); err != nil { + t.Fatal(err) + } + done := make(chan struct{}) + var worldErr error + go func() { + defer close(done) + w, err := listen(ln, 2, Options{Timeout: 30 * time.Second}) + if err != nil { + worldErr = err + return + } + w.Close() + }() + // The real rank joins behind the stray. + joined := make(chan error, 1) + go func() { + w, err := Join(ln.Addr().String(), Options{Timeout: 30 * time.Second}) + if err != nil { + joined <- err + return + } + w.Close() + joined <- nil + }() + select { + case err := <-joined: + if err != nil { + t.Fatalf("the real rank did not join behind the stray: %v", err) + } + case <-time.After(30 * time.Second): + t.Fatal("the world never assembled") + } + <-done + if worldErr != nil { + t.Fatalf("rank 0: %v", worldErr) + } +} + +// TestTCPDeadlineFailsTheCollective is the stuck-peer rule: a rank +// whose peer stops answering is errored out by its own deadline, never +// left hanging. The peer's world carries the short timeout from birth, +// so nothing changes under a running pump. +func TestTCPDeadlineFailsTheCollective(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + hubDone := make(chan *World, 1) + go func() { + // The hub assembles and stays silent: it never answers. + w, err := listen(ln, 2, Options{Timeout: 30 * time.Second}) + if err != nil { + w = nil + } + hubDone <- w + }() + peer, err := Join(ln.Addr().String(), Options{Timeout: 80 * time.Millisecond}) + if err != nil { + t.Fatal(err) + } + defer peer.Close() + hub := <-hubDone + if hub == nil { + t.Fatal("the hub did not assemble") + } + defer hub.Close() + if err := peer.Barrier(); err == nil { + t.Fatal("a barrier against a silent hub succeeded") + } +} + +// TestTCPMaxMessageRefused: a peer announcing a payload beyond the +// ceiling is an error before any allocation. +func TestTCPMaxMessageRefused(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + const ceiling = int64(1 << 20) + wCh := make(chan *World, 1) + go func() { + w, err := listen(ln, 2, Options{Timeout: 30 * time.Second, MaxMessage: ceiling}) + if err != nil { + wCh <- nil + return + } + wCh <- w + }() + // The fake peer does the handshake by hand, then announces an + // oversized frame. + peerDone := make(chan error, 1) + go func() { + conn, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + peerDone <- err + return + } + defer conn.Close() + if err := sendHello(conn); err != nil { + peerDone <- err + return + } + if _, _, err := readWelcome(conn, time.Now().Add(30*time.Second)); err != nil { + peerDone <- err + return + } + link := newTCPLink(conn) + var head [frameHeaderLen]byte + head[0] = tagBarrier + head[1] = frameVersion + binary.LittleEndian.PutUint32(head[4:], 1) // from + binary.LittleEndian.PutUint32(head[8:], 0) // dest + binary.LittleEndian.PutUint64(head[12:], uint64(ceiling+1)) + if _, err := link.wr.Write(head[:]); err != nil { + peerDone <- err + return + } + peerDone <- link.wr.Flush() + }() + w0 := <-wCh + if w0 == nil { + t.Fatal("rank 0 did not assemble") + } + defer w0.Close() + if err := <-peerDone; err != nil { + t.Fatalf("the fake peer: %v", err) + } + if _, err := w0.recvFrom(1, tagBarrier); err == nil { + t.Fatal("an oversized frame was received") + } +} + +// TestHandshake walks the joining words over an in-memory connection: +// the right magic passes both ways, a wrong magic is refused, and an +// impossible rank answer is refused. +func TestHandshake(t *testing.T) { + c, s := net.Pipe() + defer c.Close() + defer s.Close() + deadline := time.Now().Add(5 * time.Second) + go sendHello(c) + if err := readHello(s, deadline); err != nil { + t.Fatalf("the right magic was refused: %v", err) + } + go sendWelcome(c, 4, 2) + size, rank, err := readWelcome(s, deadline) + if err != nil { + t.Fatalf("the welcome did not read: %v", err) + } + if size != 4 || rank != 2 { + t.Fatalf("welcome answered size %d rank %d", size, rank) + } + var bad [handshakeLen]byte + copy(bad[0:4], []byte("XXXX")) + bad[4] = frameVersion + go c.Write(bad[:]) + if err := readHello(s, deadline); err == nil { + t.Fatal("a wrong magic was accepted") + } + go sendWelcome(c, 4, 4) + if _, _, err := readWelcome(s, deadline); err == nil { + t.Fatal("a rank beyond the world's size was accepted") + } +} + +// The compile-time guards on the error surface: every failure this +// package reports keeps the library's prefix. +func TestErrorPrefix(t *testing.T) { + err := base.Errf("spmd: test") + if err == nil || !strings.HasPrefix(err.Error(), "tensor: ") { + t.Fatalf("the package error lost its prefix: %v", err) + } + var _ io.Closer = (*World)(nil) +} + +// TestTCPNegativeLengthRefused: a frame length with its top bit set +// turns negative through the signed conversion; the receiver refuses +// it instead of allocating from it, which once panicked the hub. +func TestTCPNegativeLengthRefused(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + wCh := make(chan *World, 1) + go func() { + w, err := listen(ln, 2, Options{Timeout: 30 * time.Second}) + if err != nil { + w = nil + } + wCh <- w + }() + peerDone := make(chan error, 1) + go func() { + conn, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + peerDone <- err + return + } + defer conn.Close() + if err := sendHello(conn); err != nil { + peerDone <- err + return + } + if _, _, err := readWelcome(conn, time.Now().Add(30*time.Second)); err != nil { + peerDone <- err + return + } + link := newTCPLink(conn) + var head [frameHeaderLen]byte + head[0] = tagBarrier + head[1] = frameVersion + binary.LittleEndian.PutUint32(head[4:], 1) + binary.LittleEndian.PutUint32(head[8:], 0) + binary.LittleEndian.PutUint64(head[12:], uint64(1)<<63) + if _, err := link.wr.Write(head[:]); err != nil { + peerDone <- err + return + } + peerDone <- link.wr.Flush() + }() + w0 := <-wCh + if w0 == nil { + t.Fatal("rank 0 did not assemble") + } + defer w0.Close() + if err := <-peerDone; err != nil { + t.Fatalf("the fake peer: %v", err) + } + if _, err := w0.recvFrom(1, tagBarrier); err == nil { + t.Fatal("a negative length was received") + } +} diff --git a/stats/anova.go b/stats/anova.go new file mode 100644 index 0000000..d13a0a8 --- /dev/null +++ b/stats/anova.go @@ -0,0 +1,228 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "cmp" + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Group comparison: the parametric one-way analysis of variance and +// the rank-based Mann-Whitney U test. Both are classical frequentist +// tooling built on the distribution machinery of cdf.go, so a +// p-value never leaves the library's own incomplete beta and normal +// CDF. + +// ANOVAOneWay runs the classical one-way analysis of variance F test +// over the groups: the statistic is the between-group mean square +// divided by the within-group mean square, and the p-value is the +// upper tail of F with (k−1, N−k) degrees of freedom, evaluated +// through the regularised incomplete beta the same way WelchTTest's +// p-value is. A large F says the group means spread further than +// within-group noise explains. +// +// Every group must be a real, non-empty array of finite values, at +// least two groups are needed, and the observations must leave at +// least one within-group degree of freedom (N > k). Degenerate +// samples are refused rather than answered with a NaN: groups whose +// observations all share one identical value carry no within-group +// variance to divide by. +// +// Errors: fewer than two groups, an empty or complex group, a +// non-finite observation, N = k, or vanishing within-group variance. +func ANOVAOneWay(groups []*core.Array) (fStat, pValue float64, err error) { + if len(groups) < 2 { + return 0, 0, base.Errf("ANOVAOneWay: at least two groups are needed, got %d", len(groups)) + } + for i, g := range groups { + if g == nil { + return 0, 0, base.Errf("ANOVAOneWay: group %d is nil", i) + } + } + k := len(groups) + sizes := make([]int, k) + means := make([]float64, k) + total := 0 + sumAll := 0.0 + for i, g := range groups { + if g.Dtype() == core.Complex { + return 0, 0, base.Errf("ANOVAOneWay: group %d is complex", i) + } + s := 0.0 + // The walk is bounded by the group's own element count: a rebased + // view's payload may run past its visible elements, and those + // invisible tail slots are nobody's observations. + if fs := rawFloats(g); fs != nil { + for _, v := range fs[:g.Len()] { + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, 0, base.Errf("ANOVAOneWay: group %d holds the non-finite value %g", i, v) + } + s += v + } + } else { + for j := range g.Len() { + v := g.FloatAt(j) + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, 0, base.Errf("ANOVAOneWay: group %d holds the non-finite value %g", i, v) + } + s += v + } + } + sizes[i] = g.Len() + if sizes[i] == 0 { + return 0, 0, base.Errf("ANOVAOneWay: group %d is empty", i) + } + means[i] = s / float64(sizes[i]) + total += sizes[i] + sumAll += s + } + if total == k { + return 0, 0, base.Errf("ANOVAOneWay: one observation per group leaves no within-group degrees of freedom") + } + grand := sumAll / float64(total) + ssb, ssw := 0.0, 0.0 + for i, g := range groups { + d := means[i] - grand + ssb += float64(sizes[i]) * d * d + // Each group's within sum of squares folds through the same + // canonical partition the variance keeps, and the group totals + // then join in the groups' own order: the fold is a function of + // the group's length alone, so the split cannot move a bit. + if fs := rawFloats(g); fs != nil { + ssw += sqDeviations(fs[:g.Len()], means[i]) + } else { + ssw += sqDeviationsAt(g, means[i]) + } + } + if ssw == 0 { + if ssb == 0 { + return 0, 0, base.Errf("ANOVAOneWay: every observation shares one identical value, so the F statistic is undefined") + } + // Distinct means with no within-group noise: F grows without bound. + return math.Inf(1), 0, nil + } + fStat = (ssb / float64(k-1)) / (ssw / float64(total-k)) + // Two-sided upper tail of F(dfB, dfW) through the identity + // P(F > f) = I_{dfW/(dfW + dfB·f)}(dfW/2, dfB/2), the same + // closed form the Student-t tail uses. + x := float64(total-k) / (float64(total-k) + float64(k-1)*fStat) + pValue, err = BetaIncomplete(x, float64(total-k)/2, float64(k-1)/2) + if err != nil { + return 0, 0, base.Errf("ANOVAOneWay: %w", err) + } + return fStat, pValue, nil +} + +// MannWhitneyU runs the Mann-Whitney U test on two independent +// samples, the rank-based comparison of location that asks no +// normality: all observations are ranked together with midranks +// splitting ties, and u counts how strongly sample a outranks +// sample b through u = R_a − n(n+1)/2, with the companion statistic +// for sample b always n·m − u. The p-value is two-sided from the +// normal approximation with the continuity correction and the +// tie-corrected variance +// +// σ² = (n·m/12)·[(N+1) − Σ t(t²−1)/(N(N−1))], +// +// summing t over the tied blocks; with fewer than a dozen +// observations per sample the exact permutation distribution is +// noticeably discrete and the approximation only frames the answer. +// +// Errors: an empty or complex sample, a non-finite observation, or +// every observation in both samples tied, which leaves the U +// statistic no spread. +func MannWhitneyU(a, b *core.Array) (u, pValue float64, err error) { + if a.Len() == 0 || b.Len() == 0 { + return 0, 0, base.Errf("MannWhitneyU: both samples must be non-empty") + } + if a.Dtype() == core.Complex || b.Dtype() == core.Complex { + return 0, 0, base.Errf("MannWhitneyU: complex samples are not supported") + } + type tagged struct { + v float64 + inA bool + } + all := make([]tagged, 0, a.Len()+b.Len()) + // Both walks are bounded by the samples' own element counts: a + // rebased view's payload may run past its visible elements, and the + // invisible tail slots are nobody's observations. + if fs := rawFloats(a); fs != nil { + for _, v := range fs[:a.Len()] { + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, 0, base.Errf("MannWhitneyU: sample a holds the non-finite value %g", v) + } + all = append(all, tagged{v, true}) + } + } else { + for i := range a.Len() { + v := a.FloatAt(i) + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, 0, base.Errf("MannWhitneyU: sample a holds the non-finite value %g", v) + } + all = append(all, tagged{v, true}) + } + } + if fs := rawFloats(b); fs != nil { + for _, v := range fs[:b.Len()] { + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, 0, base.Errf("MannWhitneyU: sample b holds the non-finite value %g", v) + } + all = append(all, tagged{v, false}) + } + } else { + for i := range b.Len() { + v := b.FloatAt(i) + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, 0, base.Errf("MannWhitneyU: sample b holds the non-finite value %g", v) + } + all = append(all, tagged{v, false}) + } + } + slices.SortFunc(all, func(p, q tagged) int { return cmp.Compare(p.v, q.v) }) + // Walk the sorted pool one tie block at a time; each block shares + // the midrank of its 1-based rank span. + rankSumA := 0.0 + tieTerm := 0.0 + for i := 0; i < len(all); { + j := i + for j < len(all) && all[j].v == all[i].v { + j++ + } + mid := float64(i+j+1) / 2 + for k := i; k < j; k++ { + if all[k].inA { + rankSumA += mid + } + } + if t := j - i; t > 1 { + // Accumulate in float64 from the start: the integer form + // t·(t²−1) overflows int64 for tie blocks of 2^21 and up, + // and the wrapped term defeats the all-tied refusal below. + tieTerm += float64(t) * (float64(t)*float64(t) - 1) + } + i = j + } + na, nb := a.Len(), b.Len() + u = rankSumA - float64(na)*float64(na+1)/2 + total := float64(na + nb) + variance := float64(na) * float64(nb) / 12 * + (total + 1 - tieTerm/(total*(total-1))) + if variance <= 0 { + return 0, 0, base.Errf("MannWhitneyU: every observation is tied, so the U statistic has no spread") + } + // Continuity-corrected z from the upper side; an exact central U + // clamps to z = 0 rather than a small negative value. + z := (math.Abs(u-float64(na*nb)/2) - 0.5) / math.Sqrt(variance) + if z < 0 { + z = 0 + } + // The two-sided normal tail in one Erfc call on the magnitude, as + // the GLM fits compute it: 2·(1−Φ(z)) cancels to exactly zero once + // z passes about 8.3, where the true tail is still representable. + return u, math.Erfc(math.Abs(z) / math.Sqrt2), nil +} diff --git a/stats/anova_test.go b/stats/anova_test.go new file mode 100644 index 0000000..6cc08dc --- /dev/null +++ b/stats/anova_test.go @@ -0,0 +1,224 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestANOVAOneWayIdenticalMeans pins the null behaviour: two groups +// with the same mean give F = 0 exactly and p = 1 exactly, since the +// between-group sum of squares vanishes. +func TestANOVAOneWayIdenticalMeans(t *testing.T) { + groups := []*core.Array{ + mustFloats(t, []float64{1, 2, 3}), + mustFloats(t, []float64{3, 1, 2}), + } + f, p, err := ANOVAOneWay(groups) + if err != nil { + t.Fatalf("ANOVAOneWay: %v", err) + } + if f != 0 { + t.Errorf("F = %v, want 0", f) + } + if p != 1 { + t.Errorf("p = %v, want 1", p) + } +} + +// TestANOVAOneWayHandComputed pins a fully hand-computable case: +// groups {1,2,3} and {11,12,13} give SSB = 150, SSW = 4, F = 150 with +// (1, 4) degrees of freedom, and the upper tail has the closed form +// P(F > f) = 1 − 1.5·s + 0.5·s³ with s = √(1 − 2/77), which the test +// evaluates independently of the incomplete-beta continued fraction. +func TestANOVAOneWayHandComputed(t *testing.T) { + f, p, err := ANOVAOneWay([]*core.Array{ + mustFloats(t, []float64{1, 2, 3}), + mustFloats(t, []float64{11, 12, 13}), + }) + if err != nil { + t.Fatalf("ANOVAOneWay: %v", err) + } + if math.Abs(f-150) > 1e-9 { + t.Errorf("F = %v, want 150", f) + } + s := math.Sqrt(75.0 / 77.0) + want := 1 - 1.5*s + 0.5*s*s*s + if math.Abs(p-want) > 1e-9 { + t.Errorf("p = %.16g, want %.16g", p, want) + } +} + +// TestANOVAOneWayThreeGroups pins a three-group case whose p-value +// reduces to elementary arithmetic: F = 3 with (2, 6) degrees of +// freedom has upper tail I_{1/2}(3, 1) = (1/2)³ = 1/8 exactly. +func TestANOVAOneWayThreeGroups(t *testing.T) { + f, p, err := ANOVAOneWay([]*core.Array{ + mustFloats(t, []float64{1, 2, 3}), + mustFloats(t, []float64{2, 3, 4}), + mustFloats(t, []float64{3, 4, 5}), + }) + if err != nil { + t.Fatalf("ANOVAOneWay: %v", err) + } + if math.Abs(f-3) > 1e-9 { + t.Errorf("F = %v, want 3", f) + } + if math.Abs(p-0.125) > 1e-12 { + t.Errorf("p = %v, want 0.125", p) + } +} + +// TestANOVAOneWayDegenerate pins the zero-variance limits: distinct +// means with no noise send F to +Inf with p = 0, while one repeated +// value everywhere is refused instead of answered 0/0. +func TestANOVAOneWayDegenerate(t *testing.T) { + f, p, err := ANOVAOneWay([]*core.Array{ + mustFloats(t, []float64{1, 1}), + mustFloats(t, []float64{2, 2}), + }) + if err != nil { + t.Fatalf("ANOVAOneWay: %v", err) + } + if !math.IsInf(f, 1) || p != 0 { + t.Errorf("F = %v, p = %v, want +Inf and 0", f, p) + } + _, _, err = ANOVAOneWay([]*core.Array{ + mustFloats(t, []float64{5, 5}), + mustFloats(t, []float64{5, 5}), + }) + if err == nil || !strings.Contains(err.Error(), "identical value") { + t.Errorf("all-identical observations: error %v, want the degeneracy refusal", err) + } +} + +// TestANOVAOneWayRejects pins the input contracts. +func TestANOVAOneWayRejects(t *testing.T) { + ok := mustFloats(t, []float64{1, 2}) + cases := []struct { + name string + groups []*core.Array + want string + }{ + {"one group", []*core.Array{ok}, "at least two groups"}, + {"empty group", []*core.Array{ok, mustFloats(t, nil)}, "empty"}, + {"NaN observation", []*core.Array{ok, mustFloats(t, []float64{1, math.NaN()})}, "non-finite"}, + {"infinite observation", []*core.Array{ok, mustFloats(t, []float64{1, math.Inf(1)})}, "non-finite"}, + {"singletons", []*core.Array{mustFloats(t, []float64{1}), mustFloats(t, []float64{2})}, "no within-group degrees of freedom"}, + {"complex group", []*core.Array{ok, mustFromComplexes(t, []complex128{1, 2}, 2)}, "complex"}, + } + for _, c := range cases { + if _, _, err := ANOVAOneWay(c.groups); err == nil { + t.Errorf("%s: expected an error", c.name) + } else if !strings.Contains(err.Error(), c.want) { + t.Errorf("%s: error %q lacks %q", c.name, err, c.want) + } + } +} + +// TestMannWhitneyUSeparated pins the fully separated, tie-free case +// a = {1,2,3}, b = {4,5,6} against hand arithmetic: every rank goes +// to b, so u = 0; the tie-free variance is σ² = nm(N+1)/12 = 9·7/12 = +// 5.25, and the continuity-corrected z is 4/√5.25. +func TestMannWhitneyUSeparated(t *testing.T) { + u, p, err := MannWhitneyU(mustFloats(t, []float64{1, 2, 3}), mustFloats(t, []float64{4, 5, 6})) + if err != nil { + t.Fatalf("MannWhitneyU: %v", err) + } + if u != 0 { + t.Errorf("u = %v, want 0", u) + } + want := 2 * (1 - NormalCDF(4/math.Sqrt(5.25))) + if math.Abs(p-want) > 1e-12 { + t.Errorf("p = %.16g, want %.16g", p, want) + } + // Independently computed reference (an independent 30-digit calculation). + if math.Abs(p-0.080855598370052291) > 1e-15 { + t.Errorf("p = %.16g, want the tabulated 0.080855598370052291", p) + } +} + +// TestMannWhitneyUTies pins the tie-corrected variance on a sample +// pair with two tied blocks of three: u = 3, σ² = 76/7 by hand, and +// the continuity-corrected z is 4.5/√(76/7). +func TestMannWhitneyUTies(t *testing.T) { + u, p, err := MannWhitneyU(mustFloats(t, []float64{1, 2, 2, 3}), mustFloats(t, []float64{2, 3, 3, 4})) + if err != nil { + t.Fatalf("MannWhitneyU: %v", err) + } + if u != 3 { + t.Errorf("u = %v, want 3", u) + } + want := 2 * (1 - NormalCDF(4.5/math.Sqrt(76.0/7.0))) + if math.Abs(p-want) > 1e-12 { + t.Errorf("p = %.16g, want %.16g", p, want) + } + // Independently computed reference (an independent 30-digit calculation). + if math.Abs(p-0.17203370892182298) > 1e-15 { + t.Errorf("p = %.16g, want the tabulated 0.17203370892182298", p) + } +} + +// TestMannWhitneyUSymmetry checks the companion statistic and the +// shared p-value: swapping the samples gives n·m − u and the same p. +func TestMannWhitneyUSymmetry(t *testing.T) { + a := mustFloats(t, []float64{1, 2, 2, 3}) + b := mustFloats(t, []float64{2, 3, 3, 4}) + u1, p1, err := MannWhitneyU(a, b) + if err != nil { + t.Fatalf("MannWhitneyU(a, b): %v", err) + } + u2, p2, err := MannWhitneyU(b, a) + if err != nil { + t.Fatalf("MannWhitneyU(b, a): %v", err) + } + if u1+u2 != 16 { + t.Errorf("u + u' = %v + %v, want 16", u1, u2) + } + if p1 != p2 { + t.Errorf("p-values differ: %v vs %v", p1, p2) + } +} + +// TestMannWhitneyUCentral pins the exact-centre behaviour: identical +// samples put u at n·m/2, where the clamped z = 0 gives p = 1. +func TestMannWhitneyUCentral(t *testing.T) { + u, p, err := MannWhitneyU(mustFloats(t, []float64{1, 2, 3, 4}), mustFloats(t, []float64{1, 2, 3, 4})) + if err != nil { + t.Fatalf("MannWhitneyU: %v", err) + } + if u != 8 { + t.Errorf("u = %v, want 8", u) + } + if p != 1 { + t.Errorf("p = %v, want 1", p) + } +} + +// TestMannWhitneyURejects pins the input contracts. +func TestMannWhitneyURejects(t *testing.T) { + ok := mustFloats(t, []float64{1, 2}) + cases := []struct { + name string + a, b *core.Array + want string + }{ + {"empty a", mustFloats(t, nil), ok, "non-empty"}, + {"empty b", ok, mustFloats(t, nil), "non-empty"}, + {"complex", ok, mustFromComplexes(t, []complex128{1, 2}, 2), "complex"}, + {"NaN", ok, mustFloats(t, []float64{1, math.NaN()}), "non-finite"}, + {"all tied", mustFloats(t, []float64{1, 1}), mustFloats(t, []float64{1, 1}), "no spread"}, + } + for _, c := range cases { + if _, _, err := MannWhitneyU(c.a, c.b); err == nil { + t.Errorf("%s: expected an error", c.name) + } else if !strings.Contains(err.Error(), c.want) { + t.Errorf("%s: error %q lacks %q", c.name, err, c.want) + } + } +} diff --git a/stats/bench_kernels_test.go b/stats/bench_kernels_test.go new file mode 100644 index 0000000..e66dfe7 --- /dev/null +++ b/stats/bench_kernels_test.go @@ -0,0 +1,216 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "fmt" + "math" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the estimation kernels: the mixture expectation +// maximisation sweep, the kernel-density sweep, the windowed extrema +// and the histogram passes. Every input is a fixed arithmetic +// progression, so two runs of a benchmark measure the same work. + +// benchKernelsSeries builds a deterministic vector of length n: a few +// incommensurate frequencies plus a modular jitter, so the sorts and the +// window folds meet a spread without a generator. The phase separates +// two series of the same length. +func benchKernelsSeries(n int, phase float64) []float64 { + vals := make([]float64, n) + for i := range vals { + x := float64(i) + vals[i] = math.Sin(0.0017*x+phase)*4 + math.Cos(0.071*x+phase)*2 + float64((i*37)%101)/101 + } + return vals +} + +// benchKernelsVector wraps the series as a rank-1 array. +func benchKernelsVector(b *testing.B, n int, phase float64) *core.Array { + b.Helper() + a, err := core.FromFloats(benchKernelsSeries(n, phase), n) + if err != nil { + b.Fatal(err) + } + return a +} + +// benchKernelsCloud builds an n-by-d sample of k well-separated clusters +// whose jitter follows a fixed congruence: the same rows on every run, +// and a cloud the mixture's covariances can factor. +func benchKernelsCloud(n, d, k int) []float64 { + vals := make([]float64, n*d) + for i := range vals { + row, col := i/d, i%d + jitter := float64((row*2654435761+col*40503)%1000)/1000 - 0.5 + vals[i] = float64(row%k)*4 + jitter*1.5 + } + return vals +} + +// BenchmarkKernelsGaussianMixture measures one full fit: the k-means++ +// seeding, the Lloyd sweeps and the expectation maximisation sweeps. +func BenchmarkKernelsGaussianMixture(b *testing.B) { + const n, d, k = 3000, 5, 5 + data := benchKernelsCloud(n, d, k) + x, err := core.FromFloats(data, n, d) + if err != nil { + b.Fatal(err) + } + g := core.NewGenerator(7) + b.ReportAllocs() + for b.Loop() { + if _, err := GaussianMixture(g, x, k); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelsGMMSweep measures the expectation maximisation sweeps +// alone: the seeding runs once outside the timed loop, and every +// iteration re-copies the seeded parameters, which gmmEM updates in +// place. +func BenchmarkKernelsGMMSweep(b *testing.B) { + const n, d, k = 3000, 5, 5 + data := benchKernelsCloud(n, d, k) + weights, means, covs, err := gmmSeed(core.NewGenerator(7), data, n, d, k) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + w := slices.Clone(weights) + m := make([][]float64, k) + c := make([][]float64, k) + for j := range k { + m[j] = slices.Clone(means[j]) + c[j] = slices.Clone(covs[j]) + } + res, err := gmmEM(data, n, d, k, w, m, c) + if err != nil { + b.Fatal(err) + } + b.ReportMetric(float64(res.Iterations), "sweeps") + } +} + +// BenchmarkKernelsKernelDensity measures the O(n·points) kernel-density +// sweep with a fixed bandwidth, so no Silverman pre-pass is timed. +func BenchmarkKernelsKernelDensity(b *testing.B) { + sample := benchKernelsVector(b, 2048, 0) + points := benchKernelsVector(b, 512, 0.5) + b.ReportAllocs() + for b.Loop() { + if _, err := KernelDensity(sample, 0.35, points); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelsRollingMaxWide measures the windowed maximum at half +// the series length, the shape that separates a rescan from a deque. +func BenchmarkKernelsRollingMaxWide(b *testing.B) { + a := benchKernelsVector(b, 16384, 0) + b.ReportAllocs() + for b.Loop() { + if _, err := RollingMax(a, 8192); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelsRollingMinWide is the minimum's counterpart. +func BenchmarkKernelsRollingMinWide(b *testing.B) { + a := benchKernelsVector(b, 16384, 0) + b.ReportAllocs() + for b.Loop() { + if _, err := RollingMin(a, 8192); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelsHistogram measures the fused scan and the counted +// bins over a long sample. +func BenchmarkKernelsHistogram(b *testing.B) { + a := benchKernelsVector(b, 262144, 0) + b.ReportAllocs() + for b.Loop() { + if _, _, err := Histogram(a, 256); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelsHistogram2D measures the paired binning over a long +// sample on both axes. +func BenchmarkKernelsHistogram2D(b *testing.B) { + const n = 65536 + x := benchKernelsVector(b, n, 0) + y := benchKernelsVector(b, n, 0.5) + b.ReportAllocs() + for b.Loop() { + if _, _, _, err := Histogram2D(x, y, 64, 64); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelsGaussianProcessFit measures one fit over a design the +// size a Gaussian-process regression usually runs on: the Gram matrix, +// its factorisation, the weights and the posterior at the test points. +func BenchmarkKernelsGaussianProcessFit(b *testing.B) { + const n, m = 150, 60 + train := make([]float64, n) + test := make([]float64, m) + for i := range train { + train[i] = float64(i) / float64(n-1) + } + for i := range test { + test[i] = float64(i) / float64(m-1) + } + trainX, err := core.FromFloats(train, n, 1) + if err != nil { + b.Fatal(err) + } + trainY, err := core.FromFloats(benchKernelsSeries(n, 0), n) + if err != nil { + b.Fatal(err) + } + testX, err := core.FromFloats(test, m, 1) + if err != nil { + b.Fatal(err) + } + kernel, err := SquaredExponentialKernel(0.2) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := GaussianProcessRegression(kernel, trainX, trainY, 0.05, testX); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKernelsRollingMaxScaling reports the monotonic deque against +// the series length at a window of half the series: the linear growth +// the deque replaced the window rescan with. +func BenchmarkKernelsRollingMaxScaling(b *testing.B) { + for _, n := range []int{4096, 16384, 65536} { + a := benchKernelsVector(b, n, 0) + b.Run(fmt.Sprintf("n=%d/w=%d", n, n/2), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := RollingMax(a, n/2); err != nil { + b.Fatal(err) + } + } + }) + } +} diff --git a/stats/bench_quantilereg_test.go b/stats/bench_quantilereg_test.go new file mode 100644 index 0000000..a994139 --- /dev/null +++ b/stats/bench_quantilereg_test.go @@ -0,0 +1,38 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// BenchmarkQuantileRegressionFit measures the interior-point quantile +// fit: the barrier Newton systems, the fraction-to-the-boundary walk +// and the primal recovery, at a size where the per-iteration system +// build is a visible share of the run. +func BenchmarkQuantileRegressionFit(b *testing.B) { + const n, p = 400, 8 + x, vals := regDesign(b, n, p, 0x5eed0041) + src := &randSource{state: 0x5eed0042} + yv := make([]float64, n) + for i := range n { + eta := 0.5 + for c := range p { + eta += 0.25 * vals[i*p+c] + } + yv[i] = eta + (src.next()-0.5)*0.5 + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := QuantileRegression(x, y, 0.5); err != nil { + b.Fatal(err) + } + } +} diff --git a/stats/bench_regression_test.go b/stats/bench_regression_test.go new file mode 100644 index 0000000..b28ff93 --- /dev/null +++ b/stats/bench_regression_test.go @@ -0,0 +1,638 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "fmt" + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The randomised and cross-checked tests that guard the pair counting of +// KendallTau, the covariance diagonals of the regression entries and the +// pairwise-slope walk of Theil-Sen against the implementations they +// replace, followed by the benchmarks of those paths. + +// randSource is a deterministic linear congruential generator for the +// randomised checks: a fixed recurrence makes every failing case +// reproducible from its seed alone. +type randSource struct{ state uint64 } + +// next returns the next draw in [0, 1). +func (r *randSource) next() float64 { + r.state = r.state*6364136223846793005 + 1442695040888963407 + return float64(r.state>>11) / float64(uint64(1)<<53) +} + +// intn returns the next draw in [0, n). +func (r *randSource) intn(n int) int { return int(r.next() * float64(n)) } + +// bruteForceTally counts the pairs of two paired samples the way the +// definition reads: every ordered pair of positions, its two differences +// multiplied, the sign of the product deciding the side, and the ties of +// a sample counted over the same pair walk. It is the oracle the +// merge-sort pair counting is checked against. +func bruteForceTally(xs, ys []float64) pairTally { + var tally pairTally + for i := 1; i < len(xs); i++ { + for j := 0; j < i; j++ { + switch prod := (xs[i] - xs[j]) * (ys[i] - ys[j]); { + case prod > 0: + tally.concordant++ + case prod < 0: + tally.discordant++ + } + if xs[i] == xs[j] { + tally.tiedX++ + } + if ys[i] == ys[j] { + tally.tiedY++ + } + } + } + return tally +} + +// bruteForceTau is the O(n²) τ-b of the definition, the oracle the +// returned correlation is checked against. +func bruteForceTau(xs, ys []float64) float64 { + tally := bruteForceTally(xs, ys) + n := float64(len(xs)) + n0 := n * (n - 1) / 2 + return (float64(tally.concordant) - float64(tally.discordant)) / + math.Sqrt((n0-float64(tally.tiedX))*(n0-float64(tally.tiedY))) +} + +// TestKendallTauCountsMatchBruteForce checks the merge-sort pair +// counting against the pair enumeration it replaces. The values are +// drawn from a handful of levels so that ties in either sample are the +// common case, which is where the block scan and the tie accounting +// carry the work; continuous draws, degenerate ranges and a constant +// second sample cover the rest. Both sides are exact integers, so the +// comparison is on the returned bits rather than on a tolerance. +func TestKendallTauCountsMatchBruteForce(t *testing.T) { + src := &randSource{state: 0x9e3779b97f4a7c15} + levels := []float64{-1, 0, 1, 2, 3} + for trial := range 400 { + n := 2 + src.intn(38) + xs := make([]float64, n) + ys := make([]float64, n) + // A third of the draws are continuous, the rest come from the + // levels, so both samples are tie-heavy on average. + draw := func() float64 { + if src.intn(3) == 0 { + return src.next()*4 - 2 + } + return levels[src.intn(len(levels))] + } + for i := range n { + xs[i] = draw() + ys[i] = draw() + } + if trial%4 == 0 { + // The pairing is perfectly monotone, its τ exactly 1. + ys = append([]float64(nil), xs...) + } + if trial%9 == 0 { + // A constant chunk inside otherwise varying samples: the + // widest blocks of equal values the scan has to handle. + for i := n / 3; i < 2*n/3; i++ { + xs[i] = xs[0] + } + } + want := bruteForceTally(xs, ys) + got := tallyPairs(append([]float64(nil), xs...), append([]float64(nil), ys...)) + if got != want { + t.Fatalf("trial %d, n %d: pair counts %+v, want %+v\nx = %v\ny = %v", trial, n, got, want, xs, ys) + } + if want.tiedX == int64(n*(n-1)/2) || want.tiedY == int64(n*(n-1)/2) { + continue // a constant sample has no τ-b to compare + } + gotTau, err := KendallTau(mustFloats(t, xs), mustFloats(t, ys)) + if err != nil { + t.Fatalf("trial %d, n %d: KendallTau: %v", trial, n, err) + } + wantTau := bruteForceTau(xs, ys) + if gotTau != wantTau { + t.Fatalf("trial %d, n %d: τ = %v (%#x), want %v (%#x)\nx = %v\ny = %v", + trial, n, gotTau, math.Float64bits(gotTau), wantTau, math.Float64bits(wantTau), xs, ys) + } + } +} + +// invertDiagonal returns the diagonal of the inverse of a dense matrix +// by Gauss-Jordan elimination with partial pivoting: an independent +// route to the covariance diagonals the fits read through the shared LU +// solve. It is accurate to rounding, which is all the cross-checks below +// need from it. +func invertDiagonal(m [][]float64) []float64 { + p := len(m) + aug := make([][]float64, p) + for i := range p { + aug[i] = make([]float64, 2*p) + copy(aug[i], m[i]) + aug[i][p+i] = 1 + } + for k := range p { + piv := k + for i := k + 1; i < p; i++ { + if math.Abs(aug[i][k]) > math.Abs(aug[piv][k]) { + piv = i + } + } + aug[k], aug[piv] = aug[piv], aug[k] + d := aug[k][k] + for j := range 2 * p { + aug[k][j] /= d + } + for i := range p { + if i == k { + continue + } + f := aug[i][k] + if f == 0 { + continue + } + for j := range 2 * p { + aug[i][j] -= f * aug[k][j] + } + } + } + diag := make([]float64, p) + for i := range p { + diag[i] = aug[i][p+i] + } + return diag +} + +// copyMatrix copies a p×p system, the shape SolveSystem factors in place. +func copyMatrix(m [][]float64) [][]float64 { + out := make([][]float64, len(m)) + for i, row := range m { + out[i] = append([]float64(nil), row...) + } + return out +} + +// TestSolveSystemUnitColumnsMatchPerColumnSolves checks the shared +// factorisation the covariance diagonals read against the +// per-coefficient solves it replaced: the same normal-equations matrix, +// once against all p unit columns and once against each column alone, on +// a positive-definite system built from a deterministic design. The +// diagonal entries are compared bit for bit. +func TestSolveSystemUnitColumnsMatchPerColumnSolves(t *testing.T) { + const p = 9 + src := &randSource{state: 0xda3e39cb94b95bdb} + m := make([][]float64, p) + for i := range p { + m[i] = make([]float64, p) + } + row := make([]float64, p) + for range 40 { + row[0] = 1 + for c := 1; c < p; c++ { + row[c] = src.next() - 0.5 + } + for i := range p { + for j := range p { + m[i][j] += row[i] * row[j] + } + } + } + units := make([][]float64, p) + for j := range p { + units[j] = make([]float64, p) + units[j][j] = 1 + } + together, err := base.SolveSystem("test", copyMatrix(m), units) + if err != nil { + t.Fatal(err) + } + for j := range p { + e := make([]float64, p) + e[j] = 1 + alone, err := base.SolveSystem("test", copyMatrix(m), [][]float64{e}) + if err != nil { + t.Fatal(err) + } + if got, want := together[j][j], alone[0][j]; got != want { + t.Fatalf("diagonal %d of the batched solve = %v (%#x), want %v (%#x) from the single-column solve", + j, got, math.Float64bits(got), want, math.Float64bits(want)) + } + } +} + +// TestLogisticRegressionWaldStandardErrors rebuilds the Fisher +// information at the reported fit from the reported probabilities and +// inverts it by an independent elimination: the Wald standard errors are +// that inverse's diagonal, so the check covers the shared factorisation +// the entry now reads them through on a design no digest pins. +func TestLogisticRegressionWaldStandardErrors(t *testing.T) { + const n, p = 300, 6 + src := &randSource{state: 0x2545f4914f6cdd1d} + design := make([]float64, n*p) + response := make([]float64, n) + for r := range n { + eta := 0.3 + for c := range p { + v := src.next() - 0.5 + design[r*p+c] = v + eta += 0.8 * v + } + design[r*p] = 1 + if src.next() < 1/(1+math.Exp(-eta)) { + response[r] = 1 + } + } + res, err := LogisticRegression(mustFloats(t, design, n, p), mustFloats(t, response, n)) + if err != nil { + t.Fatal(err) + } + fisher := make([][]float64, p) + for i := range p { + fisher[i] = make([]float64, p) + } + for r := range n { + w := res.Fitted[r] * (1 - res.Fitted[r]) + for i := range p { + for j := range p { + fisher[i][j] += design[r*p+i] * design[r*p+j] * w + } + } + } + diag := invertDiagonal(fisher) + for j := range p { + want := math.Sqrt(diag[j]) + if rel := math.Abs(res.StandardErrors[j]-want) / want; rel > 1e-9 { + t.Fatalf("standard error %d = %.15g, want %.15g (relative %.3g)", + j, res.StandardErrors[j], want, rel) + } + } +} + +// TestHuberRegressionStandardErrors does the same for the Huber fit: +// σ²·(XᵀWX)⁻¹ with σ the reported robust scale and W the reported final +// weights, its diagonal checked against an independent elimination of +// the weighted normal equations the fit writes its own weights into. +func TestHuberRegressionStandardErrors(t *testing.T) { + const n, p = 300, 6 + src := &randSource{state: 0x853c49e6748fea9b} + design := make([]float64, n*p) + response := make([]float64, n) + for r := range n { + design[r*p] = 1 + response[r] = 1 + for c := 1; c < p; c++ { + v := src.next() - 0.5 + design[r*p+c] = v + response[r] += float64(c) * 0.5 * v + } + response[r] += 0.2 * (src.next() - 0.5) + } + response[7] += 5 // the gross point the robust fit exists for + res, err := HuberRegression(mustFloats(t, design, n, p), mustFloats(t, response, n)) + if err != nil { + t.Fatal(err) + } + if !(res.Scale > 0) { + t.Fatalf("the robust scale collapsed to %g on a contaminated sample", res.Scale) + } + wxx := make([][]float64, p) + for i := range p { + wxx[i] = make([]float64, p) + } + for r := range n { + w := res.Weights[r] + for i := range p { + for j := range p { + wxx[i][j] += design[r*p+i] * design[r*p+j] * w + } + } + } + diag := invertDiagonal(wxx) + for j := range p { + want := math.Sqrt(res.Scale * res.Scale * diag[j]) + if rel := math.Abs(res.StandardErrors[j]-want) / want; rel > 1e-9 { + t.Fatalf("standard error %d = %.15g, want %.15g (relative %.3g)", + j, res.StandardErrors[j], want, rel) + } + } +} + +// TestTheilSenParallelPairs checks the pairwise-slope walk against the +// serial enumeration it replaces, element for element and bit for bit, +// on a sample large enough to take the parallel path and with repeated +// predictors, which are the pairs the walk has to skip. The reference +// walk reproduces the original append order, so a slope landing in the +// wrong buffer slot moves the median and fails the test. +func TestTheilSenParallelPairs(t *testing.T) { + const n = 1400 + src := &randSource{state: 0xc0ffee123456789a} + xs := make([]float64, n) + ys := make([]float64, n) + for i := range n { + // A tenth of a unit apart, so a tenth of the pairs share a + // predictor and carry no slope. + xs[i] = math.Floor(src.next()*300) / 10 + ys[i] = 2*xs[i] + src.next() + } + gotIntercept, gotSlope, err := TheilSenRegression(mustFloats(t, xs), mustFloats(t, ys)) + if err != nil { + t.Fatal(err) + } + slopes := make([]float64, 0, n*(n-1)/2) + for i := range n { + for j := i + 1; j < n; j++ { + if dx := xs[j] - xs[i]; dx != 0 { + slopes = append(slopes, (ys[j]-ys[i])/dx) + } + } + } + slope := medianSlice(slopes) + intercepts := make([]float64, n) + for i := range n { + intercepts[i] = ys[i] - slope*xs[i] + } + intercept := medianSlice(intercepts) + if gotSlope != slope || gotIntercept != intercept { + t.Fatalf("Theil-Sen = (%v, %v), want (%v, %v) from the serial walk", + gotIntercept, gotSlope, intercept, slope) + } +} + +// regDesign builds a deterministic (n, p) design with an intercept column +// and covariates in (−0.5, 0.5), plus the row-major values for the +// cross-checks that rebuild a normal-equations matrix from them. +func regDesign(tb testing.TB, n, p int, seed uint64) (*core.Array, []float64) { + tb.Helper() + src := &randSource{state: seed} + vals := make([]float64, n*p) + for r := range n { + vals[r*p] = 1 + for c := 1; c < p; c++ { + vals[r*p+c] = src.next() - 0.5 + } + } + a, err := core.FromFloats(vals, n, p) + if err != nil { + tb.Fatal(err) + } + return a, vals +} + +// BenchmarkKendallTau measures the rank correlation at a size where the +// pair enumeration it replaced costs milliseconds: 8.4 million pairs, +// counted by the merge sort in a fraction of that. +func BenchmarkKendallTau(b *testing.B) { + const n = 4096 + src := &randSource{state: 0x123456789abcdef} + xv := make([]float64, n) + yv := make([]float64, n) + for i := range n { + // Rounded draws, so ties in both samples are part of the walk. + xv[i] = math.Floor(src.next()*512) / 512 + yv[i] = math.Floor(src.next()*512) / 512 + } + x, err := core.FromFloats(xv, n) + if err != nil { + b.Fatal(err) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := KendallTau(x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkLinearRegressionCovariance measures the least-squares fit with +// its inference on a wide design, where the covariance diagonal is a +// visible share of the work: p is large enough that the p solves it needs +// outgrow the normal equations themselves. +func BenchmarkLinearRegressionCovariance(b *testing.B) { + const n, p = 400, 48 + x, _ := regDesign(b, n, p, 0x5eed0001) + src := &randSource{state: 0x5eed0002} + yv := make([]float64, n) + for i := range n { + yv[i] = src.next()*2 - 1 + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := LinearRegression(x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkPoissonRegressionInference measures the count fit with its +// Wald inference, the covariance diagonal of the Fisher information +// included. +func BenchmarkPoissonRegressionInference(b *testing.B) { + const n, p = 600, 30 + x, vals := regDesign(b, n, p, 0x5eed0003) + src := &randSource{state: 0x5eed0004} + yv := make([]float64, n) + for i := range n { + eta := 0.4 + for c := range p { + eta += 0.3 * vals[i*p+c] + } + mu := math.Exp(eta) + yv[i] = math.Round(mu * (0.5 + src.next())) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := PoissonRegression(x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkLogisticRegressionInference measures the binary fit with its +// Wald inference, the covariance diagonal of the Fisher information +// included. +func BenchmarkLogisticRegressionInference(b *testing.B) { + const n, p = 1500, 30 + x, vals := regDesign(b, n, p, 0x5eed0005) + src := &randSource{state: 0x5eed0006} + yv := make([]float64, n) + for i := range n { + eta := 0.2 + 0.5*vals[i*p] + for c := 1; c < p; c++ { + eta += 0.5 * vals[i*p+c] + } + if src.next() < 1/(1+math.Exp(-eta)) { + yv[i] = 1 + } + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := LogisticRegression(x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkHuberRegressionInference measures the robust fit with its +// standard errors: the reweighting rounds, their medians, and the +// weighted normal equations read for the covariance diagonal. +func BenchmarkHuberRegressionInference(b *testing.B) { + const n, p = 800, 8 + x, vals := regDesign(b, n, p, 0x5eed0007) + src := &randSource{state: 0x5eed0008} + yv := make([]float64, n) + for i := range n { + yv[i] = 1 + 0.5*vals[i*p+1] - 0.3*vals[i*p+2] + 0.2*(src.next()-0.5) + } + yv[11] += 6 + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := HuberRegression(x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkTheilSenRegression measures the median-of-slopes fit: the +// pairwise walk over the pairs with distinct predictors and the median +// taken over them. +func BenchmarkTheilSenRegression(b *testing.B) { + const n = 512 + src := &randSource{state: 0x5eed0009} + xv := make([]float64, n) + yv := make([]float64, n) + for i := range n { + xv[i] = src.next()*10 - 5 + yv[i] = 1 + 2*xv[i] + src.next() + } + x, err := core.FromFloats(xv, n) + if err != nil { + b.Fatal(err) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, _, err := TheilSenRegression(x, y); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkLassoPath measures the warm-started regularisation path over +// the documented hundred-lambda grid, on an elastic net mixing that +// exercises the standardised coordinate descent. The design carries no +// intercept column: the penalty standardises every column, and a +// constant one has no slope to report. +func BenchmarkLassoPath(b *testing.B) { + const n, p = 400, 12 + src := &randSource{state: 0x5eed000a} + xv := make([]float64, n*p) + yv := make([]float64, n) + for i := range n { + for c := range p { + xv[i*p+c] = src.next() - 0.5 + } + yv[i] = 1 + 0.8*xv[i*p+1] - 0.5*xv[i*p+2] + 0.1*(src.next()-0.5) + } + x, err := core.FromFloats(xv, n, p) + if err != nil { + b.Fatal(err) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := LassoPath(x, y, 1); err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkKendallTauScaling reports the merge-sort pair count against +// the sample size: the linear-ish growth the count replaced the +// quadratic walk with. +func BenchmarkKendallTauScaling(b *testing.B) { + for _, n := range []int{512, 2048, 8192} { + src := &randSource{state: 0x51ed + uint64(n)} + xv := make([]float64, n) + yv := make([]float64, n) + for i := range n { + xv[i] = math.Floor(src.next()*512) / 512 + yv[i] = math.Floor(src.next()*512) / 512 + } + x, err := core.FromFloats(xv, n) + if err != nil { + b.Fatal(err) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := KendallTau(x, y); err != nil { + b.Fatal(err) + } + } + }) + } +} + +// BenchmarkTheilSenRegressionCap runs the fit at the observation cap, +// where the pairwise slope list holds millions of entries and the middle +// order statistic is taken by selection rather than by a full sort. +func BenchmarkTheilSenRegressionCap(b *testing.B) { + const n = TheilSenMaxObservations + src := &randSource{state: 0x5eed1234} + xv := make([]float64, n) + yv := make([]float64, n) + for i := range n { + xv[i] = src.next()*10 - 5 + yv[i] = 1 + 2*xv[i] + src.next() + } + x, err := core.FromFloats(xv, n) + if err != nil { + b.Fatal(err) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, _, err := TheilSenRegression(x, y); err != nil { + b.Fatal(err) + } + } +} diff --git a/stats/bench_test.go b/stats/bench_test.go new file mode 100644 index 0000000..3875069 --- /dev/null +++ b/stats/bench_test.go @@ -0,0 +1,265 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Benchmarks for the package's heavy paths: the O(n·m) kernel-density +// sweep, the sort-backed summaries, the regression and GLM normal +// equations and the incomplete-function machinery under the CDFs. + +// benchSeries builds a deterministic float vector of length n mixing a +// few incommensurate frequencies, so sorts and moment sums see a +// realistic spread without a generator. +func benchSeries(n int) []float64 { + vals := make([]float64, n) + for i := range vals { + x := float64(i) + vals[i] = math.Sin(0.001*x)*10 + math.Sin(0.013*x)*3 + math.Cos(0.11*x) + float64(i%7)*0.125 + } + return vals +} + +func benchArray(b *testing.B, n int, shape ...int) *core.Array { + b.Helper() + if len(shape) == 0 { + shape = []int{n} + } + a, err := core.FromFloats(benchSeries(n), shape...) + if err != nil { + b.Fatal(err) + } + return a +} + +func BenchmarkKernelDensity(b *testing.B) { + sample := benchArray(b, 1000) + points := benchArray(b, 256) + b.ReportAllocs() + for b.Loop() { + if _, err := KernelDensity(sample, 0.5, points); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkKernelDensitySilverman(b *testing.B) { + sample := benchArray(b, 2000) + points := benchArray(b, 512) + b.ReportAllocs() + for b.Loop() { + if _, err := KernelDensity(sample, 0, points); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkHistogram(b *testing.B) { + a := benchArray(b, 100000) + b.ReportAllocs() + for b.Loop() { + if _, _, err := Histogram(a, 64); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkRollingMean(b *testing.B) { + a := benchArray(b, 8192) + b.ReportAllocs() + for b.Loop() { + if _, err := RollingMean(a, 32); err != nil { + b.Fatal(err) + } + } +} + +// benchmarkRolling1M drives one rolling reduction over a million-sample +// series at the short and the long window, the pair the carried-update +// rescan trade-off is measured on. +func benchmarkRolling1M(b *testing.B, window int, sum bool) { + b.Helper() + a := benchArray(b, 1<<20) + b.ReportAllocs() + for b.Loop() { + var err error + if sum { + _, err = RollingSum(a, window) + } else { + _, err = RollingMean(a, window) + } + if err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkRollingSumWindow8(b *testing.B) { + benchmarkRolling1M(b, 8, true) +} + +func BenchmarkRollingSumWindow4096(b *testing.B) { + benchmarkRolling1M(b, 4096, true) +} + +func BenchmarkRollingMeanWindow8(b *testing.B) { + benchmarkRolling1M(b, 8, false) +} + +func BenchmarkRollingMeanWindow4096(b *testing.B) { + benchmarkRolling1M(b, 4096, false) +} + +func BenchmarkRollingMax(b *testing.B) { + a := benchArray(b, 8192) + b.ReportAllocs() + for b.Loop() { + if _, err := RollingMax(a, 32); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkMedian(b *testing.B) { + a := benchArray(b, 50000) + b.ReportAllocs() + for b.Loop() { + if _, err := Median(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCovarianceMatrix(b *testing.B) { + a := benchArray(b, 1000*16, 1000, 16) + b.ReportAllocs() + for b.Loop() { + if _, err := CovarianceMatrix(a); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLinearRegression(b *testing.B) { + n, p := 2000, 8 + xv := make([]float64, n*p) + yv := make([]float64, n) + for r := range n { + yv[r] = 0 + for c := range p { + v := float64((r*31+c*17)%97)/97 - 0.5 + xv[r*p+c] = v + yv[r] += float64(c+1) * v + } + } + x, err := core.FromFloats(xv, n, p) + if err != nil { + b.Fatal(err) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := LinearRegression(x, y); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkLogisticRegression(b *testing.B) { + n, p := 1000, 6 + xv := make([]float64, n*p) + yv := make([]float64, n) + for r := range n { + eta := -1.5 + for c := range p { + v := float64((r*13+c*29)%53)/53 - 0.5 + xv[r*p+c] = v + eta += float64(c) * v + } + if r%2 == 0 { + eta += 0.5 + } + if eta > 0 { + yv[r] = 1 + } + } + x, err := core.FromFloats(xv, n, p) + if err != nil { + b.Fatal(err) + } + y, err := core.FromFloats(yv, n) + if err != nil { + b.Fatal(err) + } + b.ReportAllocs() + for b.Loop() { + if _, err := LogisticRegression(x, y); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkGammaCDF(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := GammaCDF(2.5, 3.0, 1.2); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkBetaIncomplete(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := BetaIncomplete(0.4, 2.5, 3.5); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkStudentTQuantile(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := StudentTQuantile(0.975, 12); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkNormalQuantile(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := NormalQuantile(0.975); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkGammaQuantile(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + if _, err := GammaQuantile(0.99, 3, 1.2); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkKolmogorovSmirnovTest(b *testing.B) { + a := benchArray(b, 4000) + c := benchArray(b, 5000) + b.ReportAllocs() + for b.Loop() { + if _, _, err := KolmogorovSmirnovTest(a, c); err != nil { + b.Fatal(err) + } + } +} diff --git a/stats/cdf.go b/stats/cdf.go new file mode 100644 index 0000000..c984e85 --- /dev/null +++ b/stats/cdf.go @@ -0,0 +1,735 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import "sourcedock.dev/petrbalvin/tensor/internal/base" + +import ( + "math" +) + +// Distribution functions: cumulative distribution functions and their +// quantiles for the distributions the library draws from, so a +// p-value or a confidence interval needs no second library. +// +// The foundations are the regularised incomplete gamma and beta +// functions: every continuous CDF here reduces to one of them, and +// the two discrete CDFs reduce through the exact identities +// P(Poisson ≤ k) = ΓUpper(k+1, λ) and P(Binomial ≤ k) = +// BetaIncomplete(1−p, n−k, k+1), which hold in floating point to the +// accuracy of the incomplete functions themselves. Quantiles invert +// the CDF by a bracketed Newton iteration through the distribution's +// density, falling back to bisection whenever the derivative step +// would leave the bracket, so the convergence stays unconditional on +// every monotone CDF. + +// gammaSeriesEps bounds the series and continued-fraction iterations +// of the incomplete functions. +const gammaSeriesEps = 3e-16 + +// gammaIterations bounds the series and continued-fraction iterations +// for a shape a. Both converge in O(√a) rounds near x ≈ a, so a fixed +// cap would silently truncate the large-shape answers (a shape of +// 50000 at x ≈ a needs about 1800 series rounds); the budget grows +// with √a and every caller errors when it is still not enough. +func gammaIterations(a float64) int { + return 1000 + int(20*math.Sqrt(a)) +} + +// GammaLower returns the regularised lower incomplete gamma +// P(a, x) = γ(a, x)/Γ(a), the CDF of a Gamma(shape a, rate 1) draw. +// It evaluates the power series below x < a+1 and the continued +// fraction of the complement above, each to rounding level; for large +// shapes the series budget grows with √a, and either iteration that +// fails to converge inside its budget is an error rather than a +// silently truncated value. +func GammaLower(a, x float64) (float64, error) { + p, _, err := gammaIncomplete("GammaLower", a, x) + return p, err +} + +// GammaUpper returns the regularised upper incomplete gamma +// Q(a, x) = Γ(a, x)/Γ(a) = 1 − P(a, x), computed on the same +// foundation as GammaLower. +func GammaUpper(a, x float64) (float64, error) { + _, q, err := gammaIncomplete("GammaUpper", a, x) + return q, err +} + +// gammaIncomplete returns P(a, x) and Q(a, x) together, sharing the +// one exponent factor both need. +func gammaIncomplete(name string, a, x float64) (p, q float64, err error) { + if !(a > 0) { + return 0, 0, base.Errf("%s: shape a must be positive, got %g", name, a) + } + if !(x >= 0) { + return 0, 0, base.Errf("%s: x must not be negative, got %g", name, x) + } + if x == 0 { + return 0, 1, nil + } + if math.IsInf(x, 1) { + return 1, 0, nil + } + lg, _ := math.Lgamma(a) + front := math.Exp(-x + a*math.Log(x) - lg) + // The iteration budget is fixed by the shape alone, so it is + // computed once: the same bound the inline call evaluated, without + // re-deriving it on every round. + rounds := gammaIterations(a) + if x < a+1 { + // Power series for P: γ(a, x) = e^{−x}x^a Σ x^n/(a(a+1)…(a+n)). + // Every term is positive, so the sum carries no cancellation and + // stays accurate for arbitrarily large shapes given enough + // rounds. + sum := 1 / a + del := sum + ap := a + converged := false + for range rounds { + ap++ + del *= x / ap + sum += del + if math.Abs(del) < math.Abs(sum)*gammaSeriesEps { + converged = true + break + } + } + if !converged { + return 0, 0, base.Errf("%s: the power series did not converge for a = %g, x = %g", name, a, x) + } + return sum * front, 1 - sum*front, nil + } + // Continued fraction for Q (Numerical Recipes form), Lentz's + // modified method with the tiny floor. + const fpmin = 1e-300 + b := x + 1 - a + c := 1 / fpmin + d := 1 / b + h := d + converged := false + for i := 1; i <= rounds; i++ { + an := -float64(i) * (float64(i) - a) + b += 2 + d = an*d + b + if math.Abs(d) < fpmin { + d = fpmin + } + c = b + an/c + if math.Abs(c) < fpmin { + c = fpmin + } + d = 1 / d + del := d * c + h *= del + if math.Abs(del-1) < gammaSeriesEps { + converged = true + break + } + } + if !converged { + return 0, 0, base.Errf("%s: the continued fraction did not converge for a = %g, x = %g", name, a, x) + } + return 1 - front*h, front * h, nil +} + +// BetaIncomplete returns the regularised incomplete beta function +// I_x(a, b), the CDF of a Beta(a, b) draw, by the continued fraction +// with the standard symmetry switch at x > (a+1)/(a+b+2). +func BetaIncomplete(x, a, b float64) (float64, error) { + if !(a > 0 && b > 0) { + return 0, base.Errf("BetaIncomplete: shape parameters must be positive, got %g, %g", a, b) + } + if !(x >= 0 && x <= 1) { + return 0, base.Errf("BetaIncomplete: x must lie in [0, 1], got %g", x) + } + if x == 0 || x == 1 { + return x, nil + } + lab, _ := math.Lgamma(a + b) + la, _ := math.Lgamma(a) + lb, _ := math.Lgamma(b) + front := math.Exp(lab - la - lb + a*math.Log(x) + b*math.Log(1-x)) + // The continued fraction converges fastest away from the skew side. + if x < (a+1)/(a+b+2) { + h, err := betaContinue(a, b, x) + if err != nil { + return 0, base.Errf("BetaIncomplete: %w", err) + } + return front * h / a, nil + } + h, err := betaContinue(b, a, 1-x) + if err != nil { + return 0, base.Errf("BetaIncomplete: %w", err) + } + return 1 - front*h/b, nil +} + +// betaContinue evaluates the incomplete-beta continued fraction +// (Numerical Recipes form) by Lentz's modified method. The round +// budget grows with the shapes exactly as the incomplete gamma's does, +// and an iteration that fails to converge inside it is an error rather +// than a silently truncated value. +func betaContinue(a, b, x float64) (float64, error) { + const fpmin = 1e-300 + rounds := max(1000, gammaIterations(a)+gammaIterations(b)) + qab := a + b + qap := a + 1 + qam := a - 1 + c := 1.0 + d := 1 - qab*x/qap + if math.Abs(d) < fpmin { + d = fpmin + } + d = 1 / d + h := d + converged := false + for i := 1; i <= rounds; i++ { + m := float64(i) + m2 := 2 * m + aa := m * (b - m) * x / ((qam + m2) * (a + m2)) + d = 1 + aa*d + if math.Abs(d) < fpmin { + d = fpmin + } + c = 1 + aa/c + if math.Abs(c) < fpmin { + c = fpmin + } + d = 1 / d + h *= d * c + aa = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2)) + d = 1 + aa*d + if math.Abs(d) < fpmin { + d = fpmin + } + c = 1 + aa/c + if math.Abs(c) < fpmin { + c = fpmin + } + d = 1 / d + del := d * c + h *= del + if math.Abs(del-1) < gammaSeriesEps { + converged = true + break + } + } + if !converged { + return 0, base.Errf("the continued fraction did not converge for a = %g, b = %g, x = %g", a, b, x) + } + return h, nil +} + +// NormalCDF returns Φ(x), the standard normal cumulative distribution. +func NormalCDF(x float64) float64 { + return 0.5 * math.Erfc(-x/math.Sqrt2) +} + +// ExponentialCDF returns P(X ≤ x) for X ~ Exponential(rate). +func ExponentialCDF(x, rate float64) (float64, error) { + if !(rate > 0) { + return 0, base.Errf("ExponentialCDF: rate must be positive, got %g", rate) + } + if math.IsNaN(x) { + // The x ≤ 0 branch below is the left tail, not a rejection, so + // NaN needs its own guard: without it the answer is a silent NaN. + return 0, base.Errf("ExponentialCDF: x must be a number, got %g", x) + } + if x <= 0 { + return 0, nil + } + // −Expm1(−rate·x) is 1 − e^{−rate·x} without the cancellation: the + // literal form loses the left tail, an 11 % relative error already + // at rate·x = 1e-16 and a total loss not far below, while Expm1 + // keeps every bit of the small answer. + return -math.Expm1(-rate * x), nil +} + +// GammaCDF returns P(X ≤ x) for X ~ Gamma(shape, rate), the same +// parametrisation GammaDraws samples. +func GammaCDF(x, shape, rate float64) (float64, error) { + if !(shape > 0 && rate > 0) { + return 0, base.Errf("GammaCDF: shape and rate must be positive, got %g, %g", shape, rate) + } + return GammaLower(shape, x*rate) +} + +// ChiSquareCDF returns P(X ≤ x) for X ~ χ²(df). +func ChiSquareCDF(x float64, df int) (float64, error) { + if df < 1 { + return 0, base.Errf("ChiSquareCDF: df must be ≥ 1, got %d", df) + } + return GammaLower(float64(df)/2, x/2) +} + +// StudentTCDF returns P(T ≤ t) for T ~ Student t(df). +func StudentTCDF(t float64, df int) (float64, error) { + if df < 1 { + return 0, base.Errf("StudentTCDF: df must be ≥ 1, got %d", df) + } + z := float64(df) / (float64(df) + t*t) + upper, err := BetaIncomplete(z, float64(df)/2, 0.5) + if err != nil { + return 0, base.Errf("StudentTCDF: %w", err) + } + if t >= 0 { + return 1 - upper/2, nil + } + return upper / 2, nil +} + +// PoissonCDF returns P(N ≤ k) for N ~ Poisson(lambda), through the +// identity P(N ≤ k) = ΓUpper(k+1, λ). +func PoissonCDF(k int, lambda float64) (float64, error) { + if !(lambda > 0) { + return 0, base.Errf("PoissonCDF: lambda must be positive, got %g", lambda) + } + if k < 0 { + return 0, nil + } + return GammaUpper(float64(k+1), lambda) +} + +// BinomialCDF returns P(X ≤ k) for X ~ Binomial(trials, p), through +// the identity P(X ≤ k) = I_{1−p}(n−k, k+1). +func BinomialCDF(k, trials int, p float64) (float64, error) { + if trials < 1 { + return 0, base.Errf("BinomialCDF: trials must be ≥ 1, got %d", trials) + } + if !(p > 0 && p < 1) { + return 0, base.Errf("BinomialCDF: p must lie in (0, 1), got %g", p) + } + if k < 0 { + return 0, nil + } + if k >= trials { + return 1, nil + } + return BetaIncomplete(1-p, float64(trials-k), float64(k+1)) +} + +// continuousQuantile inverts a continuous CDF on the positive axis: +// the bracket starts at seed and doubles outward until the CDF +// straddles q, then the crossing is refined by a bracketed Newton +// iteration through the CDF's derivative pdf when one is given, and +// by plain bisection when it is not. The Newton walk carries the +// bracket with it and bisects whenever the derivative step would +// leave the bracket or would not cut it fast enough, so the +// convergence stays unconditional on every monotone CDF either way. +func continuousQuantile(name string, q float64, seed float64, + cdf func(float64) (float64, error), pdf func(float64) float64) (float64, error) { + // The guard is NaN-rejecting on purpose: NaN compares false against + // both bounds, so the written-out form would let it past and every + // bracket comparison below would then also be false. + if !(q >= 0 && q <= 1) { + return 0, base.Errf("%s: q must lie in [0, 1], got %g", name, q) + } + if q == 0 || q == 1 { + return 0, base.Errf("%s: q = %g has no finite quantile", name, q) + } + lo, hi := seed, seed + fLo, err := cdf(lo) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + fHi, err := cdf(hi) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + // Establish a bracket where the CDF crosses q. + for fLo > q { + lo /= 2 + // Halving from a finite seed underflows through the denormals + // to exactly zero, which is the whole representable range: no + // CDF value can sit below a probability of that size. + if lo == 0 { + return 0, base.Errf("%s: failed to bracket q = %g from below", name, q) + } + if fLo, err = cdf(lo); err != nil { + return 0, base.Errf("%s: %w", name, err) + } + } + for fHi < q { + hi *= 2 + if math.IsInf(hi, 0) { + return 0, base.Errf("%s: failed to bracket q = %g from above", name, q) + } + if fHi, err = cdf(hi); err != nil { + return 0, base.Errf("%s: %w", name, err) + } + } + if pdf == nil { + // Halve to rounding level. The interval at least halves every pass + // and the loop leaves on mid == lo || mid == hi, which a finite + // bracket always reaches (about 1100 passes from the widest one), + // so the iteration count is a guard against a non-monotone CDF, + // not a working part of the convergence: hitting it is an error, + // never a quietly unconverged midpoint. + converged := false + for range 4096 { + mid := (lo + hi) / 2 + if mid == lo || mid == hi { + converged = true + break + } + f, err := cdf(mid) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if f < q { + lo = mid + } else { + hi = mid + } + } + if !converged { + return 0, base.Errf("%s: the bisection for q = %g did not converge", name, q) + } + return (lo + hi) / 2, nil + } + return continuousQuantileNewton(name, q, lo, hi, cdf, pdf) +} + +// continuousQuantileNewton refines a bracket where the CDF crosses q +// by a safeguarded Newton walk: a Newton step on cdf(x) = q through +// the pdf derivative whenever that step lands inside the bracket and +// cuts it at least as fast as the bisection it replaces, a bisection +// step otherwise. The bisection guarantee survives untouched: the +// bracket at least halves every second round, the walk leaves when the +// bracket has no float strictly inside it, the same rounding-level +// exit the plain bisection takes, and a vanishing or non-finite +// derivative fails both Newton guards and bisects. +func continuousQuantileNewton(name string, q float64, lo, hi float64, + cdf func(float64) (float64, error), pdf func(float64) float64) (float64, error) { + x := 0.5 * (lo + hi) + f, err := cdf(x) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if f == q { + return x, nil + } + if f < q { + lo = x + } else { + hi = x + } + // stepSize carries the previous round's step, the yardstick the + // second Newton guard measures against: a derivative step that does + // not halve it is stalling, and the round bisects instead. + stepSize := hi - lo + converged := false + for range 4096 { + d := pdf(x) + g := f - q + // The two rtsafe guards: the derivative step is taken only when + // it lands strictly inside the bracket, x - g/d in (lo, hi), + // and cuts the last step at least in half. A vanishing or + // non-finite derivative fails the first guard and bisects. + newton := d > 0 && !math.IsInf(d, 1) && + (x-lo)*d > g && g > (x-hi)*d && + 2*math.Abs(g) <= math.Abs(stepSize*d) + if newton { + stepSize = g / d + x -= stepSize + } else { + stepSize = 0.5 * (hi - lo) + x = 0.5 * (lo + hi) + } + if x == lo || x == hi { + converged = true + break + } + if f, err = cdf(x); err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if f == q { + converged = true + break + } + if f < q { + lo = x + } else { + hi = x + } + } + if !converged { + return 0, base.Errf("%s: the quantile iteration for q = %g did not converge", name, q) + } + return x, nil +} + +// continuousQuantileUpper inverts a monotone decreasing upper-tail +// probability on the positive axis: it returns the z > 0 with +// upper(z) = q, for q below what the reflected CDF can represent +// (below 2⁻⁵³, where 1−q rounds to 1). The bracket starts at [0, 1] +// and doubles outward; upper(0) = 0.5 covers every such q. +func continuousQuantileUpper(name string, q float64, + upper func(float64) (float64, error)) (float64, error) { + lo, hi := 0.0, 1.0 + fHi, err := upper(hi) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + for fHi > q { + lo = hi + hi *= 2 + if math.IsInf(hi, 0) { + return 0, base.Errf("%s: failed to bracket q = %g from above", name, q) + } + if fHi, err = upper(hi); err != nil { + return 0, base.Errf("%s: %w", name, err) + } + } + converged := false + for range 4096 { + mid := (lo + hi) / 2 + if mid == lo || mid == hi { + converged = true + break + } + f, err := upper(mid) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if f > q { + lo = mid + } else { + hi = mid + } + } + if !converged { + return 0, base.Errf("%s: the bisection for q = %g did not converge", name, q) + } + return (lo + hi) / 2, nil +} + +// discreteQuantile returns the smallest integer k whose CDF reaches +// q, found by doubling the upper bound then bisecting the integer +// grid. The result is a whole number in float form. +func discreteQuantile(name string, q float64, cdf func(int) (float64, error)) (float64, error) { + // NaN-rejecting for the same reason continuousQuantile's guard is: + // a NaN q makes every comparison below false, and the doubling + // bracket would then run without end. + if !(q >= 0 && q <= 1) { + return 0, base.Errf("%s: q must lie in [0, 1], got %g", name, q) + } + if q == 0 { + return 0, nil + } + hi := 1 + for { + f, err := cdf(hi) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if f >= q { + break + } + // A monotone CDF reaches any q > 0 at some finite k, so a bound + // this high means the CDF never will: refuse rather than let hi + // wrap to zero and the doubling loop spin for ever. + if hi > math.MaxInt/2 { + return 0, base.Errf("%s: failed to bracket q = %g", name, q) + } + hi *= 2 + } + lo := 0 + for lo < hi { + mid := (lo + hi) / 2 + f, err := cdf(mid) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if f >= q { + hi = mid + } else { + lo = mid + 1 + } + } + return float64(lo), nil +} + +// normalPdf is the standard normal density, the derivative +// NormalQuantile's Newton walk steps through. +func normalPdf(x float64) float64 { + return math.Exp(-0.5*x*x) / math.Sqrt(2*math.Pi) +} + +// gammaShapeRatePdf is the Gamma(shape, rate) density at x, evaluated +// from its logarithm so a large shape's normalising factor never +// overflows on its way to a density the format cannot hold. x at or +// below zero answers 0: the quantile walk treats a vanishing +// derivative as a signal to bisect. +func gammaShapeRatePdf(x, shape, rate float64) float64 { + if x <= 0 { + return 0 + } + return math.Exp((shape-1)*math.Log(x) - rate*x + shape*math.Log(rate) - logGamma(shape)) +} + +// studentTPdf is the Student t(df) density at x, from its logarithm +// like the gamma density, so the wide tails' underflow lands on the 0 +// the quantile walk reads as a bisection signal. +func studentTPdf(x float64, df int) float64 { + d := float64(df) + return math.Exp(logGamma(0.5*(d+1)) - logGamma(0.5*d) - + 0.5*math.Log(d*math.Pi) - 0.5*(d+1)*math.Log1p(x*x/d)) +} + +// NormalQuantile returns the q-quantile of the standard normal +// distribution, the inverse of NormalCDF. +func NormalQuantile(q float64) (float64, error) { + if !(q >= 0 && q <= 1) { + return 0, base.Errf("NormalQuantile: q must lie in [0, 1], got %g", q) + } + if q == 0 || q == 1 { + return 0, base.Errf("NormalQuantile: q = %g has no finite quantile", q) + } + // The bracket expansion below only walks up from the seed, and + // Phi(0) = 0.5 is the infimum it can reach, so the whole lower half + // of the domain is bracketed on the reflected probability instead: + // Phi is symmetric, so the lower-tail quantile is the mirror of the + // upper one. Below 2⁻⁵³ the reflection 1−q rounds to exactly 1 and + // would be read as the q = 1 refusal, so the extreme tail bisects + // the accurate upper tail 0.5·Erfc directly, on the positive axis. + if q < 0.5 { + if 1-q == 1 { + z, err := continuousQuantileUpper("NormalQuantile", q, func(x float64) (float64, error) { + return 0.5 * math.Erfc(x/math.Sqrt2), nil + }) + if err != nil { + return 0, err + } + return -z, nil + } + z, err := continuousQuantile("NormalQuantile", 1-q, 1, func(x float64) (float64, error) { + return NormalCDF(x), nil + }, normalPdf) + if err != nil { + return 0, err + } + return -z, nil + } + return continuousQuantile("NormalQuantile", q, 1, func(x float64) (float64, error) { + return NormalCDF(x), nil + }, normalPdf) +} + +// ExponentialQuantile returns the q-quantile of Exponential(rate). +func ExponentialQuantile(q, rate float64) (float64, error) { + return continuousQuantile("ExponentialQuantile", q, 1, func(x float64) (float64, error) { + return ExponentialCDF(x, rate) + }, func(x float64) float64 { + return rate * math.Exp(-rate*x) + }) +} + +// GammaQuantile returns the q-quantile of Gamma(shape, rate). +func GammaQuantile(q, shape, rate float64) (float64, error) { + return continuousQuantile("GammaQuantile", q, shape/rate, func(x float64) (float64, error) { + return GammaCDF(x, shape, rate) + }, func(x float64) float64 { + return gammaShapeRatePdf(x, shape, rate) + }) +} + +// ChiSquareQuantile returns the q-quantile of χ²(df), the critical +// value tables quote. +func ChiSquareQuantile(q float64, df int) (float64, error) { + shape := float64(df) / 2 + return continuousQuantile("ChiSquareQuantile", q, float64(df), func(x float64) (float64, error) { + return ChiSquareCDF(x, df) + }, func(x float64) float64 { + return gammaShapeRatePdf(x, shape, 0.5) + }) +} + +// StudentTQuantile returns the q-quantile of Student t(df). The signed +// axis is bracketed through a two-sided bracket expansion, the crossing +// refined by the bracketed Newton walk through the t density. +func StudentTQuantile(q float64, df int) (float64, error) { + if !(q >= 0 && q <= 1) { + return 0, base.Errf("StudentTQuantile: q must lie in [0, 1], got %g", q) + } + if q == 0 || q == 1 { + return 0, base.Errf("StudentTQuantile: q = %g has no finite quantile", q) + } + // Bracket on the positive side, then mirror for q < 0.5. Below 2⁻⁵³ + // the reflection 1−q rounds to exactly 1 and would be read as the + // q = 1 refusal, so the extreme tail bisects the one-sided upper + // tail directly, on the positive axis. + if q < 0.5 { + if 1-q == 1 { + t, err := continuousQuantileUpper("StudentTQuantile", q, func(x float64) (float64, error) { + return studentTUpperTail(x, df) + }) + if err != nil { + return 0, err + } + return -t, nil + } + t, err := continuousQuantile("StudentTQuantile", 1-q, 1, func(x float64) (float64, error) { + return StudentTCDF(x, df) + }, func(x float64) float64 { + return studentTPdf(x, df) + }) + if err != nil { + return 0, err + } + return -t, nil + } + return continuousQuantile("StudentTQuantile", q, 1, func(x float64) (float64, error) { + return StudentTCDF(x, df) + }, func(x float64) float64 { + return studentTPdf(x, df) + }) +} + +// studentTUpperTail returns P(T > t) for Student t(df), one tail of +// the symmetric two-sided form twoSidedT uses. Beyond t²'s overflow +// the incomplete-beta argument saturates to 0 and would report a +// phantom zero tail exactly where the heavy df ≤ 2 tails stay +// representable, so those degrees answer by their asymptotic forms. +func studentTUpperTail(t float64, df int) (float64, error) { + if t > 1.3e154 { + switch { + case df == 1: + // Cauchy: the exact tail 1/2 − atan(t)/π, in a form the + // huge t keeps accurate. + return 1 / (math.Pi * t), nil + case df == 2: + return 1 / (2 * t * t), nil + } + } + z := float64(df) / (float64(df) + t*t) + p, err := BetaIncomplete(z, float64(df)/2, 0.5) + if err != nil { + return 0, err + } + return p / 2, nil +} + +// PoissonQuantile returns the smallest k with P(N ≤ k) ≥ q for +// N ~ Poisson(lambda). +func PoissonQuantile(q, lambda float64) (float64, error) { + if !(lambda > 0) { + return 0, base.Errf("PoissonQuantile: lambda must be positive, got %g", lambda) + } + return discreteQuantile("PoissonQuantile", q, func(k int) (float64, error) { + return PoissonCDF(k, lambda) + }) +} + +// BinomialQuantile returns the smallest k with P(X ≤ k) ≥ q for +// X ~ Binomial(trials, p). +func BinomialQuantile(q, p float64, trials int) (float64, error) { + if !(p > 0 && p < 1) { + return 0, base.Errf("BinomialQuantile: p must lie in (0, 1), got %g", p) + } + return discreteQuantile("BinomialQuantile", q, func(k int) (float64, error) { + return BinomialCDF(k, trials, p) + }) +} diff --git a/stats/cdf_test.go b/stats/cdf_test.go new file mode 100644 index 0000000..0ef59bf --- /dev/null +++ b/stats/cdf_test.go @@ -0,0 +1,278 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "testing" +) + +// TestGammaIncompleteClosed checks the incomplete gamma against +// closed forms: P(1, x) = 1 − e^{−x}, P(0.5, x) = erf(√x), +// Q(2, 1) = 3/e and the complement identity. +func TestGammaIncompleteClosed(t *testing.T) { + p, err := GammaLower(1, 2.5) + if err != nil { + t.Fatalf("GammaLower: %v", err) + } + if math.Abs(p-(1-math.Exp(-2.5))) > 1e-14 { + t.Fatalf("P(1, 2.5) = %.16g, want %.16g", p, 1-math.Exp(-2.5)) + } + p, err = GammaLower(0.5, 1) + if err != nil { + t.Fatalf("GammaLower: %v", err) + } + if math.Abs(p-math.Erf(1)) > 1e-14 { + t.Fatalf("P(0.5, 1) = %.16g, want erf(1) = %.16g", p, math.Erf(1)) + } + q, err := GammaUpper(2, 1) + if err != nil { + t.Fatalf("GammaUpper: %v", err) + } + if math.Abs(q-2/math.E) > 1e-14 { + t.Fatalf("Q(2, 1) = %.16g, want 2/e = %.16g", q, 2/math.E) + } + p, err = GammaLower(3, 0.7) + if err != nil { + t.Fatalf("GammaLower: %v", err) + } + q, err = GammaUpper(3, 0.7) + if err != nil { + t.Fatalf("GammaUpper: %v", err) + } + if math.Abs(p+q-1) > 1e-14 { + t.Fatalf("P + Q = %.17g, want 1", p+q) + } + // The series side of a large shape and the fraction side of a + // small one must agree with the complement to rounding level. + p, _ = GammaLower(8.5, 9.3) + q, _ = GammaUpper(8.5, 9.3) + if math.Abs(p+q-1) > 1e-13 { + t.Fatalf("P(8.5, 9.3) + Q(8.5, 9.3) = %.17g, want 1", p+q) + } + if _, err := GammaLower(0, 1); err == nil { + t.Fatal("a = 0: want an error") + } + if _, err := GammaLower(1, -1); err == nil { + t.Fatal("negative x: want an error") + } +} + +// TestBetaIncompleteClosed checks the incomplete beta against closed +// forms and the symmetry I_x(a,b) = 1 − I_{1−x}(b,a). +func TestBetaIncompleteClosed(t *testing.T) { + cases := []struct { + x, a, b, want float64 + }{ + {0.3, 1, 1, 0.3}, // uniform + {0.2, 1, 3, 0.488}, // 1 − (1−x)^b + {0.09, 2, 1, 0.0081}, // x^a + {0.5, 2, 2, 0.5}, // symmetric + {0.5, 0.5, 0.5, 0.5}, // arcsine, symmetric + {0.25, 0.5, 1, 0.5}, // I_x(1/2,1) = √x + } + for _, c := range cases { + got, err := BetaIncomplete(c.x, c.a, c.b) + if err != nil { + t.Fatalf("BetaIncomplete(%g, %g, %g): %v", c.x, c.a, c.b, err) + } + if math.Abs(got-c.want) > 1e-13 { + t.Fatalf("I_%g(%g, %g) = %.16g, want %.16g", c.x, c.a, c.b, got, c.want) + } + } + got, err := BetaIncomplete(0.7, 2, 3) + if err != nil { + t.Fatalf("BetaIncomplete: %v", err) + } + mirror, err := BetaIncomplete(0.3, 3, 2) + if err != nil { + t.Fatalf("BetaIncomplete mirror: %v", err) + } + if math.Abs(got+mirror-1) > 1e-13 { + t.Fatalf("symmetry broken: %.17g + %.17g != 1", got, mirror) + } + if _, err := BetaIncomplete(1.5, 1, 1); err == nil { + t.Fatal("x outside [0, 1]: want an error") + } +} + +// TestContinuousCDFs pins each continuous CDF on exact or tabulated +// values. +func TestContinuousCDFs(t *testing.T) { + if v := NormalCDF(1); math.Abs(v-0.8413447460685429) > 1e-15 { + t.Fatalf("Φ(1) = %.16g", v) + } + if v := NormalCDF(1.959963984540054); math.Abs(v-0.975) > 1e-14 { + t.Fatalf("Φ(1.96…) = %.16g, want 0.975", v) + } + v, err := ExponentialCDF(1, 1) + if err != nil || math.Abs(v-(1-1/math.E)) > 1e-15 { + t.Fatalf("ExponentialCDF(1, 1) = %v, %v", v, err) + } + // Gamma(2, 1) CDF: 1 − e^{−x}(1 + x). + v, err = GammaCDF(1, 2, 1) + if err != nil || math.Abs(v-(1-2/math.E)) > 1e-14 { + t.Fatalf("GammaCDF(1, 2, 1) = %v, %v", v, err) + } + // χ² with df 2 is the exponential with mean 2: 1 − e^{−x/2}. + v, err = ChiSquareCDF(1, 2) + if err != nil || math.Abs(v-(1-math.Exp(-0.5))) > 1e-14 { + t.Fatalf("ChiSquareCDF(1, 2) = %v, %v", v, err) + } + // χ² 95 % critical value with df 1: 3.841458820694124. + v, err = ChiSquareCDF(3.841458820694124, 1) + if err != nil || math.Abs(v-0.95) > 1e-13 { + t.Fatalf("ChiSquareCDF at the tabulated 95 %% point = %v, %v", v, err) + } + // Student t with df 1 is the Cauchy: 0.5 + atan(t)/π. + v, err = StudentTCDF(1, 1) + if err != nil || math.Abs(v-(0.5+math.Atan(1)/math.Pi)) > 1e-14 { + t.Fatalf("StudentTCDF(1, 1) = %v, %v", v, err) + } + // Student t with df 2: 0.5 + t/(2√(2 + t²)). + v, err = StudentTCDF(1, 2) + if err != nil || math.Abs(v-(0.5+1/(2*math.Sqrt(3)))) > 1e-14 { + t.Fatalf("StudentTCDF(1, 2) = %v, %v", v, err) + } +} + +// TestDiscreteCDFs pins the Poisson and binomial CDFs on exact sums. +func TestDiscreteCDFs(t *testing.T) { + v, err := PoissonCDF(1, 1) + if err != nil || math.Abs(v-2/math.E) > 1e-14 { + t.Fatalf("PoissonCDF(1, 1) = %v, %v, want 2/e", v, err) + } + // P(N ≤ 4) for λ = 4: e^{−4}·Σ_{k≤4} 4^k/k!. + sum := 0.0 + for k, term := 0, 1.0; k <= 4; k++ { + if k > 0 { + term *= 4 / float64(k) + } + sum += term + } + v, err = PoissonCDF(4, 4) + if err != nil || math.Abs(v-sum*math.Exp(-4)) > 1e-14 { + t.Fatalf("PoissonCDF(4, 4) = %v, %v, want %.16g", v, err, sum*math.Exp(-4)) + } + // Binomial(10, 0.5) at 5: 638/1024 by symmetry of the row. + v, err = BinomialCDF(5, 10, 0.5) + if err != nil || math.Abs(v-638.0/1024) > 1e-14 { + t.Fatalf("BinomialCDF(5, 10, 0.5) = %v, %v, want 0.623046875", v, err) + } + v, err = BinomialCDF(-1, 10, 0.5) + if err != nil || v != 0 { + t.Fatalf("BinomialCDF(-1, …) = %v, %v", v, err) + } + v, err = BinomialCDF(10, 10, 0.5) + if err != nil || v != 1 { + t.Fatalf("BinomialCDF(10, 10, 0.5) = %v, %v", v, err) + } +} + +// TestContinuousQuantiles inverts every continuous CDF and pins the +// tabulated critical values. +func TestContinuousQuantiles(t *testing.T) { + v, err := NormalQuantile(0.975) + if err != nil || math.Abs(v-1.959963984540054) > 1e-12 { + t.Fatalf("NormalQuantile(0.975) = %v, %v", v, err) + } + v, err = StudentTQuantile(0.975, 10) + if err != nil || math.Abs(v-2.228138851986273) > 1e-9 { + t.Fatalf("StudentTQuantile(0.975, 10) = %v, %v, want 2.2281…", v, err) + } + neg, err := StudentTQuantile(0.025, 10) + if err != nil || math.Abs(neg+2.228138851986273) > 1e-9 { + t.Fatalf("StudentTQuantile(0.025, 10) = %v, %v", neg, err) + } + v, err = ChiSquareQuantile(0.95, 1) + if err != nil || math.Abs(v-3.841458820694124) > 1e-9 { + t.Fatalf("ChiSquareQuantile(0.95, 1) = %v, %v", v, err) + } + v, err = ExponentialQuantile(0.6321205588285577, 1) + if err != nil || math.Abs(v-1) > 1e-12 { + t.Fatalf("ExponentialQuantile(1−1/e, 1) = %v, %v", v, err) + } + v, err = GammaQuantile(0.5, 2, 1) + if err != nil || math.Abs(v-1.678346990016661) > 1e-9 { + t.Fatalf("GammaQuantile(0.5, 2, 1) = %v, %v, want 1.67834…", v, err) + } + // Round trip: the CDF at every quantile must return q. + for _, q := range []float64{0.01, 0.1, 0.5, 0.9, 0.999} { + v, err := GammaQuantile(q, 3.5, 2) + if err != nil { + t.Fatalf("GammaQuantile(%g): %v", q, err) + } + back, err := GammaCDF(v, 3.5, 2) + if err != nil || math.Abs(back-q) > 1e-11 { + t.Fatalf("round trip q = %g: CDF(quantile) = %v, %v", q, back, err) + } + } + if _, err := NormalQuantile(1.5); err == nil { + t.Fatal("q outside [0, 1]: want an error") + } + if _, err := NormalQuantile(0); err == nil { + t.Fatal("q = 0: want an error") + } +} + +// TestDiscreteQuantiles checks the smallest-k rule on exact cases. +func TestDiscreteQuantiles(t *testing.T) { + v, err := PoissonQuantile(0.5, 1) + if err != nil || v != 1 { + t.Fatalf("PoissonQuantile(0.5, 1) = %v, %v, want 1", v, err) + } + // P(N ≤ 2) = 5/(2e) ≈ 0.9197, P(N ≤ 1) = 2/e ≈ 0.7358 for λ = 1: + // the 0.8 quantile is the smallest k reaching it, k = 2. + v, err = PoissonQuantile(0.8, 1) + if err != nil || v != 2 { + t.Fatalf("PoissonQuantile(0.8, 1) = %v, %v, want 2", v, err) + } + // Binomial(10, 0.5) median: smallest k with CDF ≥ 0.5 is 5. + v, err = BinomialQuantile(0.5, 0.5, 10) + if err != nil || v != 5 { + t.Fatalf("BinomialQuantile(0.5, 0.5, 10) = %v, %v, want 5", v, err) + } + if _, err := PoissonQuantile(0.5, 0); err == nil { + t.Fatal("lambda = 0: want an error") + } +} + +// TestGammaLowerLargeShape pins the large-shape region: both the power +// series (x < a+1) and the continued fraction (x ≥ a+1) must deliver +// accurate values near x ≈ a where the shape makes √a-sized iteration +// counts necessary, instead of silently returning truncated sums. +// The reference is the Wilson-Hilferty normal approximation, good to +// a few digits at these shapes. +func TestGammaLowerLargeShape(t *testing.T) { + wh := func(a, x float64) float64 { + z := 3 * math.Sqrt(a) * (math.Cbrt(x/a) - (1 - 1/(9*a))) + return 0.5 * math.Erfc(-z/math.Sqrt2) + } + cases := []struct { + a, x, tol float64 + }{ + {50000, 50000, 1e-4}, // series branch, √a ≈ 224 + {50000, 49500, 5e-4}, // lower tail, series + {50000, 50500, 5e-4}, // upper tail, continued fraction + {200000, 200000, 1e-4}, + {1000, 1000, 1e-5}, + } + for _, tc := range cases { + p, err := GammaLower(tc.a, tc.x) + if err != nil { + t.Fatalf("GammaLower(%g, %g): %v", tc.a, tc.x, err) + } + want := wh(tc.a, tc.x) + if math.Abs(p-want) > tc.tol { + t.Fatalf("GammaLower(%g, %g) = %.10g, want ≈ %.10g (tol %g)", tc.a, tc.x, p, want, tc.tol) + } + q, err := GammaUpper(tc.a, tc.x) + if err != nil { + t.Fatalf("GammaUpper(%g, %g): %v", tc.a, tc.x, err) + } + if p+q != 1 { + t.Fatalf("GammaLower + GammaUpper = %.17g, want exactly 1", p+q) + } + } +} diff --git a/stats/cluster.go b/stats/cluster.go new file mode 100644 index 0000000..f17dfb4 --- /dev/null +++ b/stats/cluster.go @@ -0,0 +1,771 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Clustering: k-means with k-means++ seeding and the Gaussian mixture +// fitted by expectation maximisation over the multivariate normal +// machinery of mvn.go. Both fits are deterministic for a given +// generator state: every random choice runs through the house +// generator, never through a global source. + +// kmeansMaxIterations caps the Lloyd sweeps; a fit that has not +// settled by then is reported with Converged false rather than +// pretended to have converged. +const kmeansMaxIterations = 300 + +// kmeansTolerance is the convergence tolerance of the Lloyd loop: the +// sweep stops once the largest coordinate movement of any centre is at +// most kmeansTolerance scaled by the largest absolute coordinate in +// the data (floored at 1, so a degenerate all-zero cloud still +// converges on an absolute scale). The rule mirrors the fixed points +// of Lloyd's iteration: past the tolerance the objective moves below +// rounding, and the label assignment no longer changes. +const kmeansTolerance = 1e-10 + +// KMeansResult carries the fit of k-means over a sample. +type KMeansResult struct { + // Centres are the k fitted centroids, one row of d coordinates + // each. + Centres [][]float64 + // Labels assigns every sample row its cluster, 0 to k-1. + Labels []int + // Inertia is the within-cluster sum of squared distances to the + // centres, the objective the Lloyd loop minimises. + Inertia float64 + // Iterations counts the Lloyd sweeps taken; Converged reports + // whether the centre movement fell under kmeansTolerance before + // the iteration budget ran out. + Iterations int + Converged bool +} + +// KMeans partitions the n rows of x (n rows, d columns) into k +// clusters by Lloyd's iteration, seeded with k-means++ over the house +// generator: the first centre is drawn uniformly among the samples and +// every later centre with probability proportional to the squared +// distance to the nearest centre already chosen, so a well-separated +// cluster cannot be left unseeded except against odds its separation +// sets. The fit is deterministic for a given generator state. +// +// Convergence: the sweeps stop once no centre moves by more than +// kmeansTolerance (scaled, as the constant documents) or after +// kmeansMaxIterations sweeps; Converged names which happened. +// +// Empty clusters: a sweep can strand a centre with no members. The +// stranded centre then moves onto the sample farthest from its +// current centre, the sample that the objective can improve on most; +// if several centres strand in one sweep, each takes the next +// farthest unused sample. The rule is deterministic and keeps every +// cluster live. +// +// A nil generator, a non-finite sample, k below 1 or k above n is an +// error. +func KMeans(g *core.Generator, x *core.Array, k int) (*KMeansResult, error) { + const name = "KMeans" + if g == nil { + return nil, base.Errf("%s: the generator is nil", name) + } + data, n, d, err := clusterReadSample(name, x) + if err != nil { + return nil, err + } + if k < 1 { + return nil, base.Errf("%s: k must be at least 1, got %d", name, k) + } + if k > n { + return nil, base.Errf("%s: k = %d exceeds the %d samples", name, k, n) + } + seeds := kmeansSeeds(g, data, n, d, k) + return kmeansLloyd(data, n, d, k, seeds) +} + +// kmeansSeeds draws k initial centres by the k-means++ rule: the +// first uniformly, each later one with probability proportional to the +// squared distance to the nearest centre already drawn. The weighted +// draw walks the cumulative weights against one uniform, the same +// inverse-transform the package's discrete draws use. +func kmeansSeeds(g *core.Generator, data []float64, n, d, k int) [][]float64 { + seeds := make([][]float64, k) + first := min(int(g.Unit()*float64(n)), n-1) + seeds[0] = append([]float64(nil), data[first*d:first*d+d]...) + nearest := make([]float64, n) + for i := range n { + nearest[i] = sqDistance(data[i*d:i*d+d], seeds[0]) + } + for c := 1; c < k; c++ { + total := 0.0 + for i := range n { + total += nearest[i] + } + // A total of zero means every sample sits exactly on a chosen + // centre: nothing left to separate, and any further centre is + // arbitrary. The position repeats, and the Lloyd loop's + // empty-cluster rule resolves the tie. + if total == 0 { + seeds[c] = append([]float64(nil), data[0:d]...) + } else { + target := g.Unit() * total + walk := 0.0 + pick := n - 1 + for i := range n { + walk += nearest[i] + if walk >= target { + pick = i + break + } + } + seeds[c] = append([]float64(nil), data[pick*d:pick*d+d]...) + } + for i := range n { + if dist := sqDistance(data[i*d:i*d+d], seeds[c]); dist < nearest[i] { + nearest[i] = dist + } + } + } + return seeds +} + +// kmeansChosen reports whether the point equals one of the centres +// already drawn, the guard of the degenerate all-coincident fallback. +func kmeansChosen(centres [][]float64, point []float64) bool { + for _, c := range centres { + same := true + for j, v := range c { + if v != point[j] { + same = false + break + } + } + if same { + return true + } + } + return false +} + +// sqDistance returns the squared Euclidean distance between two +// points of equal length. +func sqDistance(a, b []float64) float64 { + total := 0.0 + for j, v := range a { + diff := v - b[j] + total += diff * diff + } + return total +} + +// kmeansLloyd runs the assignment and update sweeps from the given +// centres to the documented tolerance, applying the empty-cluster +// rule on the way. It is the shared engine of the public entry point +// and of the seeding comparisons the tests make. +func kmeansLloyd(data []float64, n, d, k int, centres [][]float64) (*KMeansResult, error) { + // The scale that makes the tolerance relative: the largest + // absolute coordinate in the sample, floored at 1. + spread := 1.0 + for _, v := range data { + if s := math.Abs(v); s > spread { + spread = s + } + } + labels := make([]int, n) + converged := false + iterations := kmeansMaxIterations + // The update sweep's accumulators live outside the iteration: a sweep + // clears them instead of allocating, and the fold below then adds the + // same rows into the same slots in the same order. + counts := make([]int, k) + sums := make([][]float64, k) + for c := range k { + sums[c] = make([]float64, d) + } + for iter := 1; iter <= kmeansMaxIterations; iter++ { + for i := range n { + best, bestDist := 0, math.Inf(1) + for c := range k { + if dist := sqDistance(data[i*d:i*d+d], centres[c]); dist < bestDist { + best, bestDist = c, dist + } + } + labels[i] = best + } + // The update sweep, with the empty-cluster rule: a centre + // whose cluster stranded takes the sample farthest from its + // own centre, skipping the samples earlier stranded centres + // in the same sweep already claimed. + clear(counts) + for c := range k { + clear(sums[c]) + } + for i := range n { + counts[labels[i]]++ + row := sums[labels[i]] + for j := range d { + row[j] += data[i*d+j] + } + } + worst := 0.0 + for c := range k { + switch { + case counts[c] > 0: + for j := range d { + moved := sums[c][j] / float64(counts[c]) + if m := math.Abs(moved - centres[c][j]); m > worst { + worst = m + } + centres[c][j] = moved + } + default: + // The farthest sample from this centre: the search is + // for a maximum, so the running best starts at -Inf. + farthest, farDist := -1, math.Inf(-1) + for i := range n { + if labels[i] == c || kmeansChosen(centres, data[i*d:i*d+d]) { + continue + } + if dist := sqDistance(data[i*d:i*d+d], centres[c]); dist > farDist { + farthest, farDist = i, dist + } + } + if farthest < 0 { + return nil, base.Errf("KMeans: cluster %d cannot be repopulated: the sample holds fewer distinct points than k", c) + } + copy(centres[c], data[farthest*d:farthest*d+d]) + worst = math.Inf(1) + } + } + if worst <= kmeansTolerance*spread { + converged = true + iterations = iter + break + } + } + return &KMeansResult{ + Centres: centres, + Labels: labels, + Inertia: kmeansInertia(data, n, d, centres, labels), + Iterations: iterations, + Converged: converged, + }, nil +} + +// kmeansInertia sums the squared distances of every sample to its +// assigned centre. +func kmeansInertia(data []float64, n, d int, centres [][]float64, labels []int) float64 { + total := 0.0 + for i := range n { + total += sqDistance(data[i*d:i*d+d], centres[labels[i]]) + } + return total +} + +// gmmMaxIterations caps the EM sweeps, with the same reporting +// convention kmeansMaxIterations sets. +const gmmMaxIterations = 200 + +// gmmTolerance is the EM convergence tolerance on the relative change +// of the log likelihood: the sweeps stop once it moves by less than +// the tolerance scaled by 1 + |log likelihood|, past which no +// parameter the M step can produce moves the objective materially. +const gmmTolerance = 1e-9 + +// gmmCovarianceFloor is the regularisation floor of the M step: after +// every update, gmmCovarianceFloor of the data's mean per-dimension +// variance is added to the covariance diagonal. A component that +// collapses onto a single sample computes a singular covariance and a +// Cholesky factor that does not exist; the floor keeps the factor +// alive at a width far below any sample's spread, so it regularises +// the algebra without moving any fit the data actually supports. +const gmmCovarianceFloor = 1e-6 + +// GaussianMixtureResult carries the fit of a Gaussian mixture. +type GaussianMixtureResult struct { + // Components is the fitted component count. + Components int + // Weights are the mixture weights, in component order, summing to + // 1. + Weights []float64 + // Means holds one mean vector per component. + Means [][]float64 + // Covariances holds one d-by-d covariance per component, + // row-major, in component order. + Covariances [][]float64 + // Responsibilities are the final E-step posteriors, n rows of + // Components entries each, row-major: P(component j | sample i). + Responsibilities []float64 + // LogLikelihood is the maximised observed-data log likelihood. + LogLikelihood float64 + // BIC is -2·LogLikelihood + p·ln n at the fitted parameters, + // p the count of free parameters. + BIC float64 + // Iterations counts the EM sweeps; Converged reports whether the + // log likelihood settled under gmmTolerance before the budget. + Iterations int + Converged bool + // BICGrid is filled by GaussianMixtureBIC only: the BIC of every + // fit on the component grid 1..len(BICGrid), in grid order. The + // plain fit leaves it nil. + BICGrid []float64 +} + +// GaussianMixture fits a mixture of `components` multivariate normals +// over the rows of x by expectation maximisation, built on the MVN +// machinery of mvn.go: every component density goes through the +// Cholesky factor of its covariance. +// +// The E step runs in log space: each sample's responsibilities come +// from a log-sum-exp normalisation guarded by the row maximum, so a +// sample sitting far outside every component underflows to a clean 0/1 +// split rather than to a NaN. +// +// The M step regularises every covariance by the documented +// gmmCovarianceFloor before the next factorisation. +// +// Initialisation is k-means over the same generator: the partitions +// seed the weights, means and covariances, so the fit is deterministic +// for a given seed and every component starts populated. A component +// whose total responsibility collapses below 1e-10 (the floor at work +// on a stray sample) keeps its previous parameters and a floored +// weight rather than dividing by zero; the renormalised weights keep +// the mixture a mixture. +// +// The sample must be rank 2 and finite, the component count at least 1 +// and at most the sample size, and the sample at least two rows. +func GaussianMixture(g *core.Generator, x *core.Array, components int) (*GaussianMixtureResult, error) { + const name = "GaussianMixture" + if g == nil { + return nil, base.Errf("%s: the generator is nil", name) + } + data, n, d, err := clusterReadSample(name, x) + if err != nil { + return nil, err + } + if n < 2 { + return nil, base.Errf("%s: needs at least two samples, got %d", name, n) + } + if components < 1 { + return nil, base.Errf("%s: the component count must be at least 1, got %d", name, components) + } + if components > n { + return nil, base.Errf("%s: %d components exceed the %d samples", name, components, n) + } + weights, means, covs, serr := gmmSeed(g, data, n, d, components) + if serr != nil { + return nil, base.Errf("%s: %w", name, serr) + } + result, err := gmmEM(data, n, d, components, weights, means, covs) + if err != nil { + return nil, err + } + result.BIC = gaussianMixtureBIC(result.LogLikelihood, components, n, d) + return result, nil +} + +// GaussianMixtureBIC fits mixtures over the component grid 1 to +// maxComponents and keeps the one with the lowest BIC. The grid and +// the winner travel in the result: BICGrid holds every fit's BIC in +// grid order and Components names the selected count. Ties resolve to +// the smaller model, the parsimony the penalty exists to buy. +func GaussianMixtureBIC(g *core.Generator, x *core.Array, maxComponents int) (*GaussianMixtureResult, error) { + const name = "GaussianMixtureBIC" + if g == nil { + return nil, base.Errf("%s: the generator is nil", name) + } + data, n, d, err := clusterReadSample(name, x) + if err != nil { + return nil, err + } + if n < 2 { + return nil, base.Errf("%s: needs at least two samples, got %d", name, n) + } + if maxComponents < 1 { + return nil, base.Errf("%s: the largest component count must be at least 1, got %d", name, maxComponents) + } + if maxComponents > n { + return nil, base.Errf("%s: %d components exceed the %d samples", name, maxComponents, n) + } + grid := make([]float64, maxComponents) + best := -1 + bestBIC := math.Inf(1) + var bestResult *GaussianMixtureResult + for k := 1; k <= maxComponents; k++ { + weights, means, covs, serr := gmmSeed(g, data, n, d, k) + if serr != nil { + return nil, serr + } + result, err := gmmEM(data, n, d, k, weights, means, covs) + if err != nil { + return nil, base.Errf("%s: the fit with %d components failed (%w)", name, k, err) + } + grid[k-1] = gaussianMixtureBIC(result.LogLikelihood, k, n, d) + // Strictly better only: a tie keeps the smaller model. + if grid[k-1] < bestBIC { + bestBIC = grid[k-1] + best = k + bestResult = result + } + } + bestResult.BICGrid = grid + bestResult.BIC = grid[best-1] + return bestResult, nil +} + +// gaussianMixtureBIC assembles -2·logL + p·ln n, with p the free +// parameters of a k-component mixture: k-1 weights, k·d mean entries +// and k·d(d+1)/2 covariance entries. +func gaussianMixtureBIC(logL float64, k, n, d int) float64 { + p := (k - 1) + k*d + k*d*(d+1)/2 + return -2*logL + float64(p)*math.Log(float64(n)) +} + +// clusterReadSample validates x as a finite real rank-2 sample and +// returns it flattened row-major with its shape. +func clusterReadSample(name string, x *core.Array) ([]float64, int, int, error) { + if x.NDim() != 2 { + return nil, 0, 0, base.Errf("%s: the sample must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, 0, 0, base.Errf("%s: complex inputs are not supported", name) + } + if err := checkFinite(name, "the sample", x); err != nil { + return nil, 0, 0, err + } + n, d := x.Shape()[0], x.Shape()[1] + if d < 1 { + return nil, 0, 0, base.Errf("%s: the sample needs at least one column", name) + } + data := make([]float64, n*d) + if fs := rawFloats(x); fs != nil { + copy(data, fs) + } else { + for i := range data { + data[i] = x.FloatAt(i) + } + } + return data, n, d, nil +} + +// gmmSeed initialises the mixture from the k-means partitions of the +// same generator: each component takes its cluster's weight, mean and +// covariance, floored exactly as the M step floors, so the first E +// step factors live covariances even for a cluster that collapsed onto +// one sample. +func gmmSeed(g *core.Generator, data []float64, n, d, k int) ([]float64, [][]float64, [][]float64, error) { + seeds := kmeansSeeds(g, data, n, d, k) + fit, err := kmeansLloyd(data, n, d, k, seeds) + if err != nil { + // A sample with fewer distinct points than k drives the + // empty-cluster rule out of answers: the caller turns this + // into the entry point's own named refusal. + return nil, nil, nil, err + } + meanVariance := clusterMeanVariance(data, n, d) + weights := make([]float64, k) + means := make([][]float64, k) + covs := make([][]float64, k) + counts := make([]int, k) + for i := range n { + counts[fit.Labels[i]]++ + } + for c := range k { + weights[c] = math.Max(float64(counts[c]), 1e-10) + mean := make([]float64, d) + if counts[c] > 0 { + for i := range n { + if fit.Labels[i] != c { + continue + } + for j := range d { + mean[j] += data[i*d+j] + } + } + for j := range d { + mean[j] /= float64(counts[c]) + } + } + cov := make([]float64, d*d) + for i := range n { + if fit.Labels[i] != c { + continue + } + for a := range d { + for b := range d { + cov[a*d+b] += (data[i*d+a] - mean[a]) * (data[i*d+b] - mean[b]) + } + } + } + den := math.Max(float64(counts[c])-1, 1) + for a := range d { + for b := range d { + cov[a*d+b] /= den + } + cov[a*d+a] += gmmCovarianceFloor * meanVariance + } + weights[c] /= float64(n) + means[c] = mean + covs[c] = cov + } + // The floored weights renormalised: the mixture must sum to 1 even + // when a component started empty and took the floor. + total := 0.0 + for _, w := range weights { + total += w + } + for c := range k { + weights[c] /= total + } + return weights, means, covs, nil +} + +// clusterMeanVariance returns the mean per-dimension variance of the +// sample, the scale the covariance floor is measured in. +func clusterMeanVariance(data []float64, n, d int) float64 { + if n < 2 { + return 1 + } + total := 0.0 + for j := range d { + mean := 0.0 + for i := range n { + mean += data[i*d+j] + } + mean /= float64(n) + v := 0.0 + for i := range n { + diff := data[i*d+j] - mean + v += diff * diff + } + total += v / float64(n-1) + } + return total / float64(d) +} + +// gmmParallelMinPoints is the sample count one worker must carry before +// the expectation sweep splits across goroutines: a point costs the +// component count in density evaluations and another in exponentials, +// so a shorter chunk is cheaper on the calling goroutine than in a +// pool. +const gmmParallelMinPoints = 32 + +// gmmEM runs the expectation maximisation sweeps from the seeded +// parameters to the documented tolerance. +func gmmEM(data []float64, n, d, k int, weights []float64, means [][]float64, covs [][]float64) (*GaussianMixtureResult, error) { + const name = "GaussianMixture" + resp := make([]float64, n*k) + logL := math.Inf(-1) + converged := false + iterations := gmmMaxIterations + // The sweep's buffers live outside the sweep: one flat covariance + // copy the factor step refills, one factor per component the + // factorisation clears and refills, and one row of per-point + // normalisers the log-likelihood fold reads, so a sweep allocates + // none of them. The per-component mean and covariance slices the M + // step accumulates into are the parameter storage itself: each is + // cleared before its pass, so the accumulation sees the zero state + // a fresh allocation carried and the result carries the last + // sweep's values in the same slices. + covVals := make([]float64, d*d) + factors := make([][][]float64, k) + for c := range k { + rows := make([][]float64, d) + for i := range d { + rows[i] = make([]float64, d) + } + factors[c] = rows + if len(means[c]) != d { + means[c] = make([]float64, d) + } + if len(covs[c]) != d*d { + covs[c] = make([]float64, d*d) + } + } + consts := make([]float64, k) + logWeights := make([]float64, k) + counts := make([]float64, k) + centred := make([]float64, d) + logNorms := make([]float64, n) + // The covariance floor's scale depends on the sample alone, and the + // sample never moves, so it is measured once for the whole fit. + meanVariance := clusterMeanVariance(data, n, d) + for iter := 1; iter <= gmmMaxIterations; iter++ { + // E step: the components are factored once per sweep, and each + // factor contributes its log determinant and the density's + // normalising constant as one value, read once per component + // instead of once per point. The component's log weight is + // likewise constant across the sweep. The sample loop that + // follows then carries only the quadratic form. + for c := range k { + copy(covVals, covs[c]) + if err := mvnCholeskyFlatInto(name, covVals, factors[c], d); err != nil { + return nil, base.Errf("%s: component %d failed to factor (%w)", name, c, err) + } + l := factors[c] + logDet := 0.0 + for i := range d { + logDet += math.Log(l[i][i]) + } + consts[c] = -0.5*float64(d)*math.Log(2*math.Pi) - logDet + logWeights[c] = math.Log(weights[c]) + } + // A point owns its responsibility row and its own normaliser + // entry and nothing else, so the sample loop splits across + // goroutines with no lock and no shared scratch: the worker's two + // buffers are overwritten per point. + engine.ParallelMin(n, gmmParallelMinPoints, func(start, end int) { + logps := make([]float64, k) + solve := make([]float64, d) + for i := start; i < end; i++ { + point := data[i*d : i*d+d] + for c := range k { + logps[c] = logWeights[c] + mvnLogDensitySolve(point, means[c], factors[c], d, consts[c], solve) + } + rowMax := math.Inf(-1) + for _, lp := range logps { + if lp > rowMax { + rowMax = lp + } + } + // A floored weight can park a component at -Inf; the + // max-guarded sum survives it and the row still + // normalises. logps is reused for the exponentials once + // the row maximum has served: the posterior is the + // quotient of the stored value and the row total, so + // one exponential per component serves the row and the + // divisor is the total in [1, k] instead of an exponent + // compounded with the logarithm's rounding. + total := 0.0 + for c := range k { + e := math.Exp(logps[c] - rowMax) + logps[c] = e + total += e + } + logNorm := rowMax + math.Log(total) + for c := range k { + resp[i*k+c] = logps[c] / total + } + logNorms[i] = logNorm + } + }) + // The observed-data log likelihood folds the normalisers in + // ascending point order on the calling goroutine: the chain is + // the one the serial sweep accumulated, so the split above never + // moves a bit of it. + nextLogL := 0.0 + for _, logNorm := range logNorms { + nextLogL += logNorm + } + move := nextLogL - logL + logL = nextLogL + if math.Abs(move) <= gmmTolerance*(1+math.Abs(logL)) { + converged = true + iterations = iter + break + } + // M step: responsibilities to weights, means and floored + // covariances. A component whose total responsibility + // collapses under 1e-10 keeps its previous parameters under a + // floored weight: dividing by zero would poison the sweep, + // and the floor lets the next E step give the component + // another chance. + clear(counts) + for i := range n { + for c := range k { + counts[c] += resp[i*k+c] + } + } + totalWeight := 0.0 + for c := range k { + if counts[c] < 1e-10 { + weights[c] = 1e-10 + } else { + weights[c] = counts[c] / float64(n) + } + totalWeight += weights[c] + } + for c := range k { + weights[c] /= totalWeight + } + for c := range k { + if counts[c] < 1e-10 { + continue + } + mean := means[c] + clear(mean) + for i := range n { + r := resp[i*k+c] + for j := range d { + mean[j] += r * data[i*d+j] + } + } + for j := range d { + mean[j] /= counts[c] + } + cov := covs[c] + clear(cov) + // The centred row is built once per sample and read by every + // entry the covariance accumulates: the subtraction is the one + // the product below performed per entry, and the remaining + // factorisation of the term, (r·ca)·cb added to the entry, + // keeps its own order and shape. + for i := range n { + r := resp[i*k+c] + point := data[i*d : i*d+d] + for a := range d { + centred[a] = point[a] - mean[a] + } + for a := range d { + scaled := r * centred[a] + row := cov[a*d : a*d+d] + for b, cb := range centred { + row[b] += scaled * cb + } + } + } + for a := range d { + for b := range d { + cov[a*d+b] /= counts[c] + } + cov[a*d+a] += gmmCovarianceFloor * meanVariance + } + } + } + return &GaussianMixtureResult{ + Components: k, + Weights: weights, + Means: means, + Covariances: covs, + Responsibilities: resp, + LogLikelihood: logL, + Iterations: iterations, + Converged: converged, + }, nil +} + +// mvnLogDensitySolve evaluates the log density of one point under a +// component already factored, with the caller's normalising constant +// (the log determinant already folded in) and its own forward-solve +// scratch. Every entry of solve is written before it is read, so the +// scratch carries no state across calls; the constant is the identical +// expression the caller would otherwise recompute from the same factor. +func mvnLogDensitySolve(point, mean []float64, l [][]float64, d int, normConst float64, solve []float64) float64 { + for i := range d { + total := point[i] - mean[i] + for j := range i { + total -= l[i][j] * solve[j] + } + solve[i] = total / l[i][i] + } + quad := 0.0 + for i := range d { + quad += solve[i] * solve[i] + } + return normConst - 0.5*quad +} diff --git a/stats/cluster_test.go b/stats/cluster_test.go new file mode 100644 index 0000000..3f0d010 --- /dev/null +++ b/stats/cluster_test.go @@ -0,0 +1,436 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "cmp" + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestKMeansRecoversBlobCentres draws two well-separated Gaussian +// blobs over the house generator and requires the fit to recover both +// centres to tolerance and to partition every sample onto its own +// blob's label. +func TestKMeansRecoversBlobCentres(t *testing.T) { + g := core.NewGenerator(11) + const perBlob = 80 + sample := core.New(core.Float, 2*perBlob, 2) + truth := [2][2]float64{{0, 0}, {8, 8}} + for i := range 2 * perBlob { + blob := 0 + if i >= perBlob { + blob = 1 + } + for j := range 2 { + sample.RawFloats()[i*2+j] = truth[blob][j] + 0.6*g.NormalUnit() + } + } + res, err := KMeans(core.NewGenerator(3), sample, 2) + if err != nil { + t.Fatalf("KMeans: %v", err) + } + if !res.Converged { + t.Fatalf("the fit did not converge in %d sweeps", res.Iterations) + } + // Each true centre must be matched by a fitted one within the + // sampling noise of the blob mean: 0.6/sqrt(80) is about 0.07, so + // 0.4 leaves eight standard errors of room. + for _, want := range truth { + best := math.Inf(1) + for _, c := range res.Centres { + if d := sqDistance(c, want[:]); d < best { + best = d + } + } + if best > 0.4*0.4 { + t.Fatalf("no fitted centre within 0.4 of (%g, %g): closest %.4g", + want[0], want[1], math.Sqrt(best)) + } + } + // The labels partition the blobs: every sample of a blob carries + // one label, and the two blobs carry different labels. + first, second := res.Labels[0], res.Labels[perBlob] + if first == second { + t.Fatalf("both blobs share label %d", first) + } + for i := range perBlob { + if res.Labels[i] != first { + t.Fatalf("blob 0 sample %d carries label %d, want %d", i, res.Labels[i], first) + } + if res.Labels[perBlob+i] != second { + t.Fatalf("blob 1 sample %d carries label %d, want %d", perBlob+i, res.Labels[perBlob+i], second) + } + } +} + +// TestKMeansSeedingBeatsFixedBadSeeding pins the value of k-means++ +// seeding as a property of the seeding, not a flaky inequality. The +// fixture is three one-dimensional blobs at 0, 30 and 60 with width +// 0.4 and k = 3, so the only good fit plants one centre per blob and +// costs about 19 units of inertia, while a seeding that plants all +// three centres inside blob 0 converges to a stable bad optimum: the +// leftmost centres stay to split blob 0, the rightmost is dragged to +// 45, midway between the two unclaimed blobs, and both pay 40 points +// times 15² each, about 18,000. The bad optimum is stable because the +// boundary between the shared centre at 45 and the blob-0 centres +// falls near 22, and no point of blob 1 crosses it. k-means++ cannot +// share that fate by anything the geometry allows: after the first +// (uniform) centre, every later centre is drawn with probability +// proportional to squared distance from the nearest one already +// chosen, and every point of an unclaimed blob outweighs every point +// of a claimed one by roughly (15/0.8)² ≳ 350, so keeping two centres +// in one blob is odds of millions to one against per run. Over the ten +// seeds below the outcome is deterministic, and every run must land +// the small objective. +func TestKMeansSeedingBeatsFixedBadSeeding(t *testing.T) { + g := core.NewGenerator(23) + const perBlob = 40 + data := make([]float64, 3*perBlob) + for i := range perBlob { + data[i] = 0.4 * g.NormalUnit() + data[perBlob+i] = 30 + 0.4*g.NormalUnit() + data[2*perBlob+i] = 60 + 0.4*g.NormalUnit() + } + // The bad seeding: all three centres inside blob 0. + bad, err := kmeansLloyd(data, 3*perBlob, 1, 3, [][]float64{{-0.4}, {0}, {0.4}}) + if err != nil { + t.Fatalf("bad-seeded fit: %v", err) + } + if bad.Inertia < 1000 { + t.Fatalf("the bad seeding escaped its own trap: inertia %.4g", bad.Inertia) + } + for seed := int64(101); seed < 111; seed++ { + res, err := KMeans(core.NewGenerator(seed), mustFloats(t, data, 3*perBlob, 1), 3) + if err != nil { + t.Fatalf("KMeans seed %d: %v", seed, err) + } + if res.Inertia > bad.Inertia/100 { + t.Fatalf("seed %d: k-means++ inertia %.4g is within a factor of 100 of the bad seeding's %.4g", + seed, res.Inertia, bad.Inertia) + } + } +} + +// TestKMeansDeterministicUnderSeed runs the same sample twice under +// the same seed and requires bit-identical fits, then checks that the +// validation refusals all refuse. +func TestKMeansDeterministicUnderSeedAndValidation(t *testing.T) { + g := core.NewGenerator(9) + sample := core.New(core.Float, 30, 2) + for i := range 30 { + sample.RawFloats()[i*2] = g.Unit() + sample.RawFloats()[i*2+1] = g.Unit() + } + first, err := KMeans(core.NewGenerator(4), sample, 3) + if err != nil { + t.Fatalf("KMeans: %v", err) + } + second, err := KMeans(core.NewGenerator(4), sample, 3) + if err != nil { + t.Fatalf("KMeans repeat: %v", err) + } + for c, centre := range first.Centres { + if centre[0] != second.Centres[c][0] || centre[1] != second.Centres[c][1] { + t.Fatalf("seeded fit moved between runs: centre %d", c) + } + } + if first.Inertia != second.Inertia { + t.Fatalf("inertia moved between runs: %.17g against %.17g", first.Inertia, second.Inertia) + } + // An int sample takes the widening accessor path. + ints, err := core.FromInts([]int64{0, 0, 10, 12, 20, 18}, 3, 2) + if err != nil { + t.Fatalf("FromInts: %v", err) + } + if _, err := KMeans(core.NewGenerator(1), ints, 2); err != nil { + t.Fatalf("KMeans over an int sample: %v", err) + } + // Refusals: nil generator, wrong rank, complex, non-finite, and a + // k outside the sample. + if _, err := KMeans(nil, sample, 2); err == nil || !strings.Contains(err.Error(), "generator is nil") { + t.Fatalf("nil generator: got %v, want the nil-generator refusal", err) + } + if _, err := KMeans(core.NewGenerator(1), core.New(core.Float, 10), 2); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("rank-1 sample: got %v, want the rank refusal", err) + } + if _, err := KMeans(core.NewGenerator(1), core.New(core.Complex, 4, 1), 2); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("complex sample: got %v, want the complex refusal", err) + } + cloud := core.New(core.Float, 6, 2) + cloud.RawFloats()[3] = math.NaN() + if _, err := KMeans(core.NewGenerator(1), cloud, 2); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("non-finite sample: got %v, want the non-finite refusal", err) + } + fine := core.New(core.Float, 6, 2) + if _, err := KMeans(core.NewGenerator(1), fine, 0); err == nil || !strings.Contains(err.Error(), "at least 1") { + t.Fatalf("k = 0: got %v, want the k floor refusal", err) + } + if _, err := KMeans(core.NewGenerator(1), fine, 7); err == nil || !strings.Contains(err.Error(), "exceeds") { + t.Fatalf("k above the sample size: got %v, want the exceeds refusal", err) + } + if _, err := KMeans(core.NewGenerator(1), core.New(core.Float, 6, 0), 2); err == nil || !strings.Contains(err.Error(), "one column") { + t.Fatalf("a zero-column sample: got %v, want the column refusal", err) + } +} + +// TestKMeansEmptyClusterRule drives the documented empty-cluster +// handling directly: a seeding that leaves one centre nobody's nearest +// must see that centre move onto the sample farthest from it, and a +// sample with fewer distinct points than k must be refused rather than +// silently merged. +func TestKMeansEmptyClusterRule(t *testing.T) { + // Points 0, 1, 2 with centres 5 and 4.9: every point is nearer + // 4.9, so the centre at 5 strands and takes the farthest sample, + // 0, while the other centre keeps {1, 2} and settles on their + // mean. The settled fit is the split {0} and {1, 2}. + res, err := kmeansLloyd([]float64{0, 1, 2}, 3, 1, 2, [][]float64{{5}, {4.9}}) + if err != nil { + t.Fatalf("kmeansLloyd: %v", err) + } + if !res.Converged { + t.Fatal("the empty-cluster fit did not converge") + } + if res.Labels[0] == res.Labels[1] || res.Labels[2] == res.Labels[0] { + t.Fatalf("labels %v do not split {0} from {1, 2}", res.Labels) + } + want := [][]float64{{0}, {1.5}} + for c := range 2 { + if res.Centres[c][0] != want[c][0] { + t.Fatalf("centre %d = %.4g, want %.4g", c, res.Centres[c][0], want[c][0]) + } + } + // Three coincident points and k = 2: no repopulation is possible, + // and the refusal says so. + if _, err := KMeans(core.NewGenerator(1), mustFloats(t, []float64{3, 3, 3}, 3, 1), 2); err == nil { + t.Fatal("a sample with fewer distinct points than k was accepted") + } +} + +// TestGaussianMixtureRecoversComponents draws a long sample from a +// two-component univariate mixture of known weights, means and +// variances and requires the EM fit to recover all three, matched in +// mean order. +func TestGaussianMixtureRecoversComponents(t *testing.T) { + g := core.NewGenerator(5) + const n = 6000 + truth := []struct { + weight float64 + mean float64 + variance float64 + }{{0.3, -2, 0.25}, {0.7, 3, 2.25}} + sample := core.New(core.Float, n, 1) + for i := range n { + c := 0 + if g.Unit() >= truth[0].weight { + c = 1 + } + sample.RawFloats()[i] = truth[c].mean + math.Sqrt(truth[c].variance)*g.NormalUnit() + } + res, err := GaussianMixture(core.NewGenerator(5), sample, 2) + if err != nil { + t.Fatalf("GaussianMixture: %v", err) + } + if !res.Converged { + t.Fatalf("EM did not converge in %d sweeps", res.Iterations) + } + if res.Components != 2 { + t.Fatalf("Components = %d, want 2", res.Components) + } + // The weights must form a distribution and the responsibilities + // must be posteriors. + total := 0.0 + for _, w := range res.Weights { + total += w + } + if math.Abs(total-1) > 1e-9 { + t.Fatalf("weights sum to %.17g", total) + } + for i := range n { + rowSum := 0.0 + for c := range 2 { + rowSum += res.Responsibilities[i*2+c] + } + if math.Abs(rowSum-1) > 1e-9 { + t.Fatalf("responsibilities of sample %d sum to %.17g", i, rowSum) + } + } + // Match components by mean order, the labelling that survives the + // mixture's label symmetry. + order := []int{0, 1} + slices.SortFunc(order, func(a, b int) int { + return cmp.Compare(res.Means[a][0], res.Means[b][0]) + }) + for pos, want := range truth { + c := order[pos] + if math.Abs(res.Weights[c]-want.weight) > 0.05 { + t.Fatalf("component %d: weight %.4g, want about %.2f", c, res.Weights[c], want.weight) + } + if math.Abs(res.Means[c][0]-want.mean) > 0.1 { + t.Fatalf("component %d: mean %.4g, want about %.2f", c, res.Means[c][0], want.mean) + } + got := res.Covariances[c][0] + if math.Abs(got-want.variance) > 0.15 { + t.Fatalf("component %d: variance %.4g, want about %.2f", c, got, want.variance) + } + } +} + +// TestGaussianMixtureSeparatedResponsibilities fits a mixture whose +// components sit seven standard deviations apart and requires every +// posterior to be 0/1 to tolerance: a sample drawn from one component +// is orders of magnitude more likely under it than under the other. +func TestGaussianMixtureSeparatedResponsibilities(t *testing.T) { + g := core.NewGenerator(17) + const perComponent = 400 + sample := core.New(core.Float, 2*perComponent, 1) + for i := range perComponent { + sample.RawFloats()[i] = -7 + g.NormalUnit() + sample.RawFloats()[perComponent+i] = 7 + g.NormalUnit() + } + res, err := GaussianMixture(core.NewGenerator(8), sample, 2) + if err != nil { + t.Fatalf("GaussianMixture: %v", err) + } + // Match the fitted components to the drawn sides by mean order. + low := 0 + if res.Means[1][0] < res.Means[0][0] { + low = 1 + } + for i := range 2 * perComponent { + want := 0.0 + if i < perComponent { + want = 1 + } + got := res.Responsibilities[i*2+low] + if math.Abs(got-want) > 1e-6 { + t.Fatalf("sample %d: responsibility %.3g, want %.0f", i, got, want) + } + } +} + +// TestGaussianMixtureBICSelectsTrueCount draws from a two-component +// mixture and requires the BIC sweep over one to four components to +// put the minimum at the true count. +func TestGaussianMixtureBICSelectsTrueCount(t *testing.T) { + g := core.NewGenerator(29) + const n = 1200 + sample := core.New(core.Float, n, 1) + for i := range n { + mean := -3.0 + if i >= n/2 { + mean = 3 + } + sample.RawFloats()[i] = mean + g.NormalUnit() + } + res, err := GaussianMixtureBIC(core.NewGenerator(6), sample, 4) + if err != nil { + t.Fatalf("GaussianMixtureBIC: %v", err) + } + if len(res.BICGrid) != 4 { + t.Fatalf("BIC grid holds %d entries, want 4", len(res.BICGrid)) + } + if res.Components != 2 { + t.Fatalf("BIC selected %d components, want 2 (grid %v)", res.Components, res.BICGrid) + } + best := slices.Min(res.BICGrid) + if math.Abs(best-res.BIC) > 1e-9 { + t.Fatalf("the reported BIC %.4g is not the grid minimum %.4g", res.BIC, best) + } +} + +// TestGaussianMixtureValidation checks the refusals of both entry +// points, and that the same seed reproduces the same fit bit for bit. +func TestGaussianMixtureValidation(t *testing.T) { + g := core.NewGenerator(2) + sample := core.New(core.Float, 40, 2) + for i := range 40 { + sample.RawFloats()[i*2] = g.Unit() + sample.RawFloats()[i*2+1] = g.Unit() + } + if _, err := GaussianMixture(nil, sample, 2); err == nil || !strings.Contains(err.Error(), "generator is nil") { + t.Fatalf("nil generator: got %v, want the nil-generator refusal", err) + } + if _, err := GaussianMixture(core.NewGenerator(1), core.New(core.Float, 40), 2); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("rank-1 sample: got %v, want the rank refusal", err) + } + if _, err := GaussianMixture(core.NewGenerator(1), core.New(core.Float, 1, 2), 1); err == nil || !strings.Contains(err.Error(), "at least two samples") { + t.Fatalf("a one-sample fit: got %v, want the sample floor refusal", err) + } + if _, err := GaussianMixture(core.NewGenerator(1), sample, 0); err == nil || !strings.Contains(err.Error(), "component count") { + t.Fatalf("zero components: got %v, want the component-count refusal", err) + } + if _, err := GaussianMixture(core.NewGenerator(1), sample, 41); err == nil || !strings.Contains(err.Error(), "exceed") { + t.Fatalf("components above the sample size: got %v, want the exceeds refusal", err) + } + sick := core.New(core.Float, 4, 2) + sick.RawFloats()[2] = math.Inf(1) + if _, err := GaussianMixture(core.NewGenerator(1), sick, 1); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("non-finite sample: got %v, want the non-finite refusal", err) + } + if _, err := GaussianMixture(core.NewGenerator(1), core.New(core.Complex, 4, 1), 1); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("complex sample: got %v, want the complex refusal", err) + } + // A degenerate sample with no spread at all has no Gaussian to + // fit: the covariance collapses and the factorisation refuses it + // by name, in both entry points. + if _, err := GaussianMixture(core.NewGenerator(1), mustFloats(t, []float64{2, 2}, 2, 1), 1); err == nil || !strings.Contains(err.Error(), "positive definite") { + t.Fatalf("a zero-variance sample: got %v, want the factorisation refusal", err) + } + if _, err := GaussianMixtureBIC(core.NewGenerator(1), mustFloats(t, []float64{2, 2}, 2, 1), 2); err == nil || !strings.Contains(err.Error(), "positive definite") { + t.Fatalf("BIC: a zero-variance sample: got %v, want the factorisation refusal", err) + } + if _, err := GaussianMixtureBIC(nil, sample, 3); err == nil || !strings.Contains(err.Error(), "generator is nil") { + t.Fatalf("BIC: nil generator: got %v, want the nil-generator refusal", err) + } + if _, err := GaussianMixtureBIC(core.NewGenerator(1), sample, 0); err == nil || !strings.Contains(err.Error(), "largest component count") { + t.Fatalf("BIC: empty grid: got %v, want the component-count refusal", err) + } + if _, err := GaussianMixtureBIC(core.NewGenerator(1), core.New(core.Float, 40), 3); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("BIC: rank-1 sample: got %v, want the rank refusal", err) + } + if _, err := GaussianMixtureBIC(core.NewGenerator(1), core.New(core.Float, 1, 2), 3); err == nil || !strings.Contains(err.Error(), "at least two samples") { + t.Fatalf("BIC: a one-sample fit: got %v, want the sample floor refusal", err) + } + if _, err := GaussianMixtureBIC(core.NewGenerator(1), sample, 41); err == nil || !strings.Contains(err.Error(), "exceed") { + t.Fatalf("BIC: grid above the sample size: got %v, want the exceeds refusal", err) + } + // Determinism: the same seed must reproduce the same likelihood + // and the same weights. + a, err := GaussianMixture(core.NewGenerator(12), sample, 2) + if err != nil { + t.Fatalf("GaussianMixture: %v", err) + } + b, err := GaussianMixture(core.NewGenerator(12), sample, 2) + if err != nil { + t.Fatalf("GaussianMixture repeat: %v", err) + } + if a.LogLikelihood != b.LogLikelihood { + t.Fatalf("seeded fit moved: %.17g against %.17g", a.LogLikelihood, b.LogLikelihood) + } + for c := range 2 { + if a.Weights[c] != b.Weights[c] { + t.Fatalf("weights moved between seeded runs: %.17g against %.17g", a.Weights[c], b.Weights[c]) + } + } +} + +// TestGaussianMixtureConstantSampleRefused pins the refusal a sample +// with no spread earns when k exceeds its distinct points: the +// k-means seed cannot repopulate every centre, and the mixture +// refuses instead of panicking. +func TestGaussianMixtureConstantSampleRefused(t *testing.T) { + g := core.NewGenerator(1) + x, err := core.FromFloats([]float64{2, 2}, 2, 1) + if err != nil { + t.Fatal(err) + } + if _, err := GaussianMixture(g, x, 2); err == nil { + t.Fatal("a constant sample with k = 2 was accepted") + } +} diff --git a/stats/contingency.go b/stats/contingency.go new file mode 100644 index 0000000..438e853 --- /dev/null +++ b/stats/contingency.go @@ -0,0 +1,355 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strconv" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Contingency tables: the exact and the asymptotic tests of a +// cross-classification and the effect size that summarises one. The +// counts enter as an array and are read through the widening +// accessor, so a table carried in any real dtype answers the same +// numbers; a negative or fractional entry is refused, because a count +// it is not. The exact answers travel no further than lgamma and the +// package's own incomplete beta, and every p-value is computed in log +// space from the ratios of table probabilities, so a far-from-uniform +// table cannot lose its answer to an underflow the absolute +// probabilities would have suffered. + +// Alternative names the alternative hypothesis a directional test +// evaluates. +type Alternative int + +const ( + // TwoSided is the alternative of a difference in either direction. + TwoSided Alternative = iota + // Less is the alternative that the first named quantity is smaller. + Less + // Greater is the alternative that the first named quantity is larger. + Greater +) + +// String returns the name of the alternative, the spelling the error +// messages carry. +func (a Alternative) String() string { + switch a { + case TwoSided: + return "two-sided" + case Less: + return "less" + case Greater: + return "greater" + } + return "alternative(" + strconv.Itoa(int(a)) + ")" +} + +// maxFisherSupport caps the hypergeometric support FisherExactTest +// enumerates. Each support point costs a pair of lgamma evaluations, +// and a table whose margins put millions of possible tables between +// the support's ends has long left the exact regime the test exists +// for: there the asymptotic ChiSquareIndependence is the honest +// answer, and the exact test refuses with the cap named rather than +// spend seconds enumerating a sum whose terms have all underflowed. +const maxFisherSupport = 4_000_000 + +// maxExactCount bounds a single count of the exact tests' tables. The +// enumeration walks the support integer by integer, and the format +// holds those integers stepwise only below 2^52: past it every second +// one is missing, the walk's increment stops advancing and the loop +// cannot leave its first table. One count at most 2^51 keeps every +// margin the four counts can sum to below 2^52 and every support +// point an exact integer. +const maxExactCount = 1 << 51 + +// FisherExactTest runs Fisher's exact test on a 2×2 table of counts, +// the classical answer to a cross-classification too sparse for the +// χ² approximation. The rows are the two samples and the columns the +// two outcomes, and the test conditions on the margins: under +// independence the first cell follows the hypergeometric law, and the +// p-value sums the probabilities of the tables at least as extreme as +// the observed one, exactly. The two-sided sum collects every table +// whose probability does not exceed the observed one's, so an +// asymmetric table sums its two tails unequally, as the exact +// definition demands. +// +// Returns the p-value for the given alternative and the sample odds +// ratio (a·d)/(b·c): +∞ where a zero cell sits against a full one, +// NaN where the table carries no contrast at all (a zero margin), +// which answers a p-value of 1, the one table its margins allow. +// +// Refuses a table that is not 2×2, complex input, a non-finite, +// negative or fractional count, a count past the maxExactCount the +// format holds the integers of, an unknown alternative and a support +// wider than maxFisherSupport. +func FisherExactTest(table *core.Array, alternative Alternative) (pValue, oddsRatio float64, err error) { + const name = "FisherExactTest" + a, b, c, d, err := squareCounts(name, table) + if err != nil { + return 0, 0, err + } + switch alternative { + case TwoSided, Less, Greater: + default: + return 0, 0, base.Errf("%s: unknown alternative %s", name, alternative) + } + oddsRatio = a * d / (b * c) + row1, row2 := a+b, c+d + col1, col2 := a+c, b+d + if row1 == 0 || row2 == 0 || col1 == 0 || col2 == 0 { + // A zero margin admits exactly one table, the observed one: the + // conditioning has removed every degree of freedom and no + // arrangement is more extreme. + return 1, oddsRatio, nil + } + lo := math.Max(0, row1-col2) + hi := math.Min(row1, col1) + if hi-lo+1 > maxFisherSupport { + return 0, 0, base.Errf("%s: the margins span a support of %.0f tables, above the %d cap; use ChiSquareIndependence", + name, hi-lo+1, maxFisherSupport) + } + // The probability of the table with first cell x, up to the factor + // the margins fix: log C(col1, x) + log C(col2, row1−x). The + // common constant (the normalising choice of margins) never enters + // a ratio, so it is left unevaluated, and the two sums are folded + // in log space: on a lopsided table the ratio the observed table's + // probability bears to the mode's overflows the format many times + // over, and the quotients the direct sums would then divide (Inf by + // Inf on a one-sided tail, a finite extreme by an infinite total on + // the two-sided one) answer NaN or a flushed zero where the true + // tail is representable. The log-sum-exp fold keeps every + // accumulator finite, so the p-value is one quotient of two + // well-scaled logarithms at the end. + logTerm := func(x float64) float64 { + return lchoose(col1, x) + lchoose(col2, row1-x) + } + observed := logTerm(a) + slack := math.Log1p(1e-9) + logTotal := math.Inf(-1) + logExtreme := math.Inf(-1) + for x := lo; x <= hi; x++ { + lt := logTerm(x) + logTotal = logAddExp(logTotal, lt) + counts := false + switch alternative { + case TwoSided: + // A table counts as at least as extreme when its probability + // does not exceed the observed one's; the slack absorbs the + // rounding of equal-probability tables onto either side of 1. + counts = lt <= observed+slack + case Less: + counts = x <= a + case Greater: + counts = x >= a + } + if counts { + logExtreme = logAddExp(logExtreme, lt) + } + } + return min(1, math.Exp(logExtreme-logTotal)), oddsRatio, nil +} + +// logAddExp folds one logarithm into another without leaving log space: +// the max guard keeps the exponential's argument at most zero, so the +// fold never overflows however far apart the two terms sit, and a term +// at negative infinity (a sum still empty) passes through untouched. +func logAddExp(a, b float64) float64 { + if a == math.Inf(-1) { + return b + } + if b == math.Inf(-1) { + return a + } + m := math.Max(a, b) + return m + math.Log1p(math.Exp(-math.Abs(a-b))) +} + +// lchoose returns log C(n, k) for integral arguments held as floats. +// The lgamma sign is never negative on these arguments, all at least 1. +func lchoose(n, k float64) float64 { + la, _ := math.Lgamma(n + 1) + lb, _ := math.Lgamma(k + 1) + lc, _ := math.Lgamma(n - k + 1) + return la - lb - lc +} + +// ChiSquareIndependence runs Pearson's χ² test of independence on an +// r×c table of counts: the expected count under independence is the +// row total times the column total over the grand total, the +// statistic is Σ(O−E)²/E over the cells, and the p-value is the upper +// tail of χ² on (r−1)(c−1) degrees of freedom. A row or column that +// totals zero carries no information and only divides by zero; it is +// refused by name rather than absorbed. +// +// The approximation is the asymptotic one: with expected counts in +// the single digits the exact FisherExactTest is the honest test, and +// this one is the answer once the table is dense enough for χ² to +// hold. Refuses fewer than two rows or columns, complex input, a +// non-finite, negative or fractional count, and a zero row or column +// total. +func ChiSquareIndependence(table *core.Array) (chi2 float64, df int, pValue float64, err error) { + const name = "ChiSquareIndependence" + stat, _, nrows, ncols, serr := contingencyStatistic(name, table) + if serr != nil { + return 0, 0, 0, serr + } + chi2 = stat + df = (nrows - 1) * (ncols - 1) + pValue, err = GammaUpper(float64(df)/2, chi2/2) + if err != nil { + return 0, 0, 0, base.Errf("%s: %w", name, err) + } + return chi2, df, pValue, nil +} + +// CramersV returns Cramér's V for an r×c table, the χ² statistic of +// independence rescaled into [0, 1] by the sample size and the +// smaller margin: V = √(χ²/(n·(min(r, c)−1))). Zero means the counts +// sit exactly on independence, one means each row leans on a single +// column. The input contract is ChiSquareIndependence's own. +func CramersV(table *core.Array) (float64, error) { + const name = "CramersV" + chi2, total, rows, cols, err := contingencyStatistic(name, table) + if err != nil { + return 0, err + } + return math.Sqrt(chi2 / (total * float64(min(rows, cols)-1))), nil +} + +// McNemarTest runs the exact McNemar test on a 2×2 table of paired +// counts: the diagonal holds the agreements and the off-diagonal pair +// (b, c) the two directions of disagreement, and under the null the +// disagreements split evenly, so min(b, c) follows the binomial law +// with m = b + c trials at p = ½. The two-sided p-value doubles the +// smaller tail, evaluated through the package's incomplete beta +// rather than an m-term sum, so a table with millions of discordant +// pairs costs the same as one with a dozen. A table with no +// discordant pair has nothing to test and answers 1. Refuses a table +// that is not 2×2, complex input, a non-finite, negative or fractional +// count, and a count past the maxExactCount the format holds the +// integers of. +func McNemarTest(table *core.Array) (pValue float64, err error) { + const name = "McNemarTest" + _, b, c, _, err := squareCounts(name, table) + if err != nil { + return 0, err + } + m := b + c + if m == 0 { + return 1, nil + } + k := min(b, c) + tail, err := BetaIncomplete(0.5, m-k, k+1) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + return min(1, 2*tail), nil +} + +// squareCounts reads a 2×2 count table into its cells, the shared +// reader of FisherExactTest and McNemarTest. +func squareCounts(name string, table *core.Array) (a, b, c, d float64, err error) { + if table.NDim() != 2 || table.Shape()[0] != 2 || table.Shape()[1] != 2 { + return 0, 0, 0, 0, base.Errf("%s: the table must be 2×2, got shape %s", name, base.ShapeText(table.Shape())) + } + if table.Dtype() == core.Complex { + return 0, 0, 0, 0, base.Errf("%s: complex tables are not supported", name) + } + counts, err := countValues(name, table) + if err != nil { + return 0, 0, 0, 0, err + } + for _, v := range counts { + if v >= maxExactCount { + return 0, 0, 0, 0, base.Errf("%s: count %g exceeds 2^51, past which float64 no longer holds the integers stepwise and the exact enumeration cannot run", name, v) + } + } + return counts[0], counts[1], counts[2], counts[3], nil +} + +// contingencyStatistic computes the Pearson statistic of an r×c count +// table together with the total the effect sizes normalise by, the +// shared body of ChiSquareIndependence and CramersV. +func contingencyStatistic(name string, table *core.Array) (chi2, total float64, rows, cols int, err error) { + if table.NDim() != 2 { + return 0, 0, 0, 0, base.Errf("%s: the table must be rank 2, got shape %s", name, base.ShapeText(table.Shape())) + } + rows, cols = table.Shape()[0], table.Shape()[1] + if rows < 2 || cols < 2 { + return 0, 0, 0, 0, base.Errf("%s: the table needs at least two rows and two columns, got %d×%d", name, rows, cols) + } + if table.Dtype() == core.Complex { + return 0, 0, 0, 0, base.Errf("%s: complex tables are not supported", name) + } + counts, err := countValues(name, table) + if err != nil { + return 0, 0, 0, 0, err + } + rowTotals := make([]float64, rows) + colTotals := make([]float64, cols) + for i := range rows { + for j := range cols { + v := counts[i*cols+j] + rowTotals[i] += v + colTotals[j] += v + total += v + } + } + if total == 0 { + return 0, 0, 0, 0, base.Errf("%s: the table totals zero, there is nothing to test", name) + } + for i, r := range rowTotals { + if r == 0 { + return 0, 0, 0, 0, base.Errf("%s: row %d totals zero and carries no information", name, i+1) + } + } + for j, c := range colTotals { + if c == 0 { + return 0, 0, 0, 0, base.Errf("%s: column %d totals zero and carries no information", name, j+1) + } + } + chi2 = 0.0 + for i := range rows { + for j := range cols { + expected := rowTotals[i] * colTotals[j] / total + d := counts[i*cols+j] - expected + chi2 += d * d / expected + } + } + return chi2, total, rows, cols, nil +} + +// countValues reads a real array as counts, refusing complex input, +// a non-finite entry, a negative entry and a fractional entry: the +// accessor walk answers the same widened values for every real dtype, +// so a table carried in any of them reads identically. +func countValues(name string, table *core.Array) ([]float64, error) { + n := table.Len() + counts := make([]float64, n) + if fs := rawFloats(table); fs != nil { + // A rebased view's payload may run past its own count: only the + // visible elements are counts. + copy(counts, fs[:n]) + } else { + for i := range counts { + counts[i] = table.FloatAt(i) + } + } + for i, v := range counts { + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: count [%d] is not finite (%g)", name, i, v) + } + if v < 0 { + return nil, base.Errf("%s: count [%d] is negative (%g)", name, i, v) + } + if v != math.Trunc(v) { + return nil, base.Errf("%s: count [%d] is fractional (%g)", name, i, v) + } + } + return counts, nil +} diff --git a/stats/contingency_test.go b/stats/contingency_test.go new file mode 100644 index 0000000..8cdc78f --- /dev/null +++ b/stats/contingency_test.go @@ -0,0 +1,513 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "math/big" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The contingency tests against exact referents. The hypergeometric +// and the central binomial probabilities are rational numbers, so the +// Fisher and McNemar p-values are compared against a big.Rat +// evaluation of the same definition, not against a tolerance quote. + +// exactFactorial returns n! as a big.Int. +func exactFactorial(n int64) *big.Int { + out := big.NewInt(1) + for i := int64(2); i <= n; i++ { + out.Mul(out, big.NewInt(i)) + } + return out +} + +// exactChoose returns C(n, k) as a big.Int. +func exactChoose(n, k int64) *big.Int { + return new(big.Int).Div(exactFactorial(n), new(big.Int).Mul(exactFactorial(k), exactFactorial(n-k))) +} + +// exactFisherTwoSided evaluates the two-sided Fisher p-value of the +// 2×2 table (a b / c d) exactly: the sum of every table probability +// that does not exceed the observed one's, over the hypergeometric +// support the margins induce. +func exactFisherTwoSided(a, b, c, d int64) *big.Rat { + row1, row2 := a+b, c+d + col1 := a + c + col2 := b + d + lo := max(int64(0), row1-col2) + hi := min(row1, col1) + den := new(big.Rat).SetInt(exactChoose(row1+row2, row1)) + observed := new(big.Rat) + terms := make([]*big.Rat, 0, hi-lo+1) + for x := lo; x <= hi; x++ { + p := new(big.Rat).SetFrac( + new(big.Int).Mul(exactChoose(col1, x), exactChoose(col2, row1-x)), + den.Num()) + terms = append(terms, p) + if x == a { + observed.Set(p) + } + } + sum := new(big.Rat) + for _, p := range terms { + if p.Cmp(observed) <= 0 { + sum.Add(sum, p) + } + } + return sum +} + +// exactMcNemar evaluates the exact McNemar p-value of the discordant +// pair (b, c) exactly: twice the binomial(m, ½) tail at min(b, c). +func exactMcNemar(b, c int64) *big.Rat { + m := b + c + if m == 0 { + return big.NewRat(1, 1) + } + k := min(b, c) + den := new(big.Int).Lsh(big.NewInt(1), uint(m)) + tail := new(big.Rat) + for i := int64(0); i <= k; i++ { + term := new(big.Rat).SetFrac(exactChoose(m, i), den) + tail.Add(tail, term) + } + // The doubled tail is clamped at 1 exactly as the entry point + // clamps it: at the middle pair min(b, c) = m/2 the sum already + // covers the median and doubling overruns. + if tail.Cmp(big.NewRat(1, 2)) > 0 { + return big.NewRat(1, 1) + } + return tail.Add(tail, tail) +} + +// table2x2 builds a 2×2 float array from four counts. +func table2x2(t *testing.T, a, b, c, d float64) *core.Array { + t.Helper() + tab, err := core.FromFloats([]float64{a, b, c, d}, 2, 2) + if err != nil { + t.Fatal(err) + } + return tab +} + +func TestFisherExactAgainstRationalReferent(t *testing.T) { + tables := [][4]int64{ + {12, 5, 7, 10}, + {8, 2, 1, 5}, + {3, 1, 1, 3}, + {1, 9, 9, 1}, + {0, 6, 6, 0}, + {17, 0, 0, 23}, + {35, 12, 41, 9}, + } + for _, tab := range tables { + p, or, err := FisherExactTest(table2x2(t, float64(tab[0]), float64(tab[1]), float64(tab[2]), float64(tab[3])), TwoSided) + if err != nil { + t.Fatalf("FisherExactTest(%v): %v", tab, err) + } + want := exactFisherTwoSided(tab[0], tab[1], tab[2], tab[3]) + wantF, _ := want.Float64() + if math.Abs(p-wantF) > 1e-12 { + t.Fatalf("FisherExactTest(%v) two-sided = %.16f, want the exact %.16f", tab, p, wantF) + } + wantOR := float64(tab[0]) * float64(tab[3]) / (float64(tab[1]) * float64(tab[2])) + if math.IsNaN(wantOR) { + if !math.IsNaN(or) { + t.Fatalf("FisherExactTest(%v): odds ratio %g, want NaN", tab, or) + } + } else if or != wantOR { + t.Fatalf("FisherExactTest(%v): odds ratio %g, want %g", tab, or, wantOR) + } + } +} + +func TestFisherExactTeaTablePin(t *testing.T) { + // The classical 3/1 versus 1/3 table: the two-sided answer is + // 34/70 by hand, every table but the central one being at most as + // probable as the observed one. + p, or, err := FisherExactTest(table2x2(t, 3, 1, 1, 3), TwoSided) + if err != nil { + t.Fatal(err) + } + if math.Abs(p-34.0/70.0) > 1e-12 { + t.Fatalf("the tea table p = %.16f, want 34/70 = %.16f", p, 34.0/70.0) + } + if or != 9 { + t.Fatalf("the tea table odds ratio = %g, want 9", or) + } +} + +func TestFisherExactOneSidedTails(t *testing.T) { + tab := table2x2(t, 12, 5, 7, 10) + less, _, err := FisherExactTest(tab, Less) + if err != nil { + t.Fatal(err) + } + greater, _, err := FisherExactTest(tab, Greater) + if err != nil { + t.Fatal(err) + } + // The observed first cell (12) sits above its expectation 9.5: the + // greater tail must be the shorter one. The two-sided sum carries + // no ordering guarantee against the tails, because it collects + // only the tables at most as probable as the observed one while a + // one-sided tail sweeps the whole arm including the mode. + if !(greater < less) { + t.Fatalf("the greater tail %g does not fall below the less tail %g on an upper-tail table", greater, less) + } + // The greater tail, recomputed from the exact binomial + // coefficients over the support's upper arm: the margins are + // row1 = 17, row2 = 17, col1 = 19, col2 = 15, the support runs + // 2..17, and the tail collects x = 12..17. + sum := new(big.Rat) + den := new(big.Rat).SetInt(exactChoose(34, 17)) + for x := int64(12); x <= 17; x++ { + term := new(big.Rat).SetInt(new(big.Int).Mul(exactChoose(19, x), exactChoose(15, 17-x))) + term.Quo(term, den) + sum.Add(sum, term) + } + want, _ := sum.Float64() + if math.Abs(greater-want) > 1e-12 { + t.Fatalf("the greater tail = %.16f, want the exact %.16f", greater, want) + } +} + +func TestFisherExactDegenerateMargins(t *testing.T) { + // A zero margin admits one table only: p = 1 with the odds ratio + // the formula speaks. + p, or, err := FisherExactTest(table2x2(t, 0, 5, 0, 7), TwoSided) + if err != nil { + t.Fatal(err) + } + if p != 1 { + t.Fatalf("a zero-row table answered p = %g, want 1", p) + } + if !math.IsNaN(or) { + t.Fatalf("a zero-row table answered odds ratio %g, want NaN", or) + } + // A zero cell against a full one inside living margins: the odds + // ratio goes infinite and the two-sided p-value is the observed + // corner's own probability, 1/210 by the exact coefficients. + p, or, err = FisherExactTest(table2x2(t, 6, 0, 0, 4), TwoSided) + if err != nil { + t.Fatal(err) + } + if math.Abs(p-1.0/210.0) > 1e-12 { + t.Fatalf("the corner table answered p = %.16f, want 1/210 = %.16f", p, 1.0/210.0) + } + if or != math.Inf(1) { + t.Fatalf("the corner table answered odds ratio %g, want +Inf", or) + } +} + +func TestFisherExactRefusals(t *testing.T) { + if _, _, err := FisherExactTest(mustStatArray(t, []float64{1, 2, 3}, 3), TwoSided); err == nil { + t.Fatal("a rank-1 table was accepted") + } + if _, _, err := FisherExactTest(table2x2(t, 1, 2, 3, 4), Alternative(7)); err == nil { + t.Fatal("an unknown alternative was accepted") + } + if _, _, err := FisherExactTest(table2x2(t, 1, 2, 3, -4), TwoSided); err == nil { + t.Fatal("a negative count was accepted") + } + if _, _, err := FisherExactTest(table2x2(t, 1, 2.5, 3, 4), TwoSided); err == nil { + t.Fatal("a fractional count was accepted") + } + if _, _, err := FisherExactTest(table2x2(t, 1, math.NaN(), 3, 4), TwoSided); err == nil { + t.Fatal("a NaN count was accepted") + } + // Margins wide enough to push the support past the cap: the refusal + // names the cap instead of enumerating millions of tables. + if _, _, err := FisherExactTest(table2x2(t, 3e6, 3e6, 3e6, 3e6), TwoSided); err == nil { + t.Fatal("an over-cap support was accepted") + } +} + +// mustStatArray builds a float array for the refusal probes. +func mustStatArray(t *testing.T, vals []float64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +func TestMcNemarAgainstRationalReferent(t *testing.T) { + pairs := [][2]int64{ + {5, 15}, + {0, 10}, + {10, 0}, + {7, 7}, + {1, 1}, + {0, 0}, + {23, 41}, + } + for _, pr := range pairs { + p, err := McNemarTest(table2x2(t, 10, float64(pr[0]), float64(pr[1]), 12)) + if err != nil { + t.Fatalf("McNemarTest(%v): %v", pr, err) + } + wantF, _ := exactMcNemar(pr[0], pr[1]).Float64() + if math.Abs(p-wantF) > 1e-12 { + t.Fatalf("McNemarTest(b=%d, c=%d) = %.16f, want the exact %.16f", pr[0], pr[1], p, wantF) + } + } + // Large discordant counts: the incomplete-beta route must keep + // answering well past where an m-term sum would have become too + // expensive to run. + b, c := 2000.0, 3000.0 + p, err := McNemarTest(table2x2(t, 100, b, c, 100)) + if err != nil { + t.Fatal(err) + } + if !(0 < p && p < 1) { + t.Fatalf("McNemarTest(2000, 3000) = %g, want a probability in (0, 1)", p) + } + pSmall, err := McNemarTest(table2x2(t, 100, 0, 900, 100)) + if err != nil { + t.Fatal(err) + } + if !(pSmall < 1e-15) { + t.Fatalf("McNemarTest(0, 900) = %g, want a vanishing tail", pSmall) + } +} + +func TestMcNemarRefusals(t *testing.T) { + if _, err := McNemarTest(mustStatArray(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)); err == nil { + t.Fatal("a 2×3 table was accepted") + } + if _, err := McNemarTest(table2x2(t, 1, 2, 3, math.Inf(1))); err == nil { + t.Fatal("an infinite count was accepted") + } +} + +func TestChiSquareIndependenceTwoByTwoIdentity(t *testing.T) { + // For a 2×2 table the definitional sum collapses algebraically to + // n(ad−bc)² over the product of the four margins, an independent + // route the sum must reproduce bit for bit. + tables := [][4]float64{ + {12, 7, 5, 16}, + {10, 20, 20, 30}, + {1, 9, 9, 1}, + {153, 71, 44, 232}, + } + for _, tab := range tables { + chi2, df, p, err := ChiSquareIndependence(table2x2(t, tab[0], tab[1], tab[2], tab[3])) + if err != nil { + t.Fatalf("ChiSquareIndependence(%v): %v", tab, err) + } + if df != 1 { + t.Fatalf("ChiSquareIndependence(%v): df = %d, want 1", tab, df) + } + n := tab[0] + tab[1] + tab[2] + tab[3] + want := n * math.Pow(tab[0]*tab[3]-tab[1]*tab[2], 2) / + ((tab[0] + tab[1]) * (tab[2] + tab[3]) * (tab[0] + tab[2]) * (tab[1] + tab[3])) + if math.Abs(chi2-want) > 1e-12*math.Max(1, want) { + t.Fatalf("ChiSquareIndependence(%v): χ² = %.16f, want the closed form %.16f", tab, chi2, want) + } + if !(0 < p && p < 1) { + t.Fatalf("ChiSquareIndependence(%v): p = %g is not a probability in (0, 1)", tab, p) + } + // Cramér's V for a 2×2 table is |φ| = |ad−bc| over the square + // root of the same product of margins. + v, err := CramersV(table2x2(t, tab[0], tab[1], tab[2], tab[3])) + if err != nil { + t.Fatal(err) + } + phi := math.Abs(tab[0]*tab[3]-tab[1]*tab[2]) / + math.Sqrt((tab[0]+tab[1])*(tab[2]+tab[3])*(tab[0]+tab[2])*(tab[1]+tab[3])) + if math.Abs(v-phi) > 1e-12 { + t.Fatalf("CramersV(%v) = %.16f, want |φ| = %.16f", tab, v, phi) + } + } +} + +func TestChiSquareIndependenceLargerTable(t *testing.T) { + // A 3×3 table with the statistic recomputed from the definition + // through an independent walk over the expected counts. + counts := []float64{ + 21, 8, 3, + 9, 15, 11, + 4, 12, 30, + } + tab := mustStatArray(t, counts, 3, 3) + chi2, df, p, err := ChiSquareIndependence(tab) + if err != nil { + t.Fatal(err) + } + if df != 4 { + t.Fatalf("df = %d, want 4", df) + } + rows := make([]float64, 3) + cols := make([]float64, 3) + total := 0.0 + for i := range 3 { + for j := range 3 { + rows[i] += counts[i*3+j] + cols[j] += counts[i*3+j] + total += counts[i*3+j] + } + } + want := 0.0 + for i := range 3 { + for j := range 3 { + e := rows[i] * cols[j] / total + want += math.Pow(counts[i*3+j]-e, 2) / e + } + } + if math.Abs(chi2-want) > 1e-12*math.Max(1, want) { + t.Fatalf("χ² = %.16f, want %.16f", chi2, want) + } + v, err := CramersV(tab) + if err != nil { + t.Fatal(err) + } + wantV := math.Sqrt(want / (total * 2)) + if math.Abs(v-wantV) > 1e-12 { + t.Fatalf("Cramér's V = %.16f, want %.16f", v, wantV) + } + // Perfect independence answers χ² = 0, p = 1 and V = 0. + independent := []float64{10, 20, 10, 20, 40, 20, 10, 20, 10} + chi2, _, p, err = ChiSquareIndependence(mustStatArray(t, independent, 3, 3)) + if err != nil { + t.Fatal(err) + } + if chi2 != 0 || p != 1 { + t.Fatalf("a perfectly independent table answered χ² = %g, p = %g", chi2, p) + } +} + +func TestChiSquareIndependenceRefusals(t *testing.T) { + if _, _, _, err := ChiSquareIndependence(mustStatArray(t, []float64{1, 2, 3, 4}, 4)); err == nil { + t.Fatal("a rank-1 table was accepted") + } + if _, _, _, err := ChiSquareIndependence(mustStatArray(t, []float64{1, 2, 3}, 3, 1)); err == nil { + t.Fatal("a single-column table was accepted") + } + if _, _, _, err := ChiSquareIndependence(table2x2(t, 1, 0, 3, 0)); err == nil { + t.Fatal("a zero-column margin was accepted") + } + if _, _, _, err := ChiSquareIndependence(table2x2(t, 1, 2, 3, 2.5)); err == nil { + t.Fatal("a fractional count was accepted") + } + if _, err := CramersV(table2x2(t, 1, 2, 3, -4)); err == nil { + t.Fatal("a negative count was accepted") + } + if _, _, _, err := ChiSquareIndependence(table2x2(t, 0, 0, 0, 0)); err == nil { + t.Fatal("an all-zero table was accepted") + } +} + +func TestContingencyDtypeAgreement(t *testing.T) { + // A table carried in an integer dtype must answer the same numbers + // as the float array of the widened values. + vals := []float64{12, 5, 7, 10} + fv, _, err := FisherExactTest(table2x2(t, vals[0], vals[1], vals[2], vals[3]), TwoSided) + if err != nil { + t.Fatal(err) + } + iv := make([]int64, len(vals)) + for i, v := range vals { + iv[i] = int64(v) + } + ia, err := core.FromInts(iv, 2, 2) + if err != nil { + t.Fatal(err) + } + pv, or, err := FisherExactTest(ia, TwoSided) + if err != nil { + t.Fatal(err) + } + if pv != fv || or != 120.0/35.0 { + t.Fatalf("the int table answered (%g, %g), want (%g, %g)", pv, or, fv, 120.0/35.0) + } +} + +func TestFisherExactLopsidedOverflow(t *testing.T) { + // A corner table with living margins: the observed table sits + // hundreds of nats below the hypergeometric mode, so the ratio of + // the mode's probability to the observed one's overflows the + // format. The one-sided tail over the whole support is the certain + // event and must answer exactly 1, the opposite tail and the + // two-sided sum underflow to the correctly rounded zero, and none + // of the three may turn into the NaN a quotient of infinities + // produces. + greater, _, err := FisherExactTest(table2x2(t, 0, 1000, 1000, 0), Greater) + if err != nil { + t.Fatal(err) + } + if greater != 1 { + t.Fatalf("the certain Greater tail answered %g, want 1", greater) + } + less, _, err := FisherExactTest(table2x2(t, 0, 1000, 1000, 0), Less) + if err != nil { + t.Fatal(err) + } + twoSided, _, err := FisherExactTest(table2x2(t, 0, 1000, 1000, 0), TwoSided) + if err != nil { + t.Fatal(err) + } + if math.IsNaN(less) || math.IsNaN(twoSided) { + t.Fatalf("the far tails answered NaN (%g, %g)", less, twoSided) + } + if less != 0 || twoSided != 0 { + // 1/C(2000,1000) and 2/C(2000,1000) sit near 1e-600: the + // correctly rounded float64 of both is 0. + t.Fatalf("the underflowed tails answered (%g, %g), want (0, 0)", less, twoSided) + } + // A table whose two-sided answer lands in the representable + // subnormal window while the mode's ratio still overflows: the + // exact referent is 2/C(1040,520), near 2e-313. + p, _, err := FisherExactTest(table2x2(t, 0, 520, 520, 0), TwoSided) + if err != nil { + t.Fatal(err) + } + want, _ := new(big.Rat).SetFrac( + big.NewInt(2), + exactChoose(1040, 520)).Float64() + if !(p > 0) || math.Abs(p-want) > 1e-318 { + t.Fatalf("the subnormal two-sided tail answered %.16e, want the exact %.16e", p, want) + } +} + +func TestFisherExactUnknownAlternativeBeforeDegenerateMargins(t *testing.T) { + // The alternative is refused whatever the margins: a zero-margin + // table must not short-circuit the refusal behind a p-value of 1. + if _, _, err := FisherExactTest(table2x2(t, 0, 5, 0, 7), Alternative(7)); err == nil { + t.Fatal("an unknown alternative was accepted on a zero-margin table") + } +} + +func TestFisherExactCountBeyondTheFormatsIntegers(t *testing.T) { + // A count past 2^51 can push the margins past 2^52, where float64 + // holds only every second integer: the support walk's increment + // stops advancing and the enumeration cannot leave its first table. + // The table below answers a support of one rounded point, so the + // cap check cannot catch it; the refusal must, and it must come + // back rather than hang. + var err error + pinWatchdog(t, 10*time.Second, "FisherExactTest on [[1e18,0],[0,3]]", func() { + _, _, err = FisherExactTest(table2x2(t, 1e18, 0, 0, 3), TwoSided) + }) + if err == nil { + t.Fatal("a count past the format's integers was accepted") + } + if _, err := McNemarTest(table2x2(t, 1, 1e18, 1, 1)); err == nil { + t.Fatal("McNemarTest accepted a count past the format's integers") + } + // The last count the format holds stepwise stays answerable: the + // bound refuses nothing a legal enumeration can run on. + p, _, err := FisherExactTest(table2x2(t, float64(maxExactCount-2), 0, 0, 3), TwoSided) + if err != nil { + t.Fatalf("a count at the bound's inside was refused: %v", err) + } + if !(p >= 0 && p <= 1) { + t.Fatalf("the boundary table answered p = %g", p) + } +} diff --git a/stats/deep_tail_pin_test.go b/stats/deep_tail_pin_test.go new file mode 100644 index 0000000..3bb5b2d --- /dev/null +++ b/stats/deep_tail_pin_test.go @@ -0,0 +1,154 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestQuantileDeepTails pins the deep-left-tail quantiles that +// the old 200-pass bisection silently mis-answered and the 1e300 +// bracket floors refused: the answers must be within a rounding of the +// exact values, never an unconverged midpoint with a nil error. +func TestQuantileDeepTails(t *testing.T) { + v, err := ExponentialQuantile(1e-100, 1) + if err != nil { + t.Fatalf("ExponentialQuantile: %v", err) + } + if relErr(v, 1e-100) > 1e-6 { + t.Fatalf("ExponentialQuantile(1e-100, 1) = %g, want 1e-100", v) + } + v, err = ChiSquareQuantile(1e-30, 1) + if err != nil { + t.Fatalf("ChiSquareQuantile: %v", err) + } + // The exact χ²(1) deep lower tail: x = z² with Φ(z) = 0.5 + q/2, + // so z ≈ (q/2)/φ(0) and x ≈ (π/2)·q². + want := math.Pi / 2 * 1e-60 + if relErr(v, want) > 1e-4 { + t.Fatalf("ChiSquareQuantile(1e-30, 1) = %g, want %g", v, want) + } + v, err = GammaQuantile(1e-100, 0.5, 0.5) + if err != nil { + t.Fatalf("GammaQuantile: %v", err) + } + if v < 0.5*1e-200 || v > 2*1e-200 { + t.Fatalf("GammaQuantile(1e-100, 0.5, 0.5) = %g, want ≈ 0.785e-200 (χ²(1)/2 at q²)", v) + } + // The denormal floor: the smallest representable q still answers + // the smallest representable scale, or refuses loudly; never a + // silent wrong number. + if _, err = ExponentialQuantile(5e-324, 1); err != nil { + t.Fatalf("ExponentialQuantile(5e-324): %v", err) + } +} + +func relErr(got, want float64) float64 { + d := math.Abs(got - want) + if want == 0 { + return d + } + return d / math.Abs(want) +} + +// TestNormalQuantileExtremeTail pins that a q below 2⁻⁵³, where +// 1−q rounds to exactly 1, still answers through the accurate upper +// tail instead of the bogus "q = 1" refusal. +func TestNormalQuantileExtremeTail(t *testing.T) { + v, err := NormalQuantile(1e-17) + if err != nil { + t.Fatalf("NormalQuantile(1e-17): %v", err) + } + // Φ(−8.2907496456... ) = 1e-17 to the digits that matter here. + if relErr(0.5*math.Erfc(-v/math.Sqrt2), 1e-17) > 1e-6 { + t.Fatalf("NormalQuantile(1e-17) = %g, tail = %g", v, 0.5*math.Erfc(-v/math.Sqrt2)) + } + v, err = StudentTQuantile(1e-17, 5) + if err != nil { + t.Fatalf("StudentTQuantile(1e-17, 5): %v", err) + } + tail, err := studentTUpperTail(-v, 5) + if err != nil { + t.Fatalf("studentTUpperTail: %v", err) + } + // The quantile mirrors the one-sided tail: P(T > −t) = q. + if relErr(tail, 1e-17) > 1e-6 { + t.Fatalf("StudentTQuantile(1e-17, 5) = %g, one-sided tail = %g, want %g", v, tail, 1e-17) + } + // The heavy tails stay representable far past where t² overflows: + // Cauchy answers 1/(πq) exactly, df = 2 answers 1/sqrt(2q), where + // the incomplete-beta argument used to saturate to a phantom zero + // and clamp the quantile at the overflow wall. + vc, err := StudentTQuantile(1e-200, 1) + if err != nil { + t.Fatalf("StudentTQuantile(1e-200, 1): %v", err) + } + if want := -1 / (math.Pi * 1e-200); relErr(vc, want) > 1e-6 { + t.Fatalf("StudentTQuantile(1e-200, 1) = %g, want %g", vc, want) + } + v2, err := StudentTQuantile(1e-160, 2) + if err != nil { + t.Fatalf("StudentTQuantile(1e-160, 2): %v", err) + } + if want := -1 / math.Sqrt(2*1e-160); relErr(v2, want) > 1e-6 { + t.Fatalf("StudentTQuantile(1e-160, 2) = %g, want %g", v2, want) + } +} + +// TestRegressionConstantResponse pins R² = 1 for a constant +// response reproduced exactly, where the 1 − 0/0 form reported NaN +// with a nil error. +func TestRegressionConstantResponse(t *testing.T) { + x := mustFloats(t, []float64{1, 0, 1, 1, 1, 2, 1, 3}, 4, 2) + y := mustFloats(t, []float64{5, 5, 5, 5}, 4) + res, err := LinearRegression(x, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + if res.RSquared != 1 || res.AdjustedRSquared != 1 { + t.Fatalf("constant response: R² = %v, adj = %v, want 1 and 1", res.RSquared, res.AdjustedRSquared) + } + w := mustFloats(t, []float64{1, 2, 1, 1}, 4) + wres, err := WeightedLinearRegression(x, y, w) + if err != nil { + t.Fatalf("WeightedLinearRegression: %v", err) + } + if wres.RSquared != 1 || wres.AdjustedRSquared != 1 { + t.Fatalf("weighted constant response: R² = %v, adj = %v, want 1 and 1", wres.RSquared, wres.AdjustedRSquared) + } +} + +// TestHistogramFullRange pins the refusal of a sample holding +// both float extremes, whose edges would be ±Inf and whose counts +// silently collapsed into bin 0. +func TestHistogramFullRange(t *testing.T) { + a := mustFloats(t, []float64{-math.MaxFloat64, 0, math.MaxFloat64}) + if _, _, err := Histogram(a, 2); err == nil { + t.Fatal("Histogram: expected a range error") + } + b := mustFloats(t, []float64{-math.MaxFloat64, 1, math.MaxFloat64, 4}, 2, 2) + bx, err := core.Slice(b, 1, 0, 1) + if err != nil { + t.Fatalf("Slice: %v", err) + } + if _, _, _, err := Histogram2D(bx, b, 2, 2); err == nil { + t.Fatal("Histogram2D: expected a range error") + } +} + +// TestChiSquareGOFRejectsInfiniteExpected pins that an +// infinite expectation is refused at the entry point, under its own +// name, rather than surfacing as a NaN inside the tail function. +func TestChiSquareGOFRejectsInfiniteExpected(t *testing.T) { + obs := mustFloats(t, []float64{10, 12, 9}) + exp := mustFloats(t, []float64{math.Inf(1), 10, 10}) + _, _, _, err := ChiSquareGoodnessOfFit(obs, exp) + if err == nil || !strings.Contains(err.Error(), "expected frequencies") { + t.Fatalf("ChiSquareGoodnessOfFit: err = %v, want the expected-frequencies refusal", err) + } +} diff --git a/stats/degenerate_input_pin_test.go b/stats/degenerate_input_pin_test.go new file mode 100644 index 0000000..f935aa5 --- /dev/null +++ b/stats/degenerate_input_pin_test.go @@ -0,0 +1,74 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestPoissonDrawsZeroLambda pins the degenerate distribution: λ = 0 +// must draw 0, never the −1 the multiplication loop used to hand back. +func TestPoissonDrawsZeroLambda(t *testing.T) { + g := core.NewGenerator(7) + out, err := PoissonDraws(g, 16, 0) + if err != nil { + t.Fatalf("PoissonDraws: %v", err) + } + for i := range 16 { + if out.FloatAt(i) != 0 { + t.Fatalf("draw %d = %g, want 0", i, out.FloatAt(i)) + } + } +} + +// TestChiSquareComplexRejects pins the dtype guard. +func TestChiSquareComplexRejects(t *testing.T) { + obs, err := core.FromComplexes([]complex128{1, 2, 3}, 3) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + exp, _ := core.FromFloats([]float64{1, 2, 3}, 3) + if _, _, _, err := ChiSquareGoodnessOfFit(obs, exp); err == nil { + t.Fatal("ChiSquareGoodnessOfFit accepted complex observed") + } + if _, _, _, err := ChiSquareGoodnessOfFit(exp, obs); err == nil { + t.Fatal("ChiSquareGoodnessOfFit accepted complex expected") + } +} + +// TestBootstrapCIFreshArray pins the documented contract: a statistic +// that retains its argument must observe values frozen at call time, +// never the next resample's mutation through the shared buffer. +func TestBootstrapCIFreshArray(t *testing.T) { + data, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8}, 8) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + type snapshot struct { + arr *core.Array + mean float64 + } + var kept []snapshot + stat := func(a *core.Array) (float64, error) { + m, err := core.Mean(a) + if err != nil { + return 0, err + } + kept = append(kept, snapshot{a, m}) + return m, nil + } + if _, _, err := BootstrapCI(data, stat, 0.9, 8, 42); err != nil { + t.Fatalf("BootstrapCI: %v", err) + } + for i, s := range kept { + m, err := core.Mean(s.arr) + if err != nil { + t.Fatalf("kept[%d] mean: %v", i, err) + } + if m != s.mean { + t.Fatalf("kept[%d] mean drifted from %g to %g (buffer was mutated)", i, s.mean, m) + } + } +} diff --git a/stats/distrib.go b/stats/distrib.go new file mode 100644 index 0000000..ab1af01 --- /dev/null +++ b/stats/distrib.go @@ -0,0 +1,182 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import ( + "math" +) + +// ExponentialDraws returns n draws from the exponential distribution +// with the given rate (mean 1/rate), by inverse CDF. +func ExponentialDraws(g *core.Generator, n int, rate float64) (*core.Array, error) { + if n < 1 { + return nil, base.Errf("ExponentialDraws: n must be ≥ 1") + } + if !(rate > 0) { + return nil, base.Errf("ExponentialDraws: rate must be positive, got %v", rate) + } + out := core.New(core.Float, []int{n}...) + for i := range n { + u := 1 - g.Unit() + out.RawFloats()[i] = -math.Log(u) / rate + } + return out, nil +} + +// GammaDraws returns n draws from the gamma distribution with shape +// α > 0 and rate β > 0, by Marsaglia-Tsang for α ≥ 1 and by boosting +// with the exponential for α < 1. +func GammaDraws(g *core.Generator, n int, alpha, beta float64) (*core.Array, error) { + if n < 1 { + return nil, base.Errf("GammaDraws: n must be ≥ 1") + } + if !(alpha > 0 && beta > 0) { + return nil, base.Errf("GammaDraws: shape and rate must be positive") + } + out := core.New(core.Float, []int{n}...) + d := alpha - 1.0/3 + c := 1 / math.Sqrt(9*d) + for i := range n { + if alpha >= 1 { + out.RawFloats()[i] = gammaMarsagliaTsang(g, alpha, d, c) / beta + continue + } + // α < 1: boost to α + 1 and scale by a uniform^(1/α). + boost := alpha + 1 + dd := boost - 1.0/3 + cc := 1 / math.Sqrt(9*dd) + v := gammaMarsagliaTsang(g, boost, dd, cc) + out.RawFloats()[i] = v * math.Pow(g.Unit(), 1/alpha) / beta + } + return out, nil +} + +// gammaMarsagliaTsang draws one gamma(α, 1) for α ≥ 1 by the +// Marsaglia-Tsang squeeze: a normal draw shapes the cube root of the +// scale, an exponential-tilted accept/reject polishes the tail. +func gammaMarsagliaTsang(g *core.Generator, alpha, d, c float64) float64 { + for { + x := g.NormalUnit() + v := 1 + c*x + if v <= 0 { + continue + } + vvv := v * v * v + u := g.Unit() + if u < 1-0.0331*x*x*x*x { + return d * vvv + } + if math.Log(u) < 0.5*x*x+d*(1-vvv+math.Log(vvv)) { + return d * vvv + } + } +} + +// ChiSquareDraws returns n draws from the chi-squared distribution +// with df degrees of freedom, the gamma(df/2, 2) distribution. +func ChiSquareDraws(g *core.Generator, n int, df int) (*core.Array, error) { + if df < 1 { + return nil, base.Errf("ChiSquareDraws: df must be ≥ 1") + } + return GammaDraws(g, n, float64(df)/2, 0.5) +} + +// StudentTDraws returns n draws from Student's t with df degrees of +// freedom, as N(0,1)/√(χ²_df/df). +func StudentTDraws(g *core.Generator, n int, df int) (*core.Array, error) { + if df < 1 { + return nil, base.Errf("StudentTDraws: df must be ≥ 1") + } + chi, err := ChiSquareDraws(g, n, df) + if err != nil { + return nil, err + } + out := core.New(core.Float, []int{n}...) + for i := range n { + z := g.NormalUnit() + den := math.Sqrt(chi.FloatAt(i) / float64(df)) + if den == 0 { + // The χ² draw underflowed to exactly zero: the ratio is an + // infinity carrying the numerator's sign, not an unsigned + // one (and not the NaN a 0/0 numerator of zero would make). + out.RawFloats()[i] = math.Copysign(math.Inf(1), z) + } else { + out.RawFloats()[i] = z / den + } + } + return out, nil +} + +// PoissonDraws returns n draws from the Poisson distribution with the +// given λ, by the Knuth multiplication method for λ < 30 and the +// normal approximation above. +func PoissonDraws(g *core.Generator, n int, lambda float64) (*core.Array, error) { + if n < 1 { + return nil, base.Errf("PoissonDraws: n must be ≥ 1") + } + if !(lambda >= 0) { + return nil, base.Errf("PoissonDraws: λ must be ≥ 0") + } + out := core.New(core.Float, []int{n}...) + L := math.Exp(-lambda) + for i := range n { + if lambda < 30 { + // λ = 0 (or tiny enough that exp(−λ) rounds to 1) is the + // degenerate distribution at 0: the multiplication loop + // would never run and hand back k−1 = −1. + if L >= 1 { + out.RawFloats()[i] = 0 + continue + } + k := 0.0 + p := 1.0 + for p > L { + k++ + p *= g.Unit() + } + out.RawFloats()[i] = k - 1 + } else { + // Normal approximation for large λ. + z := g.NormalUnit() + out.RawFloats()[i] = max(0, math.Floor(lambda+math.Sqrt(lambda)*z+0.5)) + } + } + return out, nil +} + +// BinomialDraws returns n draws from the binomial distribution with +// the given number of trials and success probability: each draw runs +// trials uniform comparisons against p and counts the successes, the +// exact per-trial loop. It is exact but costs O(trials) uniforms per +// draw, so keep trials modest. +func BinomialDraws(g *core.Generator, n int, trials int, p float64) (*core.Array, error) { + if n < 1 { + // The sibling draws all refuse n < 1; without the guard this one + // returned a nil array with a nil error, which a caller that only + // checks the error then dereferenced. + return nil, base.Errf("BinomialDraws: n must be ≥ 1") + } + if trials < 1 { + return nil, base.Errf("BinomialDraws: trials must be ≥ 1") + } + if !(p >= 0 && p <= 1) { + return nil, base.Errf("BinomialDraws: p must be in [0, 1]") + } + out := core.New(core.Float, []int{n}...) + for i := range n { + count := 0.0 + for range trials { + if g.Unit() < p { + count++ + } + } + out.RawFloats()[i] = count + } + return out, nil +} diff --git a/stats/distrib2.go b/stats/distrib2.go new file mode 100644 index 0000000..d1afd1d --- /dev/null +++ b/stats/distrib2.go @@ -0,0 +1,433 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import ( + "math" +) + +// Distribution additions to the families of distrib.go and cdf.go: the +// Weibull, lognormal, Pareto and negative binomial laws, and the +// Dirichlet. The continuous laws carry closed density, CDF and +// quantile forms; the negative binomial follows the discrete house +// shape of Poisson and the binomial, counting on the same integer axis +// the Poisson counts on. The support convention of the existing +// distributions holds throughout: a CDF is 0 below the support and a +// density 0 outside it, while a faulty parameter is an error and never +// a silent value. + +// WeibullDensity returns the Weibull density with shape k > 0 and +// scale λ > 0 at x ≥ 0: (k/λ)(x/λ)^{k−1}e^{−(x/λ)^k}. Below the +// support, and at +∞, the density is 0; at x = 0 the formula speaks +// for itself, 0 for k > 1, the finite 1/λ for k = 1 and the +Inf the +// integrable singularity carries for k < 1. +func WeibullDensity(x, k, lambda float64) (float64, error) { + if !(k > 0) || math.IsInf(k, 0) { + return 0, base.Errf("WeibullDensity: shape k must be finite and positive, got %g", k) + } + if !(lambda > 0) || math.IsInf(lambda, 0) { + return 0, base.Errf("WeibullDensity: scale λ must be finite and positive, got %g", lambda) + } + if math.IsNaN(x) { + return 0, base.Errf("WeibullDensity: x must be a number, got %g", x) + } + if x < 0 || math.IsInf(x, 1) { + return 0, nil + } + z := x / lambda + return k / lambda * math.Pow(z, k-1) * math.Exp(-math.Pow(z, k)), nil +} + +// WeibullCDF returns P(X ≤ x) for X ~ Weibull(k, λ), the closed form +// 1 − e^{−(x/λ)^k}. +func WeibullCDF(x, k, lambda float64) (float64, error) { + if !(k > 0) || math.IsInf(k, 0) { + return 0, base.Errf("WeibullCDF: shape k must be finite and positive, got %g", k) + } + if !(lambda > 0) || math.IsInf(lambda, 0) { + return 0, base.Errf("WeibullCDF: scale λ must be finite and positive, got %g", lambda) + } + if math.IsNaN(x) { + return 0, base.Errf("WeibullCDF: x must be a number, got %g", x) + } + if x <= 0 { + return 0, nil + } + // Expm1 keeps the left tail, as in ExponentialCDF. + return -math.Expm1(-math.Pow(x/lambda, k)), nil +} + +// WeibullQuantile returns the q-quantile of Weibull(k, λ), the closed +// form λ(−ln(1−q))^{1/k}. +func WeibullQuantile(q, k, lambda float64) (float64, error) { + if !(k > 0) || math.IsInf(k, 0) { + return 0, base.Errf("WeibullQuantile: shape k must be finite and positive, got %g", k) + } + if !(lambda > 0) || math.IsInf(lambda, 0) { + return 0, base.Errf("WeibullQuantile: scale λ must be finite and positive, got %g", lambda) + } + // NaN-rejecting on purpose, as in NormalQuantile. + if !(q >= 0 && q <= 1) { + return 0, base.Errf("WeibullQuantile: q must lie in [0, 1], got %g", q) + } + if q == 0 || q == 1 { + return 0, base.Errf("WeibullQuantile: q = %g has no finite quantile", q) + } + return lambda * math.Pow(-math.Log1p(-q), 1/k), nil +} + +// LognormalDensity returns the lognormal density with location μ and +// log-scale σ > 0 at x > 0: 1/(xσ√(2π))e^{−(ln x−μ)²/(2σ²)}. The +// support convention gives 0 at x ≤ 0 and at +∞. +func LognormalDensity(x, mu, sigma float64) (float64, error) { + if math.IsNaN(mu) || math.IsInf(mu, 0) { + return 0, base.Errf("LognormalDensity: location μ must be finite, got %g", mu) + } + if !(sigma > 0) || math.IsInf(sigma, 0) { + return 0, base.Errf("LognormalDensity: log-scale σ must be finite and positive, got %g", sigma) + } + if math.IsNaN(x) { + return 0, base.Errf("LognormalDensity: x must be a number, got %g", x) + } + if x <= 0 || math.IsInf(x, 1) { + return 0, nil + } + z := (math.Log(x) - mu) / sigma + // The tail is assembled in log space: at a subnormal x both the + // numerator and the denominator of the closed form underflow to + // exact zero, and their division reports NaN for a point inside + // the support whose density is an honest 0. + return math.Exp(-z*z/2 - math.Log(x) - math.Log(sigma) - 0.5*math.Log(2*math.Pi)), nil +} + +// LognormalCDF returns P(X ≤ x) for X ~ lognormal(μ, σ), the normal +// CDF at (ln x − μ)/σ, the reduction the law is named for. +func LognormalCDF(x, mu, sigma float64) (float64, error) { + if math.IsNaN(mu) || math.IsInf(mu, 0) { + return 0, base.Errf("LognormalCDF: location μ must be finite, got %g", mu) + } + if !(sigma > 0) || math.IsInf(sigma, 0) { + return 0, base.Errf("LognormalCDF: log-scale σ must be finite and positive, got %g", sigma) + } + if math.IsNaN(x) { + return 0, base.Errf("LognormalCDF: x must be a number, got %g", x) + } + if x <= 0 { + return 0, nil + } + return NormalCDF((math.Log(x) - mu) / sigma), nil +} + +// LognormalQuantile returns the q-quantile of lognormal(μ, σ) through +// the existing normal quantile: e^{μ + σ·Φ^{−1}(q)}. +func LognormalQuantile(q, mu, sigma float64) (float64, error) { + if math.IsNaN(mu) || math.IsInf(mu, 0) { + return 0, base.Errf("LognormalQuantile: location μ must be finite, got %g", mu) + } + if !(sigma > 0) || math.IsInf(sigma, 0) { + return 0, base.Errf("LognormalQuantile: log-scale σ must be finite and positive, got %g", sigma) + } + z, err := NormalQuantile(q) + if err != nil { + return 0, base.Errf("LognormalQuantile: %w", err) + } + return math.Exp(mu + sigma*z), nil +} + +// ParetoDensity returns the Pareto density with scale x_m > 0 and tail +// index α > 0 at x ≥ x_m: α·x_m^α/x^{α+1}. Below the support, and at +// +∞, the density is 0. +func ParetoDensity(x, xm, alpha float64) (float64, error) { + if !(xm > 0) || math.IsInf(xm, 0) { + return 0, base.Errf("ParetoDensity: scale x_m must be finite and positive, got %g", xm) + } + if !(alpha > 0) || math.IsInf(alpha, 0) { + return 0, base.Errf("ParetoDensity: tail index α must be finite and positive, got %g", alpha) + } + if math.IsNaN(x) { + return 0, base.Errf("ParetoDensity: x must be a number, got %g", x) + } + if x < xm || math.IsInf(x, 1) { + return 0, nil + } + return alpha / x * math.Pow(xm/x, alpha), nil +} + +// ParetoCDF returns P(X ≤ x) for X ~ Pareto(x_m, α), the closed form +// 1 − (x_m/x)^α, evaluated through Expm1 so the answers just above the +// support keep their digits. +func ParetoCDF(x, xm, alpha float64) (float64, error) { + if !(xm > 0) || math.IsInf(xm, 0) { + return 0, base.Errf("ParetoCDF: scale x_m must be finite and positive, got %g", xm) + } + if !(alpha > 0) || math.IsInf(alpha, 0) { + return 0, base.Errf("ParetoCDF: tail index α must be finite and positive, got %g", alpha) + } + if math.IsNaN(x) { + return 0, base.Errf("ParetoCDF: x must be a number, got %g", x) + } + if x < xm { + return 0, nil + } + return -math.Expm1(alpha * math.Log(xm/x)), nil +} + +// ParetoQuantile returns the q-quantile of Pareto(x_m, α), the closed +// form x_m(1−q)^{−1/α}. +func ParetoQuantile(q, xm, alpha float64) (float64, error) { + if !(xm > 0) || math.IsInf(xm, 0) { + return 0, base.Errf("ParetoQuantile: scale x_m must be finite and positive, got %g", xm) + } + if !(alpha > 0) || math.IsInf(alpha, 0) { + return 0, base.Errf("ParetoQuantile: tail index α must be finite and positive, got %g", alpha) + } + if !(q >= 0 && q <= 1) { + return 0, base.Errf("ParetoQuantile: q must lie in [0, 1], got %g", q) + } + if q == 0 || q == 1 { + return 0, base.Errf("ParetoQuantile: q = %g has no finite quantile", q) + } + return xm * math.Pow(1-q, -1/alpha), nil +} + +// NegativeBinomialPMF returns P(X = k), the probability of k failures +// before the r-th success in independent trials of probability p: the +// law the Poisson draws and BinomialDraws count on the same integer +// axis. The mass is assembled in log space with lgamma, the form that +// keeps every term representable for large r and k. +func NegativeBinomialPMF(k, r int, p float64) (float64, error) { + if r < 1 { + return 0, base.Errf("NegativeBinomialPMF: r must be ≥ 1, got %d", r) + } + if !(p > 0 && p < 1) { + return 0, base.Errf("NegativeBinomialPMF: p must lie in (0, 1), got %g", p) + } + if k < 0 { + return 0, nil + } + lf := logGamma(float64(r+k)) - logGamma(float64(r)) - logGamma(float64(k+1)) + + float64(r)*math.Log(p) + float64(k)*math.Log1p(-p) + return math.Exp(lf), nil +} + +// NegativeBinomialCDF returns P(X ≤ k) for the number of failures X +// before the r-th success, by direct summation of the PMF terms under +// the multiplicative recurrence term_{j+1} = term_j·(r+j)/(j+1)·(1−p). +// Every term is positive, so the sum carries no cancellation; the same +// probability equals the regularised beta I_p(r, k+1), which the tests +// hold the summation against. The summation seeds from p^r, and a seed +// the format cannot hold rounds to an exact zero the recurrence never +// recovers from: every later term would stay zero while the true mass +// sits further out. That far regime answers through the beta identity +// instead, which keeps the whole support live. +func NegativeBinomialCDF(k, r int, p float64) (float64, error) { + if r < 1 { + return 0, base.Errf("NegativeBinomialCDF: r must be ≥ 1, got %d", r) + } + if !(p > 0 && p < 1) { + return 0, base.Errf("NegativeBinomialCDF: p must lie in (0, 1), got %g", p) + } + if k < 0 { + return 0, nil + } + term := math.Exp(float64(r) * math.Log(p)) + if term == 0 { + // p^r underflowed: the recurrence multiplies zeros, so the sum + // would answer 0 at every k. The identity I_p(r, k+1) = P(X ≤ k) + // evaluates the same probability through the incomplete beta, + // accurate across this regime. + return BetaIncomplete(p, float64(r), float64(k+1)) + } + sum := 0.0 + for j := 0; j <= k; j++ { + sum += term + term *= (float64(r) + float64(j)) / float64(j+1) * (1 - p) + } + return sum, nil +} + +// NegativeBinomialQuantile returns the smallest k with P(X ≤ k) ≥ q, +// through the same discrete bracketed search the Poisson and binomial +// quantiles use. +func NegativeBinomialQuantile(q, p float64, r int) (float64, error) { + if r < 1 { + return 0, base.Errf("NegativeBinomialQuantile: r must be ≥ 1, got %d", r) + } + if !(p > 0 && p < 1) { + return 0, base.Errf("NegativeBinomialQuantile: p must lie in (0, 1), got %g", p) + } + return discreteQuantile("NegativeBinomialQuantile", q, func(k int) (float64, error) { + return NegativeBinomialCDF(k, r, p) + }) +} + +// logGamma wraps math.Lgamma, keeping the log-space call sites free of +// the ignored-error idiom. +func logGamma(x float64) float64 { + l, _ := math.Lgamma(x) + return l +} + +// DirichletDensity returns the Dirichlet density with concentration +// vector α at the simplex point x: Πx_i^{α_i−1}/B(α), the normalising +// constant assembled with lgamma. Both vectors must be finite, every +// α_i positive and every x_i non-negative, and the x must sum to 1 (to +// within 1e-9); a boundary x_i = 0 gives +Inf below α_i = 1, the value +// 1 continues to contribute nothing at α_i = 1 exactly, and 0 above. +func DirichletDensity(alpha, x []float64) (float64, error) { + const name = "DirichletDensity" + if len(alpha) < 2 { + return 0, base.Errf("%s: needs at least two components, got %d", name, len(alpha)) + } + if len(x) != len(alpha) { + return 0, base.Errf("%s: the concentration has %d components, the point %d", name, len(alpha), len(x)) + } + total := 0.0 + for i, a := range alpha { + if !(a > 0) || math.IsInf(a, 0) { + return 0, base.Errf("%s: alpha[%d] must be finite and positive, got %g", name, i, a) + } + total += a + } + sum := 0.0 + for i, v := range x { + if math.IsNaN(v) || math.IsInf(v, 0) { + return 0, base.Errf("%s: x[%d] must be finite, got %g", name, i, v) + } + if v < 0 { + return 0, base.Errf("%s: x[%d] = %g lies outside the simplex", name, i, v) + } + sum += v + } + if math.Abs(sum-1) > 1e-9 { + return 0, base.Errf("%s: the point must sum to 1, got %g", name, sum) + } + // ln B(α) = Σ lgamma(α_i) − lgamma(α₀). + lb := -logGamma(total) + for _, a := range alpha { + lb += logGamma(a) + } + s := -lb + for i, a := range alpha { + switch { + case x[i] == 0: + if a < 1 { + return math.Inf(1), nil + } + if a > 1 { + return 0, nil + } + default: + s += (a - 1) * math.Log(x[i]) + } + } + return math.Exp(s), nil +} + +// DirichletMean returns the mean of the Dirichlet with concentration +// α: the normalised concentration α_i/α₀. +func DirichletMean(alpha []float64) ([]float64, error) { + if len(alpha) < 2 { + return nil, base.Errf("DirichletMean: needs at least two components, got %d", len(alpha)) + } + total := 0.0 + for i, a := range alpha { + if !(a > 0) || math.IsInf(a, 0) { + return nil, base.Errf("DirichletMean: alpha[%d] must be finite and positive, got %g", i, a) + } + total += a + } + mean := make([]float64, len(alpha)) + for i, a := range alpha { + mean[i] = a / total + } + return mean, nil +} + +// DirichletMode returns the interior mode (α_i−1)/(α₀−k), which exists +// only when every concentration exceeds 1; any α_i ≤ 1 pushes the mode +// onto the boundary and is refused rather than answered with a vector +// that is not a mode. +func DirichletMode(alpha []float64) ([]float64, error) { + if len(alpha) < 2 { + return nil, base.Errf("DirichletMode: needs at least two components, got %d", len(alpha)) + } + total := 0.0 + for i, a := range alpha { + if !(a > 0) || math.IsInf(a, 0) { + return nil, base.Errf("DirichletMode: alpha[%d] must be finite and positive, got %g", i, a) + } + if a <= 1 { + return nil, base.Errf("DirichletMode: alpha[%d] = %g leaves no interior mode; every concentration must exceed 1", i, a) + } + total += a + } + den := total - float64(len(alpha)) + mode := make([]float64, len(alpha)) + for i, a := range alpha { + mode[i] = (a - 1) / den + } + return mode, nil +} + +// DirichletDraws returns n draws from the Dirichlet with concentration +// α, as an (n, k) array whose rows are the draws. Each row scales the +// k independent gamma(α_i, 1) draws of GammaDraws' generator to sum to +// one; the all-underflow row (every α_i far below the float64 floor) +// would otherwise divide by zero and falls back to the uniform row. +func DirichletDraws(g *core.Generator, n int, alpha []float64) (*core.Array, error) { + const name = "DirichletDraws" + if n < 1 { + return nil, base.Errf("%s: n must be ≥ 1", name) + } + if len(alpha) < 2 { + return nil, base.Errf("%s: needs at least two components, got %d", name, len(alpha)) + } + for i, a := range alpha { + if !(a > 0) || math.IsInf(a, 0) { + return nil, base.Errf("%s: alpha[%d] must be finite and positive, got %g", name, i, a) + } + } + k := len(alpha) + flat := make([]float64, n*k) + for r := range n { + sum := 0.0 + row := flat[r*k : r*k+k] + for c, a := range alpha { + row[c] = gammaOne(g, a) + sum += row[c] + } + if sum == 0 { + // Every gamma draw underflowed: the uniform row is the + // honest stand-in, an Inf row would poison the draw. + for c := range row { + row[c] = 1 / float64(k) + } + continue + } + for c := range row { + row[c] /= sum + } + } + return floatsToArray(flat, []int{n, k}), nil +} + +// gammaOne draws one gamma(α, 1) variate, the single-draw form of the +// GammaDraws loop: Marsaglia-Tsang for α ≥ 1 and the boost with the +// exponential for α < 1, reusing the package squeeze. +func gammaOne(g *core.Generator, alpha float64) float64 { + if alpha >= 1 { + d := alpha - 1.0/3 + return gammaMarsagliaTsang(g, alpha, d, 1/math.Sqrt(9*d)) + } + boost := alpha + 1 + d := boost - 1.0/3 + v := gammaMarsagliaTsang(g, boost, d, 1/math.Sqrt(9*d)) + return v * math.Pow(g.Unit(), 1/alpha) +} diff --git a/stats/distrib2_test.go b/stats/distrib2_test.go new file mode 100644 index 0000000..21a01c1 --- /dev/null +++ b/stats/distrib2_test.go @@ -0,0 +1,436 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestWeibullClosed pins the Weibull on its closed forms: with k = 2 +// and λ = 1 the density is 2xe^{−x²}, the CDF 1 − e^{−x²} and the +// median √(ln 2); with k = 1 the law is the exponential of rate 1/λ. +func TestWeibullClosed(t *testing.T) { + d, err := WeibullDensity(1, 2, 1) + if err != nil || math.Abs(d-2/math.E) > 1e-15 { + t.Fatalf("WeibullDensity(1, 2, 1) = %v, %v, want 2/e", d, err) + } + v, err := WeibullCDF(1, 2, 1) + if err != nil || math.Abs(v-(1-1/math.E)) > 1e-15 { + t.Fatalf("WeibullCDF(1, 2, 1) = %v, %v, want 1 − 1/e", v, err) + } + v, err = WeibullCDF(3, 2, 1) + if err != nil || math.Abs(v-(1-math.Exp(-9))) > 1e-15 { + t.Fatalf("WeibullCDF(3, 2, 1) = %v, %v, want 1 − e⁻⁹", v, err) + } + q, err := WeibullQuantile(0.5, 2, 1) + if err != nil || math.Abs(q-math.Sqrt(math.Ln2)) > 1e-14 { + t.Fatalf("WeibullQuantile(0.5, 2, 1) = %v, %v, want √(ln 2)", q, err) + } + // k = 1 reduces to the exponential of rate 1/λ. + v, err = WeibullCDF(2, 1, 0.5) + exp, eerr := ExponentialCDF(2, 2) + if err != nil || eerr != nil || v != exp { + t.Fatalf("WeibullCDF(2, 1, 0.5) = %v vs ExponentialCDF %v (%v, %v)", v, exp, err, eerr) + } + // Support convention and round trips. + if v, _ = WeibullCDF(-1, 2, 1); v != 0 { + t.Fatalf("WeibullCDF below the support = %v, want 0", v) + } + if d, _ = WeibullDensity(-1, 2, 1); d != 0 { + t.Fatalf("WeibullDensity below the support = %v, want 0", d) + } + for _, q := range []float64{0.001, 0.1, 0.5, 0.9, 0.999} { + x, err := WeibullQuantile(q, 1.5, 2) + if err != nil { + t.Fatalf("WeibullQuantile(%g): %v", q, err) + } + back, err := WeibullCDF(x, 1.5, 2) + if err != nil || math.Abs(back-q) > 1e-13 { + t.Fatalf("round trip q = %g: CDF(quantile) = %v, %v", q, back, err) + } + } +} + +// TestLognormalClosed pins the lognormal through the normal: the CDF +// at 1 with μ = 0 is exactly ½, the density there is 1/√(2π), and the +// 0.975 quantile is e to the normal quantile. +func TestLognormalClosed(t *testing.T) { + v, err := LognormalCDF(1, 0, 1) + if err != nil || v != 0.5 { + t.Fatalf("LognormalCDF(1, 0, 1) = %v, %v, want exactly 0.5", v, err) + } + v, err = LognormalCDF(math.E, 0, 1) + if err != nil || math.Abs(v-NormalCDF(1)) > 1e-15 { + t.Fatalf("LognormalCDF(e, 0, 1) = %v, %v, want Φ(1)", v, err) + } + d, err := LognormalDensity(1, 0, 1) + if want := 1 / (math.Sqrt2 * math.SqrtPi); err != nil || math.Abs(d-want) > 1e-15 { + t.Fatalf("LognormalDensity(1, 0, 1) = %v, %v, want %.16g", d, err, want) + } + z, _ := NormalQuantile(0.975) + q, err := LognormalQuantile(0.975, 0, 1) + if err != nil || math.Abs(q-math.Exp(z)) > 1e-12*math.Exp(z) { + t.Fatalf("LognormalQuantile(0.975, 0, 1) = %v, %v, want e^{%.16g}", q, err, z) + } + // The CDF is the normal CDF at ln x, pointwise. + for _, x := range []float64{0.05, 0.5, 2, 20} { + got, _ := LognormalCDF(x, 0.3, 0.8) + want := NormalCDF((math.Log(x) - 0.3) / 0.8) + if math.Abs(got-want) > 1e-15 { + t.Fatalf("LognormalCDF(%g, 0.3, 0.8) = %v, want %v", x, got, want) + } + } + if v, _ = LognormalCDF(0, 0, 1); v != 0 { + t.Fatalf("LognormalCDF at 0 = %v, want 0", v) + } + if d, _ = LognormalDensity(0, 0, 1); d != 0 { + t.Fatalf("LognormalDensity at 0 = %v, want 0", d) + } + // A subnormal x underflows the closed form's numerator and + // denominator together: the log-space tail must answer 0, not the + // NaN their division would report. + for _, sigma := range []float64{0.1, 0.01, 1} { + if d, err = LognormalDensity(5e-324, 0, sigma); err != nil || d != 0 { + t.Fatalf("LognormalDensity(5e-324, 0, %g) = %v, %v, want 0", sigma, d, err) + } + } +} + +// TestParetoClosed pins the Pareto on its closed forms with x_m = 1 +// and α = 3: the CDF at 2 is 1 − 2^{−3} = 7/8, the density there +// 3/16, and the 7/8 quantile exactly 2. +func TestParetoClosed(t *testing.T) { + v, err := ParetoCDF(2, 1, 3) + if err != nil || math.Abs(v-0.875) > 1e-15 { + t.Fatalf("ParetoCDF(2, 1, 3) = %v, %v, want 0.875", v, err) + } + d, err := ParetoDensity(2, 1, 3) + if err != nil || math.Abs(d-3.0/16) > 1e-15 { + t.Fatalf("ParetoDensity(2, 1, 3) = %v, %v, want 3/16", d, err) + } + q, err := ParetoQuantile(0.875, 1, 3) + if err != nil || math.Abs(q-2) > 1e-12 { + t.Fatalf("ParetoQuantile(0.875, 1, 3) = %v, %v, want 2", q, err) + } + if v, _ = ParetoCDF(0.5, 1, 3); v != 0 { + t.Fatalf("ParetoCDF below the support = %v, want 0", v) + } + if d, _ = ParetoDensity(0.5, 1, 3); d != 0 { + t.Fatalf("ParetoDensity below the support = %v, want 0", d) + } + // The Expm1 form keeps the digits just above the support. The + // exact probability of the representable input is 3u for the + // offset u the double actually carries. + input := 1 + 1e-12 + u := input - 1 + v, err = ParetoCDF(input, 1, 3) + if err != nil || math.Abs(v-3*u) > 1e-10*3*u { + t.Fatalf("ParetoCDF just above the support = %v, %v, want ≈ %.17g", v, err, 3*u) + } +} + +// TestNegativeBinomialClosed pins the negative binomial on exact +// fractions with r = 3 and p = ½, and holds the summation CDF against +// the exact identity P(X ≤ k) = I_p(r, k+1). +func TestNegativeBinomialClosed(t *testing.T) { + d, err := NegativeBinomialPMF(0, 3, 0.5) + if err != nil || math.Abs(d-math.Pow(0.5, 3)) > 1e-16 { + t.Fatalf("NegativeBinomialPMF(0, 3, 0.5) = %v, %v, want 1/8", d, err) + } + d, err = NegativeBinomialPMF(1, 3, 0.5) + if err != nil || math.Abs(d-3*math.Pow(0.5, 4)) > 1e-16 { + t.Fatalf("NegativeBinomialPMF(1, 3, 0.5) = %v, %v, want 3/16", d, err) + } + v, err := NegativeBinomialCDF(1, 3, 0.5) + if err != nil || math.Abs(v-0.3125) > 1e-15 { + t.Fatalf("NegativeBinomialCDF(1, 3, 0.5) = %v, %v, want 5/16", v, err) + } + // Exact identity against the regularised beta. + for _, r := range []int{1, 2, 5, 12} { + for _, p := range []float64{0.2, 0.5, 0.8} { + for k := 0; k <= 25; k++ { + got, err := NegativeBinomialCDF(k, r, p) + if err != nil { + t.Fatalf("NegativeBinomialCDF(%d, %d, %g): %v", k, r, p, err) + } + want, err := BetaIncomplete(p, float64(r), float64(k+1)) + if err != nil { + t.Fatalf("BetaIncomplete: %v", err) + } + if math.Abs(got-want) > 1e-14 { + t.Fatalf("NegativeBinomialCDF(%d, %d, %g) = %.16g, want I_p identity %.16g", + k, r, p, got, want) + } + } + } + } + // The CDF at k = 2 is exactly ½, so both sides of the smallest-k + // rule are pinned. + q, err := NegativeBinomialQuantile(0.5, 0.5, 3) + if err != nil || q != 2 { + t.Fatalf("NegativeBinomialQuantile(0.5, 0.5, 3) = %v, %v, want 2", q, err) + } + q, err = NegativeBinomialQuantile(0.3125, 0.5, 3) + if err != nil || q != 1 { + t.Fatalf("NegativeBinomialQuantile(0.3125, 0.5, 3) = %v, %v, want 1", q, err) + } + // q = 1 brackets once the summed tail underflows below a half ulp + // of 1, at a few hundred failures for these parameters. + qEnd, err := NegativeBinomialQuantile(1, 0.5, 3) + if err != nil || qEnd < 10 { + t.Fatalf("NegativeBinomialQuantile(1, 0.5, 3) = %v, %v", qEnd, err) + } + if d, err = NegativeBinomialPMF(-1, 3, 0.5); err != nil || d != 0 { + t.Fatalf("NegativeBinomialPMF below the support = %v, %v", d, err) + } +} + +// TestDirichletDensityMoments pins the density on an exact rational +// value, the boundary conventions, and the mean and mode helpers. +// With α = (2, 3, 4) the normalising constant is +// B(α) = Γ2Γ3Γ4/Γ9 = 12/40320 = 1/3360, so the uniform point carries +// 3360·3^{−6} = 3360/729. +func TestDirichletDensityMoments(t *testing.T) { + third := 1.0 / 3 + d, err := DirichletDensity([]float64{2, 3, 4}, []float64{third, third, third}) + if err != nil || math.Abs(d-3360.0/729) > 1e-12*3360.0/729 { + t.Fatalf("DirichletDensity uniform = %v, %v, want %.16g", d, err, 3360.0/729) + } + // Boundary: α_i > 1 kills the density, α_i < 1 blows it up, and an + // α_i of exactly 1 contributes nothing. + if d, _ = DirichletDensity([]float64{2, 2}, []float64{0, 1}); d != 0 { + t.Fatalf("boundary density with α > 1 = %v, want 0", d) + } + if d, _ = DirichletDensity([]float64{0.5, 0.5}, []float64{0, 1}); !math.IsInf(d, 1) { + t.Fatalf("boundary density with α < 1 = %v, want +Inf", d) + } + d, err = DirichletDensity([]float64{1, 3}, []float64{0, 1}) + if err != nil || math.Abs(d-3) > 1e-14 { + t.Fatalf("boundary density with α = 1 = %v, %v, want 3", d, err) + } + mean, err := DirichletMean([]float64{2, 3, 4}) + if err != nil { + t.Fatalf("DirichletMean: %v", err) + } + for i, want := range []float64{2.0 / 9, 1.0 / 3, 4.0 / 9} { + if math.Abs(mean[i]-want) > 1e-15 { + t.Fatalf("DirichletMean[%d] = %.16g, want %.16g", i, mean[i], want) + } + } + mode, err := DirichletMode([]float64{2, 3, 4}) + if err != nil { + t.Fatalf("DirichletMode: %v", err) + } + for i, want := range []float64{1.0 / 6, 1.0 / 3, 0.5} { + if math.Abs(mode[i]-want) > 1e-15 { + t.Fatalf("DirichletMode[%d] = %.16g, want %.16g", i, mode[i], want) + } + } + if _, err := DirichletMode([]float64{1, 2}); err == nil { + t.Fatal("an α of 1 leaves no interior mode: want an error") + } +} + +// TestDirichletDraws checks the sampler statistically: every row is a +// probability vector, the column means track α/α₀, and the seed makes +// the run reproducible. +func TestDirichletDraws(t *testing.T) { + g := core.NewGenerator(7) + alpha := []float64{2, 3, 4} + const n = 200000 + draws, err := DirichletDraws(g, n, alpha) + if err != nil { + t.Fatalf("DirichletDraws: %v", err) + } + if draws.Shape()[0] != n || draws.Shape()[1] != len(alpha) { + t.Fatalf("draw shape %v, want (%d, %d)", draws.Shape(), n, len(alpha)) + } + colSum := make([]float64, len(alpha)) + for r := range n { + sum := 0.0 + for c := range alpha { + v := draws.FloatAt(r*len(alpha) + c) + if v < 0 { + t.Fatalf("draw (%d, %d) = %v is negative", r, c, v) + } + sum += v + colSum[c] += v + } + if math.Abs(sum-1) > 1e-12 { + t.Fatalf("row %d sums to %.17g, want 1", r, sum) + } + } + mean, _ := DirichletMean(alpha) + for c := range alpha { + m := colSum[c] / n + if math.Abs(m-mean[c]) > 0.005 { + t.Fatalf("column %d mean = %g, want ≈ %g", c, m, mean[c]) + } + } + again, _ := DirichletDraws(core.NewGenerator(7), n, alpha) + for i := range n * len(alpha) { + if again.FloatAt(i) != draws.FloatAt(i) { + t.Fatalf("seeded run not deterministic at %d", i) + } + } + // Concentrations below 1 take the boosted gamma branch of the + // sampler; the mean still tracks α/α₀. + sparse, err := DirichletDraws(core.NewGenerator(3), n, []float64{0.5, 0.7}) + if err != nil { + t.Fatalf("DirichletDraws: %v", err) + } + sums := []float64{0, 0} + for r := range n { + for c := range 2 { + v := sparse.FloatAt(r*2 + c) + sums[c] += v + } + } + for c, a := range []float64{0.5, 0.7} { + if m := sums[c] / n; math.Abs(m-a/1.2) > 0.005 { + t.Fatalf("sparse column %d mean = %g, want ≈ %g", c, m, a/1.2) + } + } + if _, err := DirichletDraws(core.NewGenerator(1), 1, []float64{2, -1}); err == nil { + t.Fatal("a negative concentration: want an error") + } +} + +// TestDistribution2Errors pins the parameter contracts of the new +// distributions. +func TestDistribution2Errors(t *testing.T) { + if _, err := WeibullDensity(1, 0, 1); err == nil || !strings.Contains(err.Error(), "shape k") { + t.Fatalf("k = 0: got %v, want the shape refusal", err) + } + if _, err := WeibullDensity(1, 2, 0); err == nil || !strings.Contains(err.Error(), "scale λ") { + t.Fatalf("λ = 0: got %v, want the scale refusal", err) + } + if _, err := WeibullDensity(math.NaN(), 2, 1); err == nil || !strings.Contains(err.Error(), "must be a number") { + t.Fatalf("NaN x: got %v, want the x refusal", err) + } + if _, err := WeibullCDF(1, -1, 1); err == nil || !strings.Contains(err.Error(), "shape k") { + t.Fatalf("negative k: got %v, want the shape refusal", err) + } + if _, err := WeibullQuantile(1.5, 2, 1); err == nil || !strings.Contains(err.Error(), "q must lie") { + t.Fatalf("q outside [0, 1]: got %v, want the q refusal", err) + } + if _, err := WeibullQuantile(0, 2, 1); err == nil || !strings.Contains(err.Error(), "no finite quantile") { + t.Fatalf("q = 0: got %v, want the q refusal", err) + } + if _, err := WeibullCDF(1, 2, math.Inf(1)); err == nil || !strings.Contains(err.Error(), "scale λ") { + t.Fatalf("λ = +Inf: got %v, want the scale refusal", err) + } + if v, err := WeibullCDF(math.Inf(1), 2, 1); err != nil || v != 1 { + t.Fatalf("WeibullCDF(+Inf) = %v, %v, want 1", v, err) + } + if _, err := LognormalQuantile(0.5, math.NaN(), 1); err == nil || !strings.Contains(err.Error(), "location μ") { + t.Fatalf("μ = NaN in the quantile: got %v, want the location refusal", err) + } + if _, err := LognormalQuantile(0.5, 0, math.Inf(1)); err == nil || !strings.Contains(err.Error(), "log-scale σ") { + t.Fatalf("σ = +Inf in the quantile: got %v, want the log-scale refusal", err) + } + if _, err := LognormalCDF(math.NaN(), 0, 1); err == nil || !strings.Contains(err.Error(), "must be a number") { + t.Fatalf("NaN x: got %v, want the x refusal", err) + } + if _, err := ParetoQuantile(0.5, 0, 2); err == nil || !strings.Contains(err.Error(), "scale x_m") { + t.Fatalf("x_m = 0 in the quantile: got %v, want the scale refusal", err) + } + if _, err := ParetoQuantile(0.5, 1, math.NaN()); err == nil || !strings.Contains(err.Error(), "tail index α") { + t.Fatalf("NaN α: got %v, want the tail-index refusal", err) + } + if _, err := ParetoCDF(math.NaN(), 1, 2); err == nil || !strings.Contains(err.Error(), "must be a number") { + t.Fatalf("NaN x: got %v, want the x refusal", err) + } + if _, err := WeibullQuantile(0.5, math.Inf(1), 1); err == nil || !strings.Contains(err.Error(), "shape k") { + t.Fatalf("k = +Inf in the quantile: got %v, want the shape refusal", err) + } + if _, err := WeibullQuantile(0.5, 2, math.Inf(-1)); err == nil || !strings.Contains(err.Error(), "scale λ") { + t.Fatalf("λ = −Inf in the quantile: got %v, want the scale refusal", err) + } + if _, err := LognormalDensity(math.NaN(), 0, 1); err == nil || !strings.Contains(err.Error(), "must be a number") { + t.Fatalf("NaN x in the density: got %v, want the x refusal", err) + } + if _, err := ParetoDensity(math.NaN(), 1, 2); err == nil || !strings.Contains(err.Error(), "must be a number") { + t.Fatalf("NaN x in the Pareto density: got %v, want the x refusal", err) + } + if _, err := ParetoDensity(1, math.Inf(1), 2); err == nil || !strings.Contains(err.Error(), "scale x_m") { + t.Fatalf("x_m = +Inf: got %v, want the scale refusal", err) + } + if _, err := NegativeBinomialCDF(0, 0, 0.5); err == nil || !strings.Contains(err.Error(), "r must be") { + t.Fatalf("r = 0 in the CDF: got %v, want the r refusal", err) + } + if _, err := DirichletMode([]float64{2}); err == nil || !strings.Contains(err.Error(), "at least two components") { + t.Fatalf("one component in the mode: got %v, want the component floor refusal", err) + } + if _, err := DirichletMode([]float64{2, math.Inf(1)}); err == nil || !strings.Contains(err.Error(), "alpha[1]") { + t.Fatalf("α = +Inf in the mode: got %v, want the concentration refusal", err) + } + if _, err := DirichletDraws(core.NewGenerator(1), 1, []float64{2}); err == nil || !strings.Contains(err.Error(), "at least two components") { + t.Fatalf("one component in the draws: got %v, want the component floor refusal", err) + } + if _, err := DirichletMean([]float64{2, math.NaN()}); err == nil || !strings.Contains(err.Error(), "alpha[1]") { + t.Fatalf("NaN α in the mean: got %v, want the concentration refusal", err) + } + if _, err := LognormalDensity(1, math.Inf(1), 1); err == nil || !strings.Contains(err.Error(), "location μ") { + t.Fatalf("μ = +Inf: got %v, want the location refusal", err) + } + if _, err := LognormalCDF(1, 0, -1); err == nil || !strings.Contains(err.Error(), "log-scale σ") { + t.Fatalf("σ < 0: got %v, want the log-scale refusal", err) + } + if _, err := LognormalQuantile(0.5, 0, 0); err == nil || !strings.Contains(err.Error(), "log-scale σ") { + t.Fatalf("σ = 0: got %v, want the log-scale refusal", err) + } + if _, err := ParetoDensity(1, 0, 2); err == nil || !strings.Contains(err.Error(), "scale x_m") { + t.Fatalf("x_m = 0: got %v, want the scale refusal", err) + } + if _, err := ParetoCDF(1, 1, math.Inf(1)); err == nil || !strings.Contains(err.Error(), "tail index α") { + t.Fatalf("α = +Inf: got %v, want the tail-index refusal", err) + } + if _, err := ParetoQuantile(1, 1, 2); err == nil || !strings.Contains(err.Error(), "no finite quantile") { + t.Fatalf("q = 1: got %v, want the q refusal", err) + } + if _, err := NegativeBinomialPMF(0, 0, 0.5); err == nil || !strings.Contains(err.Error(), "r must be") { + t.Fatalf("r = 0: got %v, want the r refusal", err) + } + if _, err := NegativeBinomialPMF(0, 3, 1.5); err == nil || !strings.Contains(err.Error(), "p must lie") { + t.Fatalf("p above 1 in the PMF: got %v, want the p refusal", err) + } + if _, err := NegativeBinomialCDF(0, 3, 0); err == nil || !strings.Contains(err.Error(), "p must lie") { + t.Fatalf("p = 0: got %v, want the p refusal", err) + } + if _, err := NegativeBinomialCDF(0, 3, 1); err == nil || !strings.Contains(err.Error(), "p must lie") { + t.Fatalf("p = 1: got %v, want the p refusal", err) + } + if _, err := NegativeBinomialQuantile(0.5, 0.5, 0); err == nil || !strings.Contains(err.Error(), "r must be") { + t.Fatalf("r = 0 in the quantile: got %v, want the r refusal", err) + } + if _, err := NegativeBinomialQuantile(0.5, 1.5, 3); err == nil || !strings.Contains(err.Error(), "p must lie") { + t.Fatalf("p above 1 in the quantile: got %v, want the p refusal", err) + } + if _, err := NegativeBinomialQuantile(-0.1, 0.5, 3); err == nil || !strings.Contains(err.Error(), "q must lie") { + t.Fatalf("q < 0: got %v, want the q refusal", err) + } + if _, err := DirichletDensity([]float64{2}, []float64{1}); err == nil || !strings.Contains(err.Error(), "at least two components") { + t.Fatalf("one component: got %v, want the component floor refusal", err) + } + if _, err := DirichletDensity([]float64{2, 3}, []float64{0.5}); err == nil || !strings.Contains(err.Error(), "components, the point") { + t.Fatalf("length mismatch: got %v, want the length refusal", err) + } + if _, err := DirichletDensity([]float64{2, 3}, []float64{0.5, 0.2}); err == nil || !strings.Contains(err.Error(), "must sum to 1") { + t.Fatalf("off-simplex point: got %v, want the simplex refusal", err) + } + if _, err := DirichletDensity([]float64{-1, 3}, []float64{0.5, 0.5}); err == nil || !strings.Contains(err.Error(), "alpha[0]") { + t.Fatalf("negative α: got %v, want the concentration refusal", err) + } + if _, err := DirichletMean([]float64{0, 3}); err == nil || !strings.Contains(err.Error(), "alpha[0]") { + t.Fatalf("α = 0 in the mean: got %v, want the concentration refusal", err) + } + if _, err := DirichletDraws(core.NewGenerator(1), 0, []float64{2, 3}); err == nil || !strings.Contains(err.Error(), "n must be") { + t.Fatalf("n = 0: got %v, want the n refusal", err) + } +} diff --git a/stats/distrib_test.go b/stats/distrib_test.go new file mode 100644 index 0000000..df221f8 --- /dev/null +++ b/stats/distrib_test.go @@ -0,0 +1,94 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestDistributions checks the seeded draws against theoretical +// moments with generous statistical tolerances, and the error +// contract for invalid parameters. +func TestDistributions(t *testing.T) { + g := core.NewGenerator(42) + const n = 200000 + + // Exponential(rate 2): mean ½, variance ¼. + exp, err := ExponentialDraws(g, n, 2) + if err != nil { + t.Fatalf("ExponentialDraws: %v", err) + } + mean := core.Sum(exp).Float() + if math.Abs(mean/float64(n)-0.5) > 0.01 { + t.Fatalf("exponential mean = %v, want ≈ 0.5", mean/float64(n)) + } + for i := range n { + if exp.FloatAt(i) < 0 { + t.Fatalf("exponential draw %d = %v, must be positive", i, exp.FloatAt(i)) + } + } + + // Gamma(3, 2): mean 3/2, variance 3/4. + gam, err := GammaDraws(g, n, 3, 2) + if err != nil { + t.Fatalf("GammaDraws: %v", err) + } + gsum := core.Sum(gam).Float() + if math.Abs(gsum/float64(n)-1.5) > 0.03 { + t.Fatalf("gamma mean = %v, want ≈ 1.5", gsum/float64(n)) + } + + // ChiSquare(5): mean 5. + chi, _ := ChiSquareDraws(g, n, 5) + csum := core.Sum(chi).Float() + if math.Abs(csum/float64(n)-5) > 0.1 { + t.Fatalf("chi² mean = %v, want ≈ 5", csum/float64(n)) + } + + // Poisson(4): mean 4, variance 4. + pois, _ := PoissonDraws(g, n, 4) + psum := core.Sum(pois).Float() + if math.Abs(psum/float64(n)-4) > 0.1 { + t.Fatalf("poisson mean = %v, want ≈ 4", psum/float64(n)) + } + + // Binomial(20, 0.3): mean 6. + bin, _ := BinomialDraws(g, n, 20, 0.3) + bsum := core.Sum(bin).Float() + if math.Abs(bsum/float64(n)-6) > 0.1 { + t.Fatalf("binomial mean = %v, want ≈ 6", bsum/float64(n)) + } + + // StudentT(5): mean 0. + st, _ := StudentTDraws(g, n, 5) + ssum := core.Sum(st).Float() + if math.Abs(ssum/float64(n)) > 0.05 { + t.Fatalf("t mean = %v, want ≈ 0", ssum/float64(n)) + } +} + +// TestDistributionsErrors pins the parameter contracts. +func TestDistributionsErrors(t *testing.T) { + g := core.NewGenerator(1) + if _, err := ExponentialDraws(g, 1, -1); err == nil { + t.Fatal("expected an error for a negative rate") + } + if _, err := GammaDraws(g, 1, 0, 1); err == nil { + t.Fatal("expected an error for a zero shape") + } + if _, err := ChiSquareDraws(g, 1, 0); err == nil { + t.Fatal("expected an error for df = 0") + } + if _, err := StudentTDraws(g, 1, -1); err == nil { + t.Fatal("expected an error for negative df") + } + if _, err := PoissonDraws(g, 1, -1); err == nil { + t.Fatal("expected an error for a negative λ") + } + if _, err := BinomialDraws(g, 1, 10, 1.5); err == nil { + t.Fatal("expected an error for p > 1") + } +} diff --git a/stats/doc.go b/stats/doc.go new file mode 100644 index 0000000..be638b1 --- /dev/null +++ b/stats/doc.go @@ -0,0 +1,90 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Package stats is the library's statistics: probability distributions, +// descriptive summaries, classical inference, and the models that fit a +// response to a design. +// +// # Distributions +// +// Every univariate law the package carries has a cumulative +// distribution function and a quantile, and the laws the library draws +// from have a matched generator on the house generator +// (tensor.Generator): the normal, exponential, gamma, chi-square, +// Student t, Poisson and binomial families, the Weibull, lognormal and +// Pareto laws, the negative binomial, and the noncentral chi-square, F +// and t distributions. The multivariate normal has a log density and +// draws through the Cholesky factor of its covariance, and the +// Dirichlet has a density, a mean, an interior mode and draws. +// +// Everything rests on three foundations: GammaLower, GammaUpper and +// BetaIncomplete, the regularised incomplete gamma and beta functions. +// A quantile inverts its CDF by a bracketed Newton iteration through +// the distribution's density, falling back to bisection whenever the +// derivative step would leave the bracket, rather than by a closed +// form; the convergence is unconditional on every monotone CDF, and +// every iteration that fails to converge inside its budget is an +// error, never a silently truncated value. A draw is deterministic +// for a given generator state. +// +// # Descriptives and inference +// +// Median, Std, Var and VarSample are the summaries, with the robust +// pair MedianAbsoluteDeviation and TrimmedMean, the sampling +// Quantile, the Histogram family, and the windowed RollingMean, +// RollingSum, RollingMin and RollingMax. +// +// The inference entries are classical frequentist tooling built on the +// package's own distribution functions, so a p-value travels no further +// than the incomplete gamma and beta above: the association matrices +// CovarianceMatrix and CorrelationMatrix, the group comparisons +// WelchTTest, ANOVAOneWay and MannWhitneyU, the goodness-of-fit test +// ChiSquareGoodnessOfFit, the two-sample KolmogorovSmirnovTest, the +// contingency table tests FisherExactTest, ChiSquareIndependence and +// McNemarTest with Cramér's V, the percentile BootstrapCI, the rank +// correlations SpearmanRho and KendallTau beside Pearson, and the +// multiple-testing corrections Bonferroni, Holm and BenjaminiHochberg. +// +// # Models +// +// LinearRegression is ordinary least squares with the full classical +// inference (standard errors, t-tests, R², adjusted R² and the model F +// test), WeightedLinearRegression its weighted counterpart, and +// LogisticRegression and PoissonRegression the generalised linear +// models, fitted by Newton-Raphson on the exact likelihood with Wald +// inference from the inverse Fisher information. LinearMixedModel adds +// the grouped random effects, estimating their covariance and the +// residual variance by residual maximum likelihood. +// +// The regularised family is Lasso, ElasticNet and LassoPath, fitted by +// coordinate descent on the standardised design; the robust family is +// HuberRegression (with HuberRegressionTuned for the tuning constant) +// and TheilSenRegression; QuantileRegression fits the tau-th +// conditional quantile by the Frisch-Newton interior-point method. +// PCA rotates an observation cloud onto its principal components and +// carries the whitening transforms between the two representations, +// KMeans and GaussianMixture (with GaussianMixtureBIC) partition or +// model it, HierarchicalClustering records the full merge tree of the +// agglomerative construction for cutting afterwards, +// FitHiddenMarkovModel fits a hidden Markov model over discrete +// sequences with Forward, Smooth and Viterbi answering the filtered +// and smoothed posteriors and the most likely path of a fitted or +// hand-built model, and GaussianProcessRegression conditions the prior +// a Kernel defines on the observations, with MarginalLogLikelihood +// exposed as the objective of a hyperparameter fit. +// +// # Contracts +// +// Every entry point returns a value with an error, and every error +// carries the library's "tensor: " prefix and names the entry point +// that raised it. The estimation entries refuse complex input and, +// by name, any non-finite observation: a single NaN would otherwise +// spread silently through a whole result. Arrays are built through the +// root package's constructors (tensor.FromFloats and its siblings), and +// a regression design carries n rows and p columns with the intercept +// included by the caller as a constant column when one is wanted. +// +// The package depends only on the library's own internal packages, so a +// caller that fits the hyperparameters of a Gaussian process drives +// MarginalLogLikelihood from outside with the house minimiser. +package stats diff --git a/stats/dtypes_census_test.go b/stats/dtypes_census_test.go new file mode 100644 index 0000000..db2aefd --- /dev/null +++ b/stats/dtypes_census_test.go @@ -0,0 +1,710 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The dtype census for stats: every array-taking public entry probed +// with Bool, the narrow integers and the Int anchor against a float64 +// baseline carrying exactly the widened probe values. The statistics +// entries widen through accessor walks by design; the ordering family +// rides the core Sort deferral; nothing panics or silently misreads. + +var stDtypes = []core.Dtype{core.Bool, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32, core.Int} + +type stMaker func(vals []float64, shape ...int) *core.Array + +func stCast(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 stMakers(t *testing.T, dt core.Dtype) (probe, base stMaker) { + t.Helper() + makeCast := func(vals []float64) []float64 { + out := make([]float64, len(vals)) + for i, v := range vals { + out[i] = stCast(dt, v) + } + return out + } + probe = func(vals []float64, shape ...int) *core.Array { + cast := makeCast(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(makeCast(vals), shape...) + if err != nil { + t.Fatalf("baseline maker: %v", err) + } + return a + } + return probe, base +} + +// stElems reads an array through the widening accessors. +func stElems(t *testing.T, a *core.Array) []complex128 { + t.Helper() + out := make([]complex128, a.Len()) + for i := range out { + if a.Dtype() == core.Complex { + out[i] = a.ComplexAt(i) + continue + } + out[i] = complex(a.FloatAt(i), 0) + } + return out +} + +// stArrays pins array outputs: paired value errors, or equal dtype and +// bit-identical widened values. +func stArrays(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 := stElems(t, p), stElems(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]) + } + } + } +} + +// stFloat pins a float64 result against the baseline's. +func stFloat(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 pv != bv { + t.Fatalf("%s(%s) = %v, want the baseline %v", label, dt, pv, bv) + } +} + +// stInt pins an int result against the baseline's. +func stInt(t *testing.T, label string, dt core.Dtype, pv int, perr error, bv int, 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 pv != bv { + t.Fatalf("%s(%s) = %d, want the baseline %d", label, dt, pv, bv) + } +} + +// stFloats pins one float64 slice against the baseline's. +func stFloats(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]) + } + } +} + +// stWantErr pins a refusal carrying the given fragments. +func stWantErr(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) + } + } +} + +// TestDtypesCensusStats probes every array-taking public entry. +func TestDtypesCensusStats(t *testing.T) { + sample := []float64{1, 2, 3, 4, 5, 6} + xcol := []float64{1, 2, 3, 4, 5, 6} + ys := []float64{2, 3, 5, 7, 11, 13} + ybin := []float64{0, 1, 0, 1, 1, 0} + xmat := []float64{1, 2, 2, 3, 3, 1, 4, 5, 5, 4, 6, 6} + rows := []struct { + name string + run func(t *testing.T, probe, base stMaker, dt core.Dtype) + }{ + {"Median", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := Median(probe(sample, 6)) + b, berr := Median(base(sample, 6)) + stFloat(t, "Median", dt, p, perr, b, berr) + }}, + {"Std", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := Std(probe(sample, 6)) + b, berr := Std(base(sample, 6)) + stFloat(t, "Std", dt, p, perr, b, berr) + }}, + {"Var", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := Var(probe(sample, 6)) + b, berr := Var(base(sample, 6)) + stFloat(t, "Var", dt, p, perr, b, berr) + }}, + {"VarSample", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := VarSample(probe(sample, 6)) + b, berr := VarSample(base(sample, 6)) + stFloat(t, "VarSample", dt, p, perr, b, berr) + }}, + {"Histogram", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + // The sample spans no exact bin boundary at three bins, so + // the int exact-integer binning and the widened float + // binning agree element for element. + pc, pe, perr := Histogram(probe(sample, 6), 3) + bc, be, berr := Histogram(base(sample, 6), 3) + stArrays(t, "Histogram", dt, []*core.Array{pc, pe}, perr, []*core.Array{bc, be}, berr) + }}, + {"BinCounts", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := BinCounts(probe(sample, 6), 3) + b, berr := BinCounts(base(sample, 6), 3) + stArrays(t, "BinCounts", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + // Quantile orders the sample, and ordering is a core narrow-dtype + // deferral: Int computes (Sort accepts it), the narrow widths + // and bool ride the loud Sort refusal through. + {"Quantile", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := Quantile(probe(sample, 6), []float64{0.25, 0.5, 0.75}) + b, berr := Quantile(base(sample, 6), []float64{0.25, 0.5, 0.75}) + if dt == core.Int { + stArrays(t, "Quantile", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + return + } + if berr != nil { + t.Fatalf("Quantile float baseline: %v", berr) + } + stWantErr(t, "Quantile/"+dt.String(), perr, "Sort", dt.String(), "convert with Astype") + }}, + {"RollingMean", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := RollingMean(probe(sample, 6), 3) + b, berr := RollingMean(base(sample, 6), 3) + stArrays(t, "RollingMean", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"RollingSum", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := RollingSum(probe(sample, 6), 3) + b, berr := RollingSum(base(sample, 6), 3) + stArrays(t, "RollingSum", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"RollingMax", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := RollingMax(probe(sample, 6), 3) + b, berr := RollingMax(base(sample, 6), 3) + stArrays(t, "RollingMax", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"RollingMin", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := RollingMin(probe(sample, 6), 3) + b, berr := RollingMin(base(sample, 6), 3) + stArrays(t, "RollingMin", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"ANOVAOneWay", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pf, pp, perr := ANOVAOneWay([]*core.Array{probe([]float64{1, 2, 3}, 3), probe([]float64{4, 5, 6}, 3), probe([]float64{2, 4, 6}, 3)}) + bf, bp, berr := ANOVAOneWay([]*core.Array{base([]float64{1, 2, 3}, 3), base([]float64{4, 5, 6}, 3), base([]float64{2, 4, 6}, 3)}) + stFloat(t, "ANOVAOneWay F", dt, pf, perr, bf, berr) + stFloat(t, "ANOVAOneWay p", dt, pp, nil, bp, nil) + }}, + {"FisherExactTest", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pp, _, perr := FisherExactTest(probe([]float64{12, 5, 7, 10}, 2, 2), TwoSided) + bp, _, berr := FisherExactTest(base([]float64{12, 5, 7, 10}, 2, 2), TwoSided) + stFloat(t, "FisherExactTest", dt, pp, perr, bp, berr) + }}, + {"ChiSquareIndependence", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pchi, pdf, pp, perr := ChiSquareIndependence(probe([]float64{21, 8, 3, 9, 15, 11, 4, 12, 30}, 3, 3)) + bchi, bdf, bp, berr := ChiSquareIndependence(base([]float64{21, 8, 3, 9, 15, 11, 4, 12, 30}, 3, 3)) + stFloat(t, "ChiSquareIndependence chi2", dt, pchi, perr, bchi, berr) + stInt(t, "ChiSquareIndependence df", dt, pdf, perr, bdf, berr) + stFloat(t, "ChiSquareIndependence p", dt, pp, nil, bp, nil) + }}, + {"CramersV", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pv, perr := CramersV(probe([]float64{12, 5, 7, 10}, 2, 2)) + bv, berr := CramersV(base([]float64{12, 5, 7, 10}, 2, 2)) + stFloat(t, "CramersV", dt, pv, perr, bv, berr) + }}, + {"McNemarTest", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pp, perr := McNemarTest(probe([]float64{10, 5, 15, 12}, 2, 2)) + bp, berr := McNemarTest(base([]float64{10, 5, 15, 12}, 2, 2)) + stFloat(t, "McNemarTest", dt, pp, perr, bp, berr) + }}, + {"LinearMixedModel", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + run := func(m stMaker) ([]float64, float64, error) { + res, err := LinearMixedModel(m([]float64{1, 2, 3, 4, 5, 6}, 6), + m([]float64{1, 1, 1, 1, 1, 1}, 6, 1), + m([]float64{1, 1, 1, 1, 1, 1}, 6, 1), + []int{0, 0, 0, 1, 1, 1}) + if err != nil { + return nil, 0, err + } + return res.Coefficients, res.ResidualVariance, nil + } + pc, pv, perr := run(probe) + bc, bv, berr := run(base) + stFloats(t, "LinearMixedModel coefficients", dt, pc, perr, bc, berr) + stFloat(t, "LinearMixedModel sigma2", dt, pv, perr, bv, berr) + }}, + {"HierarchicalClustering", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + run := func(m stMaker) ([]float64, []int, error) { + d, err := HierarchicalClustering(m([]float64{1, 2, 3, 7, 8, 12}, 6, 1), SingleLinkage) + if err != nil { + return nil, nil, err + } + labels, err := d.Cut(3) + return d.Heights, labels, err + } + ph, pl, perr := run(probe) + bh, bl, berr := run(base) + stFloats(t, "HierarchicalClustering heights", dt, ph, perr, bh, berr) + if (pl == nil) != (bl == nil) { + t.Fatalf("HierarchicalClustering(%s): one cut answered nil", dt) + } + if pl != nil && berr == nil { + for i := range pl { + if pl[i] != bl[i] { + t.Fatalf("HierarchicalClustering(%s): cut label %d = %d, want %d", dt, i, pl[i], bl[i]) + } + } + } + }}, + {"MannWhitneyU", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pu, pp, perr := MannWhitneyU(probe([]float64{1, 2, 3, 4}, 4), probe([]float64{3, 4, 5, 6}, 4)) + bu, bp, berr := MannWhitneyU(base([]float64{1, 2, 3, 4}, 4), base([]float64{3, 4, 5, 6}, 4)) + stFloat(t, "MannWhitneyU u", dt, pu, perr, bu, berr) + stFloat(t, "MannWhitneyU p", dt, pp, nil, bp, nil) + }}, + {"WelchTTest", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pt, pdf, pp, perr := WelchTTest(probe([]float64{1, 2, 3, 4}, 4), probe([]float64{3, 4, 5, 6}, 4)) + bt, bdf, bp, berr := WelchTTest(base([]float64{1, 2, 3, 4}, 4), base([]float64{3, 4, 5, 6}, 4)) + stFloat(t, "WelchTTest t", dt, pt, perr, bt, berr) + stFloat(t, "WelchTTest df", dt, pdf, nil, bdf, nil) + stFloat(t, "WelchTTest p", dt, pp, nil, bp, nil) + }}, + {"ChiSquareGoodnessOfFit", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pc, pdf, pp, perr := ChiSquareGoodnessOfFit(probe([]float64{5, 6, 7, 8}, 4), probe([]float64{6, 6, 6, 6}, 4)) + bc, bdf, bp, berr := ChiSquareGoodnessOfFit(base([]float64{5, 6, 7, 8}, 4), base([]float64{6, 6, 6, 6}, 4)) + stFloat(t, "ChiSquareGoodnessOfFit chi2", dt, pc, perr, bc, berr) + stInt(t, "ChiSquareGoodnessOfFit df", dt, pdf, nil, bdf, nil) + stFloat(t, "ChiSquareGoodnessOfFit p", dt, pp, nil, bp, nil) + }}, + {"KolmogorovSmirnovTest", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pd, pp, perr := KolmogorovSmirnovTest(probe(sample, 6), probe([]float64{1, 3, 3, 5, 5, 7}, 6)) + bd, bp, berr := KolmogorovSmirnovTest(base(sample, 6), base([]float64{1, 3, 3, 5, 5, 7}, 6)) + stFloat(t, "KolmogorovSmirnovTest d", dt, pd, perr, bd, berr) + stFloat(t, "KolmogorovSmirnovTest p", dt, pp, nil, bp, nil) + }}, + {"BootstrapCI", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + mean := func(a *core.Array) (float64, error) { return core.Mean(a) } + pl, pu, perr := BootstrapCI(probe(sample, 6), mean, 0.95, 50, 11) + bl, bu, berr := BootstrapCI(base(sample, 6), mean, 0.95, 50, 11) + stFloat(t, "BootstrapCI lower", dt, pl, perr, bl, berr) + stFloat(t, "BootstrapCI upper", dt, pu, nil, bu, nil) + }}, + {"CovarianceMatrix", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := CovarianceMatrix(probe(xmat, 6, 2)) + b, berr := CovarianceMatrix(base(xmat, 6, 2)) + stArrays(t, "CovarianceMatrix", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"CorrelationMatrix", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := CorrelationMatrix(probe(xmat, 6, 2)) + b, berr := CorrelationMatrix(base(xmat, 6, 2)) + stArrays(t, "CorrelationMatrix", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"Histogram2D", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pc, px, py, perr := Histogram2D(probe(sample, 6), probe([]float64{1, 1, 2, 2, 3, 3}, 6), 3, 3) + bc, bx, by, berr := Histogram2D(base(sample, 6), base([]float64{1, 1, 2, 2, 3, 3}, 6), 3, 3) + if berr != nil || perr != nil { + stFloat(t, "Histogram2D", dt, 0, perr, 0, berr) + return + } + stArrays(t, "Histogram2D counts", dt, []*core.Array{pc}, nil, []*core.Array{bc}, nil) + stFloats(t, "Histogram2D x edges", dt, px, nil, bx, nil) + stFloats(t, "Histogram2D y edges", dt, py, nil, by, nil) + }}, + {"KernelDensity", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := KernelDensity(probe(sample, 6), 1.0, probe([]float64{1.5, 2.5, 3.5}, 3)) + b, berr := KernelDensity(base(sample, 6), 1.0, base([]float64{1.5, 2.5, 3.5}, 3)) + stArrays(t, "KernelDensity", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"SpearmanRho", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := SpearmanRho(probe(sample, 6), probe([]float64{2, 1, 4, 3, 6, 5}, 6)) + b, berr := SpearmanRho(base(sample, 6), base([]float64{2, 1, 4, 3, 6, 5}, 6)) + stFloat(t, "SpearmanRho", dt, p, perr, b, berr) + }}, + {"KendallTau", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := KendallTau(probe(sample, 6), probe([]float64{2, 1, 4, 3, 6, 5}, 6)) + b, berr := KendallTau(base(sample, 6), base([]float64{2, 1, 4, 3, 6, 5}, 6)) + stFloat(t, "KendallTau", dt, p, perr, b, berr) + }}, + {"MedianAbsoluteDeviation", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := MedianAbsoluteDeviation(probe(sample, 6)) + b, berr := MedianAbsoluteDeviation(base(sample, 6)) + stFloat(t, "MedianAbsoluteDeviation", dt, p, perr, b, berr) + }}, + {"TrimmedMean", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := TrimmedMean(probe(sample, 6), 0.2) + b, berr := TrimmedMean(base(sample, 6), 0.2) + stFloat(t, "TrimmedMean", dt, p, perr, b, berr) + }}, + {"LinearRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := LinearRegression(probe(xcol, 6, 1), probe(ys, 6)) + br, berr := LinearRegression(base(xcol, 6, 1), base(ys, 6)) + if berr != nil || perr != nil { + stFloat(t, "LinearRegression", dt, 0, perr, 0, berr) + return + } + stFloats(t, "LinearRegression coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + stFloats(t, "LinearRegression fitted", dt, pr.Fitted, nil, br.Fitted, nil) + stFloat(t, "LinearRegression R2", dt, pr.RSquared, nil, br.RSquared, nil) + }}, + {"WeightedLinearRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := WeightedLinearRegression(probe(xcol, 6, 1), probe(ys, 6), probe([]float64{1, 2, 1, 2, 1, 2}, 6)) + br, berr := WeightedLinearRegression(base(xcol, 6, 1), base(ys, 6), base([]float64{1, 2, 1, 2, 1, 2}, 6)) + if berr != nil || perr != nil { + stFloat(t, "WeightedLinearRegression", dt, 0, perr, 0, berr) + return + } + stFloats(t, "WeightedLinearRegression coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + stFloats(t, "WeightedLinearRegression fitted", dt, pr.Fitted, nil, br.Fitted, nil) + }}, + {"PoissonRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := PoissonRegression(probe(xcol, 6, 1), probe([]float64{1, 2, 1, 3, 4, 5}, 6)) + br, berr := PoissonRegression(base(xcol, 6, 1), base([]float64{1, 2, 1, 3, 4, 5}, 6)) + if berr != nil || perr != nil { + stFloat(t, "PoissonRegression", dt, 0, perr, 0, berr) + return + } + stFloats(t, "PoissonRegression coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + stFloats(t, "PoissonRegression fitted", dt, pr.Fitted, nil, br.Fitted, nil) + }}, + {"LogisticRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := LogisticRegression(probe(xcol, 6, 1), probe(ybin, 6)) + br, berr := LogisticRegression(base(xcol, 6, 1), base(ybin, 6)) + if berr != nil || perr != nil { + stFloat(t, "LogisticRegression", dt, 0, perr, 0, berr) + return + } + stFloats(t, "LogisticRegression coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + stFloats(t, "LogisticRegression fitted", dt, pr.Fitted, nil, br.Fitted, nil) + }}, + {"QuantileRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := QuantileRegression(probe(xcol, 6, 1), probe(ys, 6), 0.5) + br, berr := QuantileRegression(base(xcol, 6, 1), base(ys, 6), 0.5) + if berr != nil || perr != nil { + stFloat(t, "QuantileRegression", dt, 0, perr, 0, berr) + return + } + stFloats(t, "QuantileRegression coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + stFloat(t, "QuantileRegression check loss", dt, pr.CheckLoss, nil, br.CheckLoss, nil) + }}, + {"HuberRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := HuberRegression(probe(xcol, 6, 1), probe(ys, 6)) + br, berr := HuberRegression(base(xcol, 6, 1), base(ys, 6)) + if berr != nil || perr != nil { + stFloat(t, "HuberRegression", dt, 0, perr, 0, berr) + return + } + stFloats(t, "HuberRegression coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + stFloats(t, "HuberRegression weights", dt, pr.Weights, nil, br.Weights, nil) + stFloat(t, "HuberRegression scale", dt, pr.Scale, nil, br.Scale, nil) + }}, + {"HuberRegressionTuned", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := HuberRegressionTuned(probe(xcol, 6, 1), probe(ys, 6), 1.5) + br, berr := HuberRegressionTuned(base(xcol, 6, 1), base(ys, 6), 1.5) + if berr != nil || perr != nil { + stFloat(t, "HuberRegressionTuned", dt, 0, perr, 0, berr) + return + } + stFloats(t, "HuberRegressionTuned coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + }}, + {"TheilSenRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pi, ps, perr := TheilSenRegression(probe(xcol, 6, 1), probe(ys, 6)) + bi, bs, berr := TheilSenRegression(base(xcol, 6, 1), base(ys, 6)) + stFloat(t, "TheilSen intercept", dt, pi, perr, bi, berr) + stFloat(t, "TheilSen slope", dt, ps, nil, bs, nil) + }}, + {"Lasso", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := Lasso(probe(xmat, 6, 2), probe(ys, 6), 0.1) + br, berr := Lasso(base(xmat, 6, 2), base(ys, 6), 0.1) + if berr != nil || perr != nil { + stFloat(t, "Lasso", dt, 0, perr, 0, berr) + return + } + stFloats(t, "Lasso coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + stFloat(t, "Lasso intercept", dt, pr.Intercept, nil, br.Intercept, nil) + }}, + {"ElasticNet", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := ElasticNet(probe(xmat, 6, 2), probe(ys, 6), 0.1, 0.5) + br, berr := ElasticNet(base(xmat, 6, 2), base(ys, 6), 0.1, 0.5) + if berr != nil || perr != nil { + stFloat(t, "ElasticNet", dt, 0, perr, 0, berr) + return + } + stFloats(t, "ElasticNet coefficients", dt, pr.Coefficients, nil, br.Coefficients, nil) + }}, + {"LassoPath", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := LassoPath(probe(xmat, 6, 2), probe(ys, 6), 0.5) + br, berr := LassoPath(base(xmat, 6, 2), base(ys, 6), 0.5) + if berr != nil || perr != nil { + stFloat(t, "LassoPath", dt, 0, perr, 0, berr) + return + } + stFloats(t, "LassoPath lambdas", dt, pr.Lambdas, nil, br.Lambdas, nil) + stFloats(t, "LassoPath intercepts", dt, pr.Intercepts, nil, br.Intercepts, nil) + }}, + {"PCA", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := PCA(probe(xmat, 6, 2)) + br, berr := PCA(base(xmat, 6, 2)) + if berr != nil || perr != nil { + stFloat(t, "PCA", dt, 0, perr, 0, berr) + return + } + stFloats(t, "PCA mean", dt, pr.Mean, nil, br.Mean, nil) + stFloats(t, "PCA explained variance", dt, pr.ExplainedVariance, nil, br.ExplainedVariance, nil) + stFloats(t, "PCA variance ratio", dt, pr.ExplainedVarianceRatio, nil, br.ExplainedVarianceRatio, nil) + stArrays(t, "PCA loadings", dt, []*core.Array{pr.Loadings, pr.Scores}, nil, + []*core.Array{br.Loadings, br.Scores}, nil) + }}, + {"KMeans", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := KMeans(core.NewGenerator(11), probe(xmat, 6, 2), 2) + br, berr := KMeans(core.NewGenerator(11), base(xmat, 6, 2), 2) + if berr != nil || perr != nil { + stFloat(t, "KMeans", dt, 0, perr, 0, berr) + return + } + if len(pr.Centres) != len(br.Centres) { + t.Fatalf("KMeans(%s): %d centres, want %d", dt, len(pr.Centres), len(br.Centres)) + } + for c := range pr.Centres { + stFloats(t, "KMeans centre", dt, pr.Centres[c], nil, br.Centres[c], nil) + } + for i := range pr.Labels { + if pr.Labels[i] != br.Labels[i] { + t.Fatalf("KMeans(%s): label %d = %d, want %d", dt, i, pr.Labels[i], br.Labels[i]) + } + } + stFloat(t, "KMeans inertia", dt, pr.Inertia, nil, br.Inertia, nil) + }}, + {"GaussianMixture", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := GaussianMixture(core.NewGenerator(11), probe(xmat, 6, 2), 2) + br, berr := GaussianMixture(core.NewGenerator(11), base(xmat, 6, 2), 2) + if berr != nil || perr != nil { + stFloat(t, "GaussianMixture", dt, 0, perr, 0, berr) + return + } + stFloats(t, "GaussianMixture weights", dt, pr.Weights, nil, br.Weights, nil) + stFloat(t, "GaussianMixture log likelihood", dt, pr.LogLikelihood, nil, br.LogLikelihood, nil) + stFloat(t, "GaussianMixture BIC", dt, pr.BIC, nil, br.BIC, nil) + }}, + {"GaussianMixtureBIC", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + pr, perr := GaussianMixtureBIC(core.NewGenerator(11), probe(xmat, 6, 2), 2) + br, berr := GaussianMixtureBIC(core.NewGenerator(11), base(xmat, 6, 2), 2) + if berr != nil || perr != nil { + stFloat(t, "GaussianMixtureBIC", dt, 0, perr, 0, berr) + return + } + stFloats(t, "GaussianMixtureBIC grid", dt, pr.BICGrid, nil, br.BICGrid, nil) + }}, + {"GaussianProcessRegression", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + kp, kerr := SquaredExponentialKernel(1.5) + if kerr != nil { + t.Fatalf("SquaredExponentialKernel: %v", kerr) + } + pr, perr := GaussianProcessRegression(kp, probe([]float64{1, 2, 3, 4}, 4, 1), + probe([]float64{1, 2, 3, 4}, 4), 0.1, probe([]float64{1, 2}, 2, 1)) + br, berr := GaussianProcessRegression(kp, base([]float64{1, 2, 3, 4}, 4, 1), + base([]float64{1, 2, 3, 4}, 4), 0.1, base([]float64{1, 2}, 2, 1)) + if berr != nil || perr != nil { + stFloat(t, "GaussianProcessRegression", dt, 0, perr, 0, berr) + return + } + stFloats(t, "GaussianProcessRegression mean", dt, pr.Mean, nil, br.Mean, nil) + stFloats(t, "GaussianProcessRegression covariance", dt, pr.Covariance, nil, br.Covariance, nil) + stFloat(t, "GaussianProcessRegression log likelihood", dt, pr.LogLikelihood, nil, br.LogLikelihood, nil) + }}, + {"MarginalLogLikelihood", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + kp, kerr := SquaredExponentialKernel(1.5) + if kerr != nil { + t.Fatalf("SquaredExponentialKernel: %v", kerr) + } + p, perr := MarginalLogLikelihood(kp, probe([]float64{1, 2, 3, 4}, 4, 1), probe([]float64{1, 2, 3, 4}, 4), 0.1) + b, berr := MarginalLogLikelihood(kp, base([]float64{1, 2, 3, 4}, 4, 1), base([]float64{1, 2, 3, 4}, 4), 0.1) + stFloat(t, "MarginalLogLikelihood", dt, p, perr, b, berr) + }}, + {"MultivariateNormalLogDensity", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := MultivariateNormalLogDensity(probe([]float64{1, 2}, 2), probe([]float64{2, 0, 0, 1}, 2, 2), probe([]float64{1, 3}, 2)) + b, berr := MultivariateNormalLogDensity(base([]float64{1, 2}, 2), base([]float64{2, 0, 0, 1}, 2, 2), base([]float64{1, 3}, 2)) + stFloat(t, "MultivariateNormalLogDensity", dt, p, perr, b, berr) + }}, + {"MultivariateNormalDraws", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + p, perr := MultivariateNormalDraws(core.NewGenerator(11), 3, probe([]float64{1, 2}, 2), probe([]float64{2, 0, 0, 1}, 2, 2)) + b, berr := MultivariateNormalDraws(core.NewGenerator(11), 3, base([]float64{1, 2}, 2), base([]float64{2, 0, 0, 1}, 2, 2)) + stArrays(t, "MultivariateNormalDraws", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + // The PCA transform surface reads its operand through the + // accessor walks whatever dtype it carries; the fit itself is + // setup built from uncast float data, not the surface under + // test. + {"PCAResult.Whiten", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + raw, err := core.FromFloats(xmat, 6, 2) + if err != nil { + t.Fatal(err) + } + fit, ferr := PCA(raw) + if ferr != nil { + t.Fatalf("PCA setup: %v", ferr) + } + p, perr := fit.Whiten(probe(xmat, 6, 2)) + b, berr := fit.Whiten(base(xmat, 6, 2)) + stArrays(t, "PCAResult.Whiten", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + {"PCAResult.Unwhiten", func(t *testing.T, probe, base stMaker, dt core.Dtype) { + raw, err := core.FromFloats(xmat, 6, 2) + if err != nil { + t.Fatal(err) + } + fit, ferr := PCA(raw) + if ferr != nil { + t.Fatalf("PCA setup: %v", ferr) + } + p, perr := fit.Unwhiten(probe(xmat, 6, 2)) + b, berr := fit.Unwhiten(base(xmat, 6, 2)) + stArrays(t, "PCAResult.Unwhiten", dt, []*core.Array{p}, perr, []*core.Array{b}, berr) + }}, + } + for _, row := range rows { + for _, dt := range stDtypes { + t.Run(row.name+"/"+dt.String(), func(t *testing.T) { + probe, base := stMakers(t, dt) + row.run(t, probe, base, dt) + }) + } + } +} diff --git a/stats/estimation_guards_pin_test.go b/stats/estimation_guards_pin_test.go new file mode 100644 index 0000000..b56f67b --- /dev/null +++ b/stats/estimation_guards_pin_test.go @@ -0,0 +1,152 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Estimation guards and exactness pins: silent NaN estimates, a +// perfect-significance null model, and int64 samples that round +// through float64. + +// TestCovarianceMatrixRejectsNaN: one NaN observation flowed into the +// column means and produced an all-NaN matrix with a nil error. +func TestCovarianceMatrixRejectsNaN(t *testing.T) { + obs, err := core.FromFloats([]float64{1, 2, math.NaN(), 4, 5, 6}, 3, 2) + if err != nil { + t.Fatal(err) + } + if _, err := CovarianceMatrix(obs); err == nil || !strings.Contains(err.Error(), "CovarianceMatrix") { + t.Fatalf("CovarianceMatrix on a NaN observation: err = %v", err) + } + if _, err := CorrelationMatrix(obs); err == nil || !strings.Contains(err.Error(), "CorrelationMatrix") { + t.Fatalf("CorrelationMatrix on a NaN observation: err = %v", err) + } +} + +// TestRegressionInterceptOnly: an intercept-only design skipped the F +// block and left FPValue at its zero value, reporting the null model +// as maximally significant. +func TestRegressionInterceptOnly(t *testing.T) { + y, err := core.FromFloats([]float64{2, 4, 6, 8}, 4) + if err != nil { + t.Fatal(err) + } + one, err := core.FromFloats([]float64{1, 1, 1, 1}, 4, 1) + if err != nil { + t.Fatal(err) + } + for name, run := range map[string]func() (*LinearRegressionResult, error){ + "LinearRegression": func() (*LinearRegressionResult, error) { return LinearRegression(one, y) }, + "WeightedLinearRegression": func() (*LinearRegressionResult, error) { + w, _ := core.FromFloats([]float64{1, 1, 1, 1}, 4) + return WeightedLinearRegression(one, y, w) + }, + } { + res, err := run() + if err != nil { + t.Fatalf("%s: %v", name, err) + } + if res.FPValue != 1 { + t.Fatalf("%s: intercept-only FPValue = %g, want 1", name, res.FPValue) + } + if res.FStatistic != 0 { + t.Fatalf("%s: intercept-only FStatistic = %g, want 0", name, res.FStatistic) + } + } +} + +// TestHistogramInt64Exact: the samples 2^53, 2^53+1, 2^53+2 all +// widened to the same float64 and landed in the wrong bins; the exact +// integer path bins them by their true values. +func TestHistogramInt64Exact(t *testing.T) { + const base = 1 << 53 + x, err := core.FromInts([]int64{base + 1, base, base + 2}, 3) + if err != nil { + t.Fatal(err) + } + counts, _, err := Histogram(x, 2) + if err != nil { + t.Fatal(err) + } + if c0, _ := core.IntAt(counts, 0); c0 != 1 { + t.Fatalf("bin 0 count = %d, want 1 (only the minimum)", c0) + } + if c1, _ := core.IntAt(counts, 1); c1 != 2 { + t.Fatalf("bin 1 count = %d, want 2", c1) + } +} + +// TestMedianQuantileInt64Exact: the median of {2^53, 2^53+1} came +// back as 2^53 although 2^53+0.5 is representable; the same held for +// the quantile at 0.5, whose interpolation collapsed on the rounded +// pair. +func TestMedianQuantileInt64Exact(t *testing.T) { + const base = 1 << 53 + x, err := core.FromInts([]int64{base, base + 1}, 2) + if err != nil { + t.Fatal(err) + } + m, err := Median(x) + if err != nil { + t.Fatal(err) + } + if m != base+0.5 { + t.Fatalf("Median = %g, want %g", m, base+0.5) + } + q, err := Quantile(x, []float64{0.5}) + if err != nil { + t.Fatal(err) + } + if q.FloatAt(0) != base+0.5 { + t.Fatalf("Quantile(0.5) = %g, want %g", q.FloatAt(0), base+0.5) + } +} + +// TestGammaUpperOwnErrorName: GammaUpper reported the GammaLower +// prefix on the shape refusal. +func TestGammaUpperOwnErrorName(t *testing.T) { + if _, err := GammaUpper(0, 1); err == nil || !strings.Contains(err.Error(), "GammaUpper:") { + t.Fatalf("GammaUpper(0, 1): err = %v", err) + } +} + +// TestWelchTTestRejectsNaN: a NaN sample only surfaced when the tail +// function rejected the NaN statistic, under its own name. +func TestWelchTTestRejectsNaN(t *testing.T) { + a, err := core.FromFloats([]float64{math.NaN(), 2, 3}, 3) + if err != nil { + t.Fatal(err) + } + b, err := core.FromFloats([]float64{1, 2, 3}, 3) + if err != nil { + t.Fatal(err) + } + if _, _, _, err := WelchTTest(a, b); err == nil || !strings.Contains(err.Error(), "WelchTTest") { + t.Fatalf("WelchTTest on a NaN sample: err = %v", err) + } +} + +// TestMannWhitneyUTieOverflow: a tie block of 2^21+1 observations +// overflowed the integer tie term, wrapped it negative and bypassed +// the every-observation-tied refusal. +func TestMannWhitneyUTieOverflow(t *testing.T) { + const n = 1<<21 + 1 + vals := make([]float64, n) + for i := range vals { + vals[i] = 7 + } + a, err := core.FromFloats(vals, n) + if err != nil { + t.Fatal(err) + } + if _, _, err := MannWhitneyU(a, a); err == nil || !strings.Contains(err.Error(), "tied") { + t.Fatalf("MannWhitneyU on one giant tie block: err = %v", err) + } +} diff --git a/stats/example_test.go b/stats/example_test.go new file mode 100644 index 0000000..a0fb366 --- /dev/null +++ b/stats/example_test.go @@ -0,0 +1,156 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats_test + +// The godoc examples: the flagship workflows of the package as +// runnable, checked snippets. pkg.go.dev renders them beside the API, +// and `go test` executes them, so the documentation cannot rot. + +import ( + "fmt" + "log" + + tensor "sourcedock.dev/petrbalvin/tensor" + "sourcedock.dev/petrbalvin/tensor/stats" +) + +// The normal quantile inverts NormalCDF: these are the two-sided +// critical values at the 2 percent and the 10 percent level. +func ExampleNormalQuantile() { + for _, q := range []float64{0.01, 0.05, 0.95, 0.99} { + z, err := stats.NormalQuantile(q) + if err != nil { + log.Fatal(err) + } + fmt.Printf("%.4f\n", z) + } + // Output: + // -2.3263 + // -1.6449 + // 1.6449 + // 2.3263 +} + +// Welch's t-test compares two independent samples without assuming +// equal variances, and reports the two-sided p-value. +func ExampleWelchTTest() { + a, _ := tensor.FromFloats([]float64{5.1, 4.9, 5.4, 5.0, 5.3}, 5) + b, _ := tensor.FromFloats([]float64{6.2, 6.0, 5.8, 6.3, 6.1}, 5) + t, df, p, err := stats.WelchTTest(a, b) + if err != nil { + log.Fatal(err) + } + fmt.Printf("t = %.3f, df = %.2f, p = %.2e\n", t, df, p) + // Output: + // t = -7.431, df = 7.96, p = 7.61e-05 +} + +// Ordinary least squares with the full classical inference: the +// coefficients, their standard errors, R² and the model F test. +func ExampleLinearRegression() { + // y = 4 + 3x over x = 0..5, the intercept carried as the constant + // first column, as every regression entry point of the package + // expects it. + design, _ := tensor.FromFloats([]float64{ + 1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, + }, 6, 2) + y, _ := tensor.FromFloats([]float64{4.1, 7.0, 9.9, 13.2, 15.9, 19.1}, 6) + fit, err := stats.LinearRegression(design, y) + if err != nil { + log.Fatal(err) + } + fmt.Printf("intercept %.3f ± %.3f (p = %.2e)\n", + fit.Coefficients[0], fit.StandardErrors[0], fit.PValues[0]) + fmt.Printf("slope %.3f ± %.3f (p = %.2e)\n", + fit.Coefficients[1], fit.StandardErrors[1], fit.PValues[1]) + fmt.Printf("R² %.4f, F = %.1f on (%d, %d) df\n", + fit.RSquared, fit.FStatistic, fit.DModel, fit.DResidual) + // Output: + // intercept 4.033 ± 0.098 (p = 2.08e-06) + // slope 3.000 ± 0.032 (p = 8.12e-08) + // R² 0.9995, F = 8590.9 on (1, 4) df +} + +// Logistic regression fits a binary response by maximum likelihood +// through the logit link and reports Wald inference. +func ExampleLogisticRegression() { + design, _ := tensor.FromFloats([]float64{ + 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, 1, 6, 1, 7, 1, 8, 1, 9, 1, 10, + }, 10, 2) + y, _ := tensor.FromFloats([]float64{0, 0, 1, 0, 1, 1, 0, 1, 1, 1}, 10) + fit, err := stats.LogisticRegression(design, y) + if err != nil { + log.Fatal(err) + } + fmt.Printf("intercept %.4f (SE %.4f)\n", fit.Coefficients[0], fit.StandardErrors[0]) + fmt.Printf("slope %.4f (SE %.4f)\n", fit.Coefficients[1], fit.StandardErrors[1]) + fmt.Printf("log likelihood %.4f, %d Newton steps\n", + fit.LogLikelihood, fit.Iterations) + // Output: + // intercept -2.2903 (SE 1.7992) + // slope 0.5279 (SE 0.3377) + // log likelihood -4.9014, 6 Newton steps +} + +// PCA decomposes a correlated cloud onto its principal components and +// whitens it to unit covariance. +func ExamplePCA() { + data, _ := tensor.FromFloats([]float64{ + 1, 2.1, 2, 3.9, 3, 6.2, 4, 7.8, 5, 10.1, + }, 5, 2) + fit, err := stats.PCA(data) + if err != nil { + log.Fatal(err) + } + fmt.Printf("explained variance ratio: %.4f %.4f\n", + fit.ExplainedVarianceRatio[0], fit.ExplainedVarianceRatio[1]) + whitened, err := fit.Whiten(data) + if err != nil { + log.Fatal(err) + } + back, err := fit.Unwhiten(whitened) + if err != nil { + log.Fatal(err) + } + fmt.Printf("whitened first row: %.4f %.4f\n", + whitened.FloatAt(0), whitened.FloatAt(1)) + fmt.Printf("unwhitened recovers: %.4f %.4f\n", back.FloatAt(0), back.FloatAt(1)) + // Output: + // explained variance ratio: 0.9996 0.0004 + // whitened first row: -1.2486 -0.4190 + // unwhitened recovers: 1.0000 2.1000 +} + +// KMeans partitions a sample into k clusters, deterministically for a +// given generator state. +func ExampleKMeans() { + data, _ := tensor.FromFloats([]float64{ + 0, 0, 0.2, 0, 0, 0.2, + 9, 9, 9.2, 9, 9, 9.2, + }, 6, 2) + fit, err := stats.KMeans(tensor.NewGenerator(7), data, 2) + if err != nil { + log.Fatal(err) + } + fmt.Printf("labels: %v\n", fit.Labels) + fmt.Printf("inertia: %.4f\n", fit.Inertia) + // Output: + // labels: [0 0 0 1 1 1] + // inertia: 0.1067 +} + +// KernelDensity smooths a sample into a continuous density; a +// non-positive bandwidth asks for Silverman's rule. +func ExampleKernelDensity() { + sample, _ := tensor.FromFloats([]float64{-1, 0, 0.5, 1, 1.5}, 5) + points, _ := tensor.FromFloats([]float64{-1, 0, 1}, 3) + density, err := stats.KernelDensity(sample, 0, points) + if err != nil { + log.Fatal(err) + } + fmt.Printf("%.4f %.4f %.4f\n", + density.FloatAt(0), density.FloatAt(1), density.FloatAt(2)) + // Output: + // 0.1852 0.3018 0.3772 +} diff --git a/stats/glm.go b/stats/glm.go new file mode 100644 index 0000000..79a921b --- /dev/null +++ b/stats/glm.go @@ -0,0 +1,699 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Generalised linear models. The logistic regression fits a binary +// response through the logit link by Newton-Raphson on the exact +// likelihood, which is the iteratively reweighted least squares the +// literature names, and reports Wald inference from the observed +// Fisher information. + +// LogisticRegressionResult carries the fit of a binary response. +type LogisticRegressionResult struct { + // Coefficients are the maximum-likelihood estimates β̂ on the + // logit scale. + Coefficients []float64 + // StandardErrors are the Wald standard errors from the inverse + // Fisher information at the optimum. + StandardErrors []float64 + // ZStatistics are β̂/SE per coefficient. + ZStatistics []float64 + // PValues are the two-sided normal-tail probabilities. + PValues []float64 + // Fitted holds the predicted probability for every sample, clamped + // into [1e-12, 1 − 1e-12] exactly as the fitting loop clamps it, so + // a far-out covariate saturates the probability without turning the + // likelihood into a logarithm of zero. + Fitted []float64 + // LogLikelihood is the maximised Bernoulli log likelihood. + LogLikelihood float64 + // Iterations counts the Newton steps taken; Converged reports + // whether the coefficient update fell under the tolerance. + Iterations int + Converged bool +} + +// logisticProbability returns the Bernoulli probability of the linear +// predictor eta, clamped away from the saturating ends: the weights and +// the logarithms the fit evaluates both need a live value there, so the +// clamp is what keeps a far-out covariate from driving the reported +// likelihood to a NaN. The fitting loop and the final pass share it, so +// the reported Fitted values are the ones the loop maximised. +func logisticProbability(eta float64) float64 { + pr := 1 / (1 + math.Exp(-eta)) + if pr < 1e-12 { + pr = 1e-12 + } + if pr > 1-1e-12 { + pr = 1 - 1e-12 + } + return pr +} + +// PoissonRegressionResult carries the fit of a count response. +type PoissonRegressionResult struct { + // Coefficients are the maximum-likelihood estimates β̂ on the log + // scale. + Coefficients []float64 + // StandardErrors are the Wald standard errors from the inverse + // Fisher information at the optimum. + StandardErrors []float64 + // ZStatistics are β̂/SE per coefficient. + ZStatistics []float64 + // PValues are the two-sided normal-tail probabilities. + PValues []float64 + // Fitted holds the predicted mean count for every sample, clamped + // into [1e-12, 1e300] exactly as the fitting loop clamps it, so a + // far-out covariate saturates the mean without turning the + // likelihood into a logarithm of zero or an infinity. + Fitted []float64 + // LogLikelihood is the maximised Poisson log likelihood, evaluated + // on the clamped Fitted values. + LogLikelihood float64 + // Iterations counts the Newton steps taken; Converged reports + // whether the coefficient update fell under the tolerance. + Iterations int + Converged bool +} + +// mirrorUpper fills the lower triangle of a symmetric matrix from the +// upper one. Each mirrored entry accumulates exactly the product chain +// the upper entry did: the row-wise product commutes bitwise and both +// entries sum the rows in the same order, so today's direct +// accumulation already holds equal bits on both sides of the diagonal +// and the copy reproduces them. +func mirrorUpper(m [][]float64) { + for i := range m { + for j := range i { + m[i][j] = m[j][i] + } + } +} + +// poissonMean returns the Poisson mean of the linear predictor eta, +// clamped away from the values where the fit cannot keep going: the +// exponential overflows above η ≈ 709.78 and underflows to zero below +// η ≈ −745, and both the y·log μ term and the Fisher weights need a +// live finite mean there. The fitting loop and the final pass share +// it, so the reported Fitted values are the ones the loop maximised. +func poissonMean(eta float64) float64 { + const ceiling = 1e300 + mu := math.Exp(eta) + if math.IsInf(mu, 1) || mu > ceiling { + mu = ceiling + } + if mu < 1e-12 { + mu = 1e-12 + } + return mu +} + +// PoissonRegression fits y = Poisson(exp(X·β)) by maximum likelihood. +// The design carries n rows and p columns exactly as LinearRegression's +// (intercept included by the caller as a constant column when wanted), +// y holds non-negative integer counts, and the fit runs Newton-Raphson +// until the largest coefficient update drops under 1e-10, halving any +// step that does not raise the likelihood: the unbounded Poisson +// weights let an undamped step overshoot into oscillation. On the +// canonical log link the observed information equals the Fisher +// information, so this is the iteratively reweighted least squares the +// literature names, with the weights equal to the means. A design +// whose Fisher information is singular, a duplicated column among +// them, is reported as an error; data that drives the iteration +// without settling exhausts the iteration budget and is reported +// rather than returned as a diverged fit. +func PoissonRegression(x, y *core.Array) (*PoissonRegressionResult, error) { + const name = "PoissonRegression" + if x.NDim() != 2 { + return nil, base.Errf("%s: the design must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 { + return nil, base.Errf("%s: the response must be rank 1", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + n, p := x.Shape()[0], x.Shape()[1] + if y.Len() != n { + return nil, base.Errf("%s: the design has %d rows but the response %d", name, n, y.Len()) + } + if n <= p { + return nil, base.Errf("%s: need n > p, got %d observations and %d columns", name, n, p) + } + if p == 0 { + return nil, base.Errf("%s: the design must carry at least one column", name) + } + // A non-finite design would flow through the exponential and the + // Newton step into a "converged" all-NaN fit: as in + // LinearRegression, non-finite input has no answer to report. + if err := checkFinite(name, "the design", x); err != nil { + return nil, err + } + for i := range y.Len() { + v := y.FloatAt(i) + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: response sample %d is not finite", name, i) + } + if v < 0 { + return nil, base.Errf("%s: response sample %d is %g, want a non-negative count", name, i, v) + } + if v != math.Trunc(v) { + return nil, base.Errf("%s: response sample %d is %g, want an integer count", name, i, v) + } + } + beta := make([]float64, p) + mean := make([]float64, n) + // The design and the response are read through raw payload slices + // when dense: the elements are the ones FloatAt returns, so every + // product and sum below keeps its exact operand bits. + fx := rawFloats(x) + fy := rawFloats(y) + // The log likelihood without the constant log(y!): every candidate + // point pays the same constant, so the comparison the step damping + // makes needs only this part. It reads the mean through the same + // clamped exponential the fit maximises. + logLikeAt := func(b []float64) float64 { + total := 0.0 + for r := range n { + eta := 0.0 + if fx != nil { + row := fx[r*p : r*p+p] + for j, xj := range row { + eta += xj * b[j] + } + } else { + for j := range p { + eta += x.FloatAt(r*p+j) * b[j] + } + } + mu := poissonMean(eta) + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + total += yv*math.Log(mu) - mu + } + return total + } + const maxIter = 100 + converged := false + iterations := maxIter + // The normal-equation buffers are allocated once and cleared per + // iteration: the accumulation adds into them, so every pass starts + // from an explicitly zeroed state, the one a fresh allocation had. + fisher := make([][]float64, p) + for i := range p { + fisher[i] = make([]float64, p) + } + gradient := make([]float64, p) + applied := make([]float64, p) + for iter := 1; iter <= maxIter; iter++ { + currentLogLike := 0.0 + for r := range n { + eta := 0.0 + if fx != nil { + row := fx[r*p : r*p+p] + for j, xj := range row { + eta += xj * beta[j] + } + } else { + for j := range p { + eta += x.FloatAt(r*p+j) * beta[j] + } + } + mean[r] = poissonMean(eta) + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + currentLogLike += yv*math.Log(mean[r]) - mean[r] + } + // The Fisher information is symmetric and each lower-triangle + // entry equals its upper twin bit for bit (mirrorUpper), so the + // accumulation runs the upper triangle alone and mirrors it once. + for i := range p { + clear(fisher[i]) + } + clear(gradient) + for r := range n { + mu := mean[r] + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + if fx != nil { + row := fx[r*p : r*p+p] + for i, xr := range row { + gradient[i] += xr * (yv - mu) + // Upper triangle, both operands pre-sliced from i: + // the same products in the same order, bounds checks + // elided. + fi := fisher[i][i:] + for j, xj := range row[i:] { + fi[j] += xr * xj * mu + } + } + } else { + for i := range p { + xr := x.FloatAt(r*p + i) + gradient[i] += xr * (yv - mu) + for j := i; j < p; j++ { + fisher[i][j] += xr * x.FloatAt(r*p+j) * mu + } + } + } + } + mirrorUpper(fisher) + step, err := base.SolveSystem(name, fisher, [][]float64{gradient}) + if err != nil { + return nil, base.Errf("%s: the Fisher information is singular (%w)", name, err) + } + // Backtracking on the log likelihood: unlike the logistic + // weights, the Poisson ones are unbounded, so an undamped step + // can overshoot into the clamped means where the next step + // explodes and the iteration oscillates instead of converging. + // The step is halved while the likelihood does not rise, the + // same damping the root finder applies to its residual norm. + // The acceptance is >=, the standard Armijo condition: a flat + // likelihood must accept the step rather than spend the halving + // budget shrinking it into what only looks like convergence. + worst := 0.0 + damping := 1.0 + for halving := 0; ; halving++ { + worst = 0.0 + for j := range p { + applied[j] = beta[j] + damping*step[0][j] + if d := math.Abs(damping * step[0][j]); d > worst { + worst = d + } + } + if logLikeAt(applied) >= currentLogLike || halving == 30 { + break + } + damping /= 2 + } + copy(beta, applied) + if worst < 1e-10 { + converged = true + iterations = iter + // One final pass for the fitted means at the settled + // coefficients, clamped exactly as the loop clamped them: + // the unclamped exponential overflows to an infinity for + // |eta| past the log of the ceiling, and the likelihood + // below is evaluated on these values. + for r := range n { + eta := 0.0 + if fx != nil { + row := fx[r*p : r*p+p] + for j, xj := range row { + eta += xj * beta[j] + } + } else { + for j := range p { + eta += x.FloatAt(r*p+j) * beta[j] + } + } + mean[r] = poissonMean(eta) + } + break + } + } + if !converged { + return nil, base.Errf("%s: %d iterations did not converge", name, maxIter) + } + logLike := 0.0 + for r := range n { + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + logGamma, _ := math.Lgamma(yv + 1) + logLike += yv*math.Log(mean[r]) - mean[r] - logGamma + } + // Wald inference from the inverse Fisher information at the + // optimum. The iteration's last Fisher matrix belongs to the + // previous point, one damped step behind, so it is rebuilt from + // the settled means before the solves, into the reused buffer and + // on the same mirrored upper triangle. + for i := range p { + clear(fisher[i]) + } + for r := range n { + mu := mean[r] + if fx != nil { + row := fx[r*p : r*p+p] + for i, xr := range row { + fi := fisher[i][i:] + for j, xj := range row[i:] { + fi[j] += xr * xj * mu + } + } + } else { + for i := range p { + xr := x.FloatAt(r*p + i) + for j := i; j < p; j++ { + fisher[i][j] += xr * x.FloatAt(r*p+j) * mu + } + } + } + } + mirrorUpper(fisher) + out := &PoissonRegressionResult{ + Coefficients: beta, + Fitted: mean, + LogLikelihood: logLike, + Iterations: iterations, + Converged: true, + } + out.StandardErrors = make([]float64, p) + out.ZStatistics = make([]float64, p) + out.PValues = make([]float64, p) + // One unit vector per coefficient, all solved through a single + // factorisation of the Fisher information: the per-coefficient + // solves refactored the same matrix p times, while the shared solve + // substitutes each column through the identical factor. + unit := make([][]float64, p) + for j := range p { + unit[j] = make([]float64, p) + unit[j][j] = 1 + } + inv, err := base.SolveSystem(name, fisher, unit) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + for j := range p { + // The Wald variance is the diagonal of the inverse Fisher + // information. A near-collinear design can drive the solve to a + // tiny negative diagonal entry through rounding alone: the bare + // square root would then be a NaN reported beside a nil error. + // An exact zero stays a zero standard error; a negative entry + // means the design is near-collinear and is refused. + v := inv[j][j] + switch { + case v > 0: + out.StandardErrors[j] = math.Sqrt(v) + case v == 0: + out.StandardErrors[j] = 0 + default: + return nil, base.Errf("%s: the design is near-collinear: the Wald variance of coefficient %d came out negative (%g)", name, j, v) + } + if out.StandardErrors[j] == 0 { + // A zero Wald variance: an exact fit reports total evidence + // for a live coefficient and nothing to test for a zero + // one, never the 0/0 NaN pair. + if beta[j] != 0 { + out.ZStatistics[j] = math.Copysign(math.Inf(1), beta[j]) + out.PValues[j] = 0 + } else { + out.ZStatistics[j] = 0 + out.PValues[j] = 1 + } + continue + } + out.ZStatistics[j] = beta[j] / out.StandardErrors[j] + z := out.ZStatistics[j] + // The two-sided normal tail in one Erfc call on the magnitude. + // The algebraic form 2·(1−Φ(z)) cancels to exactly zero once z + // passes about 8.3, where the true tail nears 1e-17 and keeps + // another three hundred orders below before the smallest + // float64. + out.PValues[j] = math.Erfc(math.Abs(z) / math.Sqrt2) + } + return out, nil +} + +// LogisticRegression fits y = Bernoulli(sigmoid(X·β)) by maximum +// likelihood. The design carries n rows and p columns exactly as +// LinearRegression's (intercept included by the caller as a constant +// column when wanted), y holds zeros and ones, and the fit runs +// Newton-Raphson until the largest coefficient update drops under +// 1e-10. Perfectly separable data has no finite optimum: the run +// reports an error rather than diverging coefficients. +func LogisticRegression(x, y *core.Array) (*LogisticRegressionResult, error) { + const name = "LogisticRegression" + if x.NDim() != 2 { + return nil, base.Errf("%s: the design must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 { + return nil, base.Errf("%s: the response must be rank 1", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + n, p := x.Shape()[0], x.Shape()[1] + if y.Len() != n { + return nil, base.Errf("%s: the design has %d rows but the response %d", name, n, y.Len()) + } + if n <= p { + return nil, base.Errf("%s: need n > p, got %d observations and %d columns", name, n, p) + } + if p == 0 { + return nil, base.Errf("%s: the design must carry at least one column", name) + } + // A non-finite design would flow through the sigmoid and the + // Newton step into a "converged" all-NaN fit: as in + // LinearRegression, non-finite input has no answer to report. + if err := checkFinite(name, "the design", x); err != nil { + return nil, err + } + for i := range y.Len() { + v := y.FloatAt(i) + if v != 0 && v != 1 { + return nil, base.Errf("%s: response sample %d is %g, want 0 or 1", name, i, v) + } + } + beta := make([]float64, p) + prob := make([]float64, n) + // The design and the response are read through raw payload slices + // when dense: the elements are the ones FloatAt returns, so every + // product and sum below keeps its exact operand bits. + fx := rawFloats(x) + fy := rawFloats(y) + const maxIter = 100 + converged := false + iterations := maxIter + // The normal-equation buffers are allocated once and cleared per + // iteration; every accumulation pass starts from the zero state a + // fresh allocation carried. + hessian := make([][]float64, p) + for i := range p { + hessian[i] = make([]float64, p) + } + gradient := make([]float64, p) + for iter := 1; iter <= maxIter; iter++ { + for r := range n { + eta := 0.0 + if fx != nil { + row := fx[r*p : r*p+p] + for j, xj := range row { + eta += xj * beta[j] + } + } else { + for j := range p { + eta += x.FloatAt(r*p+j) * beta[j] + } + } + // The sigmoid clamped away from its saturating ends: the + // weights and the log both need a live second derivative. + pr := logisticProbability(eta) + prob[r] = pr + } + // The observed information is symmetric and each lower-triangle + // entry equals its upper twin bit for bit (mirrorUpper), so the + // accumulation runs the upper triangle alone and mirrors it once. + for i := range p { + clear(hessian[i]) + } + clear(gradient) + for r := range n { + pr := prob[r] + w := pr * (1 - pr) + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + if fx != nil { + row := fx[r*p : r*p+p] + for i, xr := range row { + gradient[i] += xr * (yv - pr) + // Upper triangle, both operands pre-sliced from i: + // the same products in the same order, bounds checks + // elided. + hi := hessian[i][i:] + for j, xj := range row[i:] { + hi[j] += xr * xj * w + } + } + } else { + for i := range p { + xr := x.FloatAt(r*p + i) + gradient[i] += xr * (yv - pr) + for j := i; j < p; j++ { + hessian[i][j] += xr * x.FloatAt(r*p+j) * w + } + } + } + } + mirrorUpper(hessian) + step, err := base.SolveSystem(name, hessian, [][]float64{gradient}) + if err != nil { + return nil, base.Errf("%s: the Fisher information is singular (%w)", name, err) + } + worst := 0.0 + for j := range p { + beta[j] += step[0][j] + if math.Abs(step[0][j]) > worst { + worst = math.Abs(step[0][j]) + } + } + if worst < 1e-10 { + converged = true + iterations = iter + // One final pass for the fitted probabilities at the + // settled coefficients, clamped exactly as the loop + // clamped them: the unclamped form reaches exactly 0 and 1 + // for |eta| > ~37, and the likelihood below is evaluated on + // these values. + for r := range n { + eta := 0.0 + if fx != nil { + row := fx[r*p : r*p+p] + for j, xj := range row { + eta += xj * beta[j] + } + } else { + for j := range p { + eta += x.FloatAt(r*p+j) * beta[j] + } + } + prob[r] = logisticProbability(eta) + } + break + } + } + if !converged { + return nil, base.Errf("%s: %d iterations did not converge; the response may be perfectly separable", name, maxIter) + } + // Wald inference from the inverse Fisher information at the + // optimum. + // + // The matrix is built row by row into the reused buffer, the way + // the fitting loop builds its own: a row streams the design once + // instead of once per coefficient, and the row's weight is formed + // once. Each entry sums its products over the rows in ascending + // order on the mirrored upper triangle, so the entries are the ones + // the direct walk accumulated. + fisher := hessian + for i := range p { + clear(fisher[i]) + } + for r := range n { + w := prob[r] * (1 - prob[r]) + if fx != nil { + row := fx[r*p : r*p+p] + for i, xi := range row { + fi := fisher[i][i:] + for j, xj := range row[i:] { + fi[j] += xi * xj * w + } + } + } else { + for i := range p { + xi := x.FloatAt(r*p + i) + fi := fisher[i] + for j := i; j < p; j++ { + fi[j] += xi * x.FloatAt(r*p+j) * w + } + } + } + } + mirrorUpper(fisher) + out := &LogisticRegressionResult{ + Coefficients: beta, + Fitted: prob, + Iterations: iterations, + Converged: true, + } + logLike := 0.0 + for r := range n { + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + logLike += yv*math.Log(prob[r]) + (1-yv)*math.Log(1-prob[r]) + } + out.LogLikelihood = logLike + out.StandardErrors = make([]float64, p) + out.ZStatistics = make([]float64, p) + out.PValues = make([]float64, p) + // One unit vector per coefficient, all solved through a single + // factorisation of the Fisher information: the per-coefficient + // solves refactored the same matrix p times, while the shared solve + // substitutes each column through the identical factor. + unit := make([][]float64, p) + for j := range p { + unit[j] = make([]float64, p) + unit[j][j] = 1 + } + inv, err := base.SolveSystem(name, fisher, unit) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + for j := range p { + // The Wald variance is the diagonal of the inverse Fisher + // information. A near-collinear design can drive the solve to a + // tiny negative diagonal entry through rounding alone: the bare + // square root would then be a NaN reported beside a nil error. + // An exact zero stays a zero standard error; a negative entry + // means the design is near-collinear and is refused. + v := inv[j][j] + switch { + case v > 0: + out.StandardErrors[j] = math.Sqrt(v) + case v == 0: + out.StandardErrors[j] = 0 + default: + return nil, base.Errf("%s: the design is near-collinear: the Wald variance of coefficient %d came out negative (%g)", name, j, v) + } + if out.StandardErrors[j] == 0 { + // A zero Wald variance: an exact fit reports total evidence + // for a live coefficient and nothing to test for a zero + // one, never the 0/0 NaN pair. + if beta[j] != 0 { + out.ZStatistics[j] = math.Copysign(math.Inf(1), beta[j]) + out.PValues[j] = 0 + } else { + out.ZStatistics[j] = 0 + out.PValues[j] = 1 + } + continue + } + out.ZStatistics[j] = beta[j] / out.StandardErrors[j] + z := out.ZStatistics[j] + // The two-sided normal tail in one Erfc call on the magnitude, + // as in PoissonRegression: 2·(1−Φ(z)) cancels to exactly zero + // once z passes about 8.3. + out.PValues[j] = math.Erfc(math.Abs(z) / math.Sqrt2) + } + return out, nil +} diff --git a/stats/glmkdemvn_test.go b/stats/glmkdemvn_test.go new file mode 100644 index 0000000..747ce75 --- /dev/null +++ b/stats/glmkdemvn_test.go @@ -0,0 +1,283 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestLogisticRegressionRecoversCoefficients fits a generated binary +// response whose truth is known: with 2000 samples the Newton fit +// must land within a few standard errors of the generating +// coefficients, the Wald statistics must match the coefficients, and +// the fitted probabilities must increase in the direction of the +// true slope. +func TestLogisticRegressionRecoversCoefficients(t *testing.T) { + g := core.NewGenerator(7) + const n = 2000 + design := core.New(core.Float, n, 2) + y := core.New(core.Float, n) + for i := range n { + xv := -2 + 4*g.Unit() + design.RawFloats()[i*2] = 1 + design.RawFloats()[i*2+1] = xv + pr := 1 / (1 + math.Exp(-(0.5 + 1.5*xv))) + bit := 0.0 + if g.Unit() < pr { + bit = 1 + } + y.RawFloats()[i] = bit + } + res, err := LogisticRegression(design, y) + if err != nil { + t.Fatalf("LogisticRegression: %v", err) + } + if !res.Converged { + t.Fatal("the fit reported no convergence") + } + if math.Abs(res.Coefficients[0]-0.5) > 4*res.StandardErrors[0] { + t.Fatalf("intercept = %.4g (%.4g SE), outside four SEs of 0.5", + res.Coefficients[0], res.StandardErrors[0]) + } + if math.Abs(res.Coefficients[1]-1.5) > 4*res.StandardErrors[1] { + t.Fatalf("slope = %.4g (%.4g SE), outside four SEs of 1.5", + res.Coefficients[1], res.StandardErrors[1]) + } + if res.ZStatistics[1] <= 3 { + t.Fatalf("slope z = %.4g, want a clearly non-zero effect", res.ZStatistics[1]) + } + last := res.Fitted[n-1] + first := res.Fitted[0] + if !(last > first) { + t.Fatalf("fitted probabilities not increasing: %g then %g", first, last) + } + // The maximised likelihood must beat the null model's. + nullLike := float64(n) * math.Log(0.5) + if res.LogLikelihood <= nullLike { + t.Fatalf("log likelihood %.4g does not beat the null %.4g", res.LogLikelihood, nullLike) + } +} + +// TestLogisticRegressionRefusals checks the response and design +// guards, including the separable-data refusal. +func TestLogisticRegressionRefusals(t *testing.T) { + design := core.New(core.Float, 4, 2) + y := core.New(core.Float, 4) + if _, err := LogisticRegression(core.New(core.Float, 4), y); err == nil { + t.Fatal("rank-1 design accepted") + } + bad := core.New(core.Float, 4) + bad.RawFloats()[2] = 0.5 + if _, err := LogisticRegression(design, bad); err == nil { + t.Fatal("non-binary response accepted") + } + // Perfect separation: y = 1 exactly when x > 0 has no finite + // optimum, and the run must say so instead of diverging. + xsep := core.New(core.Float, 8, 2) + sep := core.New(core.Float, 8) + for i := range 8 { + xsep.RawFloats()[i*2] = 1 + xsep.RawFloats()[i*2+1] = float64(i) - 3.5 + if i >= 4 { + sep.RawFloats()[i] = 1 + } + } + if _, err := LogisticRegression(xsep, sep); err == nil { + t.Fatal("separable data accepted") + } +} + +// TestMultivariateNormalDensity checks the density against the +// two-dimensional formula with a diagonal covariance, where the +// answer is a product of one-dimensional normals. +func TestMultivariateNormalDensity(t *testing.T) { + mean := smallVector(t, []float64{1, -2}) + cov, err := core.FromFloats([]float64{4, 0, 0, 9}, 2, 2) + if err != nil { + t.Fatalf("cov: %v", err) + } + x := smallVector(t, []float64{3, 1}) + got, err := MultivariateNormalLogDensity(mean, cov, x) + if err != nil { + t.Fatalf("MultivariateNormalLogDensity: %v", err) + } + dx := float64(3-1) / 2 + dy := float64(1-(-2)) / 3 + want := math.Log(1/(2*math.Pi*6)) - 0.5*dx*dx - 0.5*dy*dy + if math.Abs(got-want) > 1e-12 { + t.Fatalf("log density = %.14g, want %.14g", got, want) + } + atMean, err := MultivariateNormalLogDensity(mean, cov, mean) + if err != nil { + t.Fatalf("density at the mean: %v", err) + } + if math.Abs(atMean-math.Log(1/(2*math.Pi*6))) > 1e-12 { + t.Fatalf("density at the mean = %.14g, want the normaliser", atMean) + } + if _, err := MultivariateNormalLogDensity(mean, smallVector(t, []float64{1, 2, 3}), x); err == nil { + t.Fatal("mismatched covariance accepted") + } + // A negative pivot must be refused by name. + bad, _ := core.FromFloats([]float64{1, 0, 0, -4}, 2, 2) + if _, err := MultivariateNormalLogDensity(mean, bad, x); err == nil { + t.Fatal("indefinite covariance accepted") + } +} + +// TestMultivariateNormalDraws checks the sampler's moments: with +// 60000 draws the sample mean and covariance must sit close to the +// parameters, well inside the Monte Carlo error of the moment. +func TestMultivariateNormalDraws(t *testing.T) { + g := core.NewGenerator(11) + mean := smallVector(t, []float64{2, -1}) + cov, err := core.FromFloats([]float64{1, 0.5, 0.5, 4}, 2, 2) + if err != nil { + t.Fatalf("cov: %v", err) + } + const n = 60000 + draws, err := MultivariateNormalDraws(g, n, mean, cov) + if err != nil { + t.Fatalf("MultivariateNormalDraws: %v", err) + } + if draws.Shape()[0] != n || draws.Shape()[1] != 2 { + t.Fatalf("shape %v, want [%d 2]", draws.Shape(), n) + } + m1, m2, c, v1, v2 := 0.0, 0.0, 0.0, 0.0, 0.0 + for i := range n { + x := draws.FloatAt(i * 2) + y := draws.FloatAt(i*2 + 1) + m1 += x + m2 += y + v1 += x * x + v2 += y * y + c += x * y + } + m1 /= n + m2 /= n + if math.Abs(m1-2) > 0.03 || math.Abs(m2+1) > 0.05 { + t.Fatalf("means = %.4f, %.4f, want 2, -1", m1, m2) + } + if v1/n-m1*m1 > 1.06 || v2/n-m2*m2 > 4.25 { + t.Fatalf("variances = %.4f, %.4f, want about 1 and 4", v1/n-m1*m1, v2/n-m2*m2) + } + covHat := c/n - m1*m2 + if math.Abs(covHat-0.5) > 0.05 { + t.Fatalf("covariance = %.4f, want 0.5", covHat) + } +} + +// TestKernelDensity checks the estimate on a standard normal sample: +// it must integrate to one over a wide grid, peak near the true mode +// and stay non-negative everywhere. +func TestKernelDensity(t *testing.T) { + g := core.NewGenerator(23) + const n = 3000 + sampleVals := make([]float64, n) + for i := range n { + sampleVals[i] = g.NormalUnit() + } + sample := smallVector(t, sampleVals) + const lo, hi = -5.0, 5.0 + const grid = 400 + pointVals := make([]float64, grid) + for i := range grid { + pointVals[i] = lo + (hi-lo)*float64(i)/float64(grid-1) + } + points := smallVector(t, pointVals) + density, err := KernelDensity(sample, 0, points) + if err != nil { + t.Fatalf("KernelDensity: %v", err) + } + vals := density.RawFloats()[:density.Len()] + total := 0.0 + for i, v := range vals { + if v < 0 { + t.Fatalf("negative density at %g", pointVals[i]) + } + if i > 0 { + total += 0.5 * (v + vals[i-1]) * ((hi - lo) / (grid - 1)) + } + } + if math.Abs(total-1) > 0.01 { + t.Fatalf("the estimate integrates to %.4f, want 1", total) + } + peak, peakAt := 0.0, 0.0 + for i, v := range vals { + if v > peak { + peak, peakAt = v, pointVals[i] + } + } + if math.Abs(peakAt) > 0.25 { + t.Fatalf("the estimate peaks at %.3f, want the true mode near 0", peakAt) + } + if peak < 0.3 || peak > 0.5 { + t.Fatalf("peak height %.4f, want the normal's 0.399 within a KDE's honesty", peak) + } +} + +// TestLogisticRegressionRefusesNonFiniteDesign pins the finite-input +// gate: a NaN coefficient in the design used to flow through the +// sigmoid and the Newton step into a fit that reported convergence on +// an all-NaN result. +func TestLogisticRegressionRefusesNonFiniteDesign(t *testing.T) { + design := core.New(core.Float, 4, 2) + for i := range 8 { + design.RawFloats()[i] = float64(i%4) + float64(i/4) + } + design.RawFloats()[5] = math.NaN() + y := core.New(core.Float, 4) + y.RawFloats()[0], y.RawFloats()[1] = 0, 1 + y.RawFloats()[2], y.RawFloats()[3] = 1, 0 + if _, err := LogisticRegression(design, y); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("LogisticRegression with a NaN design: %v", err) + } +} + +// TestKernelDensityViewReadsVisibleElements pins the dense payload +// bound on both sweeps: a rebased view shares its parent's payload, +// which runs past the view's own count, so the sample and the points +// are cut to the elements a caller can see. A non-finite value in the +// invisible tail must neither poison the estimate nor refuse the call, +// and the density must be the one the standalone sample gives, bit for +// bit. +func TestKernelDensityViewReadsVisibleElements(t *testing.T) { + visible := []float64{0.5, 1.5, 2.5, 3.5} + pointVals := []float64{0, 1, 2, 3, 4} + dense := mustFromFloats(t, append(append([]float64(nil), visible...), math.NaN(), math.Inf(1)), 6) + view, err := core.Slice(dense, 0, 0, len(visible)) + if err != nil { + t.Fatalf("Slice: %v", err) + } + pointParent := mustFromFloats(t, append(append([]float64(nil), pointVals...), math.NaN(), math.Inf(-1)), len(pointVals)+2) + pointView, err := core.Slice(pointParent, 0, 0, len(pointVals)) + if err != nil { + t.Fatalf("Slice: %v", err) + } + got, err := KernelDensity(view, 0.5, pointView) + if err != nil { + t.Fatalf("KernelDensity over views with a non-finite tail: %v", err) + } + want, err := KernelDensity(mustFloats(t, visible, len(visible)), 0.5, mustFloats(t, pointVals, len(pointVals))) + if err != nil { + t.Fatalf("KernelDensity over the visible elements: %v", err) + } + if got.Len() != want.Len() { + t.Fatalf("the estimate carries %d values, want %d", got.Len(), want.Len()) + } + for i := range want.Len() { + if gb, wb := math.Float64bits(got.FloatAt(i)), math.Float64bits(want.FloatAt(i)); gb != wb { + t.Fatalf("density[%d] = %v (%#x) over the view, %v (%#x) over the visible elements", + i, got.FloatAt(i), gb, want.FloatAt(i), wb) + } + } + // The sample's own view is the whole parent's, tail included: the + // guard must still refuse it by name. + if _, err := KernelDensity(dense, 0.5, pointView); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("KernelDensity over the whole parent: %v, want the non-finite refusal", err) + } +} diff --git a/stats/gmm_posterior_reference_test.go b/stats/gmm_posterior_reference_test.go new file mode 100644 index 0000000..16466f1 --- /dev/null +++ b/stats/gmm_posterior_reference_test.go @@ -0,0 +1,176 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "math/big" + "testing" +) + +// The expectation step's posterior normalisation against an exact +// referent. The fit computes each row's responsibilities from the log +// densities and their row maximum; the quotients it publishes must +// match the posteriories a 256-bit evaluation of the same numbers +// produces, and match them at least as closely as the form that +// re-exponentiates through the row normaliser. The log densities here +// come from a closed-form bivariate evaluator, never from the mixture +// code, so the referent shares no arithmetic with the implementation. + +// bigLn2 is the natural logarithm of two to a hundred places, the +// reduction constant of the reference exponential. The head the series +// accuracy needs is far shorter than the tail quoted. +const bigLn2 = "0.6931471805599453094172321214581765680755001343602552541206800094933936219696947156058633269964186875" + +// referenceExp evaluates eˣ at prec mantissa bits: the argument folds +// to x = k·ln 2 + r with |r| ≤ ln 2, the series Σ rⁿ/n! converges on +// that range in under a hundred terms, and the shift multiplies by 2^k. +func referenceExp(x float64, prec uint) *big.Float { + p := prec + 64 + ln2, ok := new(big.Float).SetPrec(p).SetString(bigLn2) + if !ok { + panic("referenceExp: the reduction constant failed to parse") + } + xf := new(big.Float).SetPrec(p).SetFloat64(x) + k := int(math.Round(x / math.Ln2)) + r := new(big.Float).SetPrec(p).Mul(ln2, new(big.Float).SetPrec(p).SetInt64(int64(k))) + r.Sub(xf, r) + sum := new(big.Float).SetPrec(p).SetInt64(1) + term := new(big.Float).SetPrec(p).SetInt64(1) + for n := 1; n < 1000; n++ { + term.Mul(term, r) + term.Quo(term, new(big.Float).SetPrec(p).SetUint64(uint64(n))) + sum.Add(sum, term) + if e := term.MantExp(nil); term.Sign() == 0 || e < -int(p)-4 { + pow2 := new(big.Float).SetPrec(p) + pow2.SetMantExp(new(big.Float).SetPrec(p).SetInt64(1), k) + return sum.Mul(sum, pow2) + } + } + panic("referenceExp: the series failed to converge") +} + +// bivariateLogDensity evaluates the normal log density through the +// closed-form inverse and determinant of the 2×2 covariance, the +// evaluator that shares nothing with the fit's forward solve. +func bivariateLogDensity(x, y float64, mean [2]float64, cov [4]float64, weight float64) float64 { + det := cov[0]*cov[3] - cov[1]*cov[2] + dx, dy := x-mean[0], y-mean[1] + inv00, inv01, inv11 := cov[3]/det, -cov[1]/det, cov[0]/det + quad := dx*dx*inv00 + 2*dx*dy*inv01 + dy*dy*inv11 + return math.Log(weight) - 0.5*(2*math.Log(2*math.Pi)+math.Log(det)) - 0.5*quad +} + +func TestGaussianMixturePosteriorReferenceForms(t *testing.T) { + // The reference exponential agrees with the hardware one wherever + // the hardware answer is representable: this pins the reduction + // constant and the series against an independent evaluator before + // anything else is claimed on top of them. + for _, x := range []float64{-50, -7.3, -1, -0.01, 0, 0.5, 3.7, 40} { + want := math.Exp(x) + got, _ := referenceExp(x, 256).Float64() + if math.Abs(got-want) > 1e-13*math.Max(1, math.Abs(want)) { + t.Fatalf("referenceExp(%g) = %g, math.Exp gives %g", x, got, want) + } + } + + weights := []float64{0.5, 0.3, 0.2} + means := [][2]float64{{-2, 0}, {1, 1}, {3, -1}} + covs := [][4]float64{ + {1, 0.3, 0.3, 0.8}, + {0.5, 0, 0, 1.2}, + {2, -0.5, -0.5, 0.7}, + } + points := [][2]float64{ + {-2.1, 0.05}, // near the first component + {0.9, 1.2}, // near the second + {3.2, -0.8}, // near the third + {0.2, 0.4}, // between the first two + {-1, 0.8}, // a broad mixture row + {5.5, -2.5}, // moderately outside every component + {-6.5, 2.5}, // the same, opposite side + {18, 24}, // far tail: every density microscopic + } + const k = 3 + + worstNew, worstOld := 0.0, 0.0 + for _, pt := range points { + // The log densities the step would see, from the independent + // closed form. + lps := make([]float64, k) + for c := range k { + lps[c] = bivariateLogDensity(pt[0], pt[1], means[c], covs[c], weights[c]) + } + rowMax := math.Inf(-1) + for _, lp := range lps { + if lp > rowMax { + rowMax = lp + } + } + // The exact posteriories of these very numbers: exponentials + // and one quotient, all at 256 bits. + exacts := make([]*big.Float, k) + exactResp := make([]float64, k) + totalExact := new(big.Float).SetPrec(256) + for c := range lps { + exacts[c] = referenceExp(lps[c], 256) + totalExact.Add(totalExact, exacts[c]) + } + for c := range lps { + exactResp[c], _ = new(big.Float).SetPrec(256).Quo(exacts[c], totalExact).Float64() + } + // The published forms on the same numbers: the quotient of the + // stored exponentials, and the retired form that re-exponentiates + // through the row normaliser. + expSum := 0.0 + expVals := make([]float64, k) + for c := range lps { + expVals[c] = math.Exp(lps[c] - rowMax) + expSum += expVals[c] + } + logNorm := rowMax + math.Log(expSum) + // The far tail underflows every form to zero or near it, where + // a relative comparison carries no information; the row still + // participates through the sum-to-one pin. + live := false + for c := range lps { + if exactResp[c] > 1e-200 { + live = true + } + } + rowNew, rowOld := 0.0, 0.0 + sum := 0.0 + for c := range lps { + newResp := expVals[c] / expSum + oldResp := math.Exp(lps[c] - logNorm) + sum += newResp + if !live { + continue + } + exact := exactResp[c] + if exact <= 1e-200 { + continue + } + en := math.Abs(newResp-exact) / exact + eo := math.Abs(oldResp-exact) / exact + rowNew = math.Max(rowNew, en) + rowOld = math.Max(rowOld, eo) + } + worstNew = math.Max(worstNew, rowNew) + worstOld = math.Max(worstOld, rowOld) + if live && math.Abs(sum-1) > 1e-14 { + t.Fatalf("point %v: the quotients sum to %.17g, want 1", pt, sum) + } + if live && rowNew > 1e-12 { + t.Fatalf("point %v: the quotient form's worst relative error is %g, want under 1e-12", pt, rowNew) + } + if live && rowOld > 1e-12 { + t.Fatalf("point %v: the re-exponentiated form's worst relative error is %g, want under 1e-12", pt, rowOld) + } + } + if worstNew > worstOld { + t.Fatalf("the quotient form is the less accurate one at the fixture: worst relative error %g against the re-exponentiated form's %g", worstNew, worstOld) + } + t.Logf("worst relative error against the 256-bit referent: quotient form %g, re-exponentiated form %g", worstNew, worstOld) +} diff --git a/stats/gp.go b/stats/gp.go new file mode 100644 index 0000000..eb989c4 --- /dev/null +++ b/stats/gp.go @@ -0,0 +1,398 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Gaussian-process regression: a prior over functions fixed by a +// covariance kernel, conditioned exactly on the observations through +// the Cholesky factor of the Gram matrix, the same house route the +// multivariate normal takes. The marginal log likelihood is exposed as +// a plain function of the hyperparameters: the dependency graph gives +// stats no edge to optim, so the fitting loop stays with the caller, +// and the house minimiser (optim.Minimise) drives +// MarginalLogLikelihood from outside the package over the two or three +// hyperparameters a kernel carries. + +// Kernel is the covariance function of a Gaussian process: it returns +// the prior covariance k(x, y) of the process at one pair of input +// points, the points of equal dimension. Every house kernel carries a +// unit amplitude: k(x, x) = 1, and the constructors below enforce the +// parameter contract, a finite positive scale (and period where one +// belongs), so a Kernel value is always safe to evaluate. +type Kernel interface { + // Covariance returns k(x, y) for two points of equal length. + Covariance(x, y []float64) float64 +} + +// gpKernelScale validates one finite positive kernel parameter, the +// contract every constructor states. +func gpKernelScale(name, label string, v float64) error { + if !(v > 0) || math.IsInf(v, 0) { + return base.Errf("%s: %s must be finite and positive, got %g", name, label, v) + } + return nil +} + +// squaredExponential is the smooth limit kernel +// k(r) = exp(-r²/(2·lengthScale²)), infinitely differentiable, the +// default prior for functions believed smooth. +type squaredExponential struct { + lengthScale float64 +} + +// SquaredExponentialKernel returns the squared-exponential (RBF) +// kernel of unit amplitude over the Euclidean distance of the input +// points, with the given length scale. +func SquaredExponentialKernel(lengthScale float64) (Kernel, error) { + const name = "SquaredExponentialKernel" + if err := gpKernelScale(name, "the length scale", lengthScale); err != nil { + return nil, err + } + return squaredExponential{lengthScale: lengthScale}, nil +} + +// Covariance implements Kernel. +func (k squaredExponential) Covariance(x, y []float64) float64 { + r2 := sqDistance(x, y) + return math.Exp(-r2 / (2 * k.lengthScale * k.lengthScale)) +} + +// matern32 is the once differentiable Matérn with ν = 3/2, +// k(r) = (1 + √3·r/lengthScale)·exp(-√3·r/lengthScale), the spectral +// density of which is S(ω) = 4a³/(a² + ω²)² for a = √3/lengthScale: +// the closed form above and that spectral form are a Fourier pair, +// k(r) = (1/π)∫₀^∞ S(ω)·cos(ω·r) dω, the defining equation the tests +// hold the closed form against. +type matern32 struct { + lengthScale float64 +} + +// Matern32Kernel returns the Matérn kernel with ν = 3/2 and the given +// length scale, of unit amplitude over the Euclidean distance. +func Matern32Kernel(lengthScale float64) (Kernel, error) { + const name = "Matern32Kernel" + if err := gpKernelScale(name, "the length scale", lengthScale); err != nil { + return nil, err + } + return matern32{lengthScale: lengthScale}, nil +} + +// Covariance implements Kernel. +func (k matern32) Covariance(x, y []float64) float64 { + a := math.Sqrt(3) * math.Sqrt(sqDistance(x, y)) / k.lengthScale + return (1 + a) * math.Exp(-a) +} + +// matern52 is the twice differentiable Matérn with ν = 5/2, +// k(r) = (1 + √5·r/lengthScale + 5r²/(3·lengthScale²))·exp(-√5·r/lengthScale), +// the roughness the incidence of real fields usually lands between the +// two lower Matérns and the smooth exponential square. +type matern52 struct { + lengthScale float64 +} + +// Matern52Kernel returns the Matérn kernel with ν = 5/2 and the given +// length scale, of unit amplitude over the Euclidean distance. +func Matern52Kernel(lengthScale float64) (Kernel, error) { + const name = "Matern52Kernel" + if err := gpKernelScale(name, "the length scale", lengthScale); err != nil { + return nil, err + } + return matern52{lengthScale: lengthScale}, nil +} + +// Covariance implements Kernel. +func (k matern52) Covariance(x, y []float64) float64 { + r := math.Sqrt(sqDistance(x, y)) + a := math.Sqrt(5) * r / k.lengthScale + return (1 + a + a*a/3) * math.Exp(-a) +} + +// periodic is the periodic kernel +// k(r) = exp(-2·sin²(π·r/period)/lengthScale²), the prior over +// functions that repeat exactly with the given period, the length +// scale setting how sharply neighbouring periods decorrelate. +type periodic struct { + lengthScale float64 + period float64 +} + +// PeriodicKernel returns the periodic kernel with the given length +// scale and period, of unit amplitude over the Euclidean distance. +func PeriodicKernel(lengthScale, period float64) (Kernel, error) { + const name = "PeriodicKernel" + if err := gpKernelScale(name, "the length scale", lengthScale); err != nil { + return nil, err + } + if err := gpKernelScale(name, "the period", period); err != nil { + return nil, err + } + return periodic{lengthScale: lengthScale, period: period}, nil +} + +// Covariance implements Kernel. +func (k periodic) Covariance(x, y []float64) float64 { + r := math.Sqrt(sqDistance(x, y)) + s := math.Sin(math.Pi * r / k.period) + return math.Exp(-2 * s * s / (k.lengthScale * k.lengthScale)) +} + +// GaussianProcessResult carries the posterior of a Gaussian process at +// the test points. +type GaussianProcessResult struct { + // Mean is the posterior mean function evaluated at the test + // points. + Mean []float64 + // Covariance is the posterior covariance between the test + // points, m-by-m row-major in the order the test points were + // given, the full uncertainty the posterior carries. + Covariance []float64 + // Variance is the diagonal of Covariance, the marginal posterior + // variance per test point, clamped at zero: rounding can push a + // training-point diagonal a hair below zero where the truth is + // zero, and a negative variance reports nothing. + Variance []float64 + // LogLikelihood is the log marginal likelihood of the training + // observations under the prior, the objective a hyperparameter + // fit maximises (see MarginalLogLikelihood). + LogLikelihood float64 +} + +// GaussianProcessRegression conditions the prior defined by the +// kernel on the n noisy observations y over the n training rows of +// trainX (n rows, d columns), and evaluates the posterior mean and +// covariance at the m rows of testX. The noise variance sits on the +// Gram diagonal, K = k(X, X) + noiseVariance·I: zero is a legitimate +// noiseless fit, and with it the posterior interpolates the training +// data exactly and has zero variance there, while duplicated training +// rows name the singular row in an error. The posterior algebra is the +// Cholesky route: the Gram matrix factors through the multivariate +// normal machinery of mvn.go, the weights alpha solve the triangular +// systems against it, and the predictive covariance subtracts the +// forward-solved cross covariances from the prior. +// +// The inputs must be real and finite, the training and test designs +// of equal width, y of the training length, and the noise variance +// finite and non-negative. +func GaussianProcessRegression(kernel Kernel, trainX, trainY *core.Array, noiseVariance float64, testX *core.Array) (*GaussianProcessResult, error) { + const name = "GaussianProcessRegression" + l, alpha, train, n, d, err := gpFit(name, kernel, trainX, trainY, noiseVariance) + if err != nil { + return nil, err + } + test, m, testD, err := gpReadDesign(name, "the test design", testX) + if err != nil { + return nil, err + } + if testD != d { + return nil, base.Errf("%s: the training design is %d wide but the test design %d", name, d, testD) + } + out := &GaussianProcessResult{ + Mean: make([]float64, m), + Covariance: make([]float64, m*m), + Variance: make([]float64, m), + } + // The cross covariances A: A[i][j] = k(test i, train j), and the + // forward-solved V = L⁻¹Aᵀ whose rows pair off against each other + // in the predictive covariance. + cross := make([][]float64, m) + solved := make([][]float64, m) + for i := range m { + cross[i] = make([]float64, n) + for j := range n { + cross[i][j] = kernel.Covariance(test[i*d:i*d+d], train[j*d:j*d+d]) + } + solved[i] = gpForwardSolve(l, cross[i]) + } + for i := range m { + total := 0.0 + for j := range n { + total += cross[i][j] * alpha[j] + } + out.Mean[i] = total + for j := i; j < m; j++ { + // The prior term is the kernel's own diagonal entry, + // k(x*, x*), which the unit-amplitude house kernels fix at + // 1 but a custom kernel may set otherwise. + prior := kernel.Covariance(test[j*d:j*d+d], test[j*d:j*d+d]) + if j > i { + prior = kernel.Covariance(test[i*d:i*d+d], test[j*d:j*d+d]) + } + cov := prior + for r := range n { + cov -= solved[i][r] * solved[j][r] + } + out.Covariance[i*m+j] = cov + out.Covariance[j*m+i] = cov + } + diag := out.Covariance[i*m+i] + if diag < 0 { + diag = 0 + } + out.Variance[i] = diag + } + out.LogLikelihood = gpLogLikelihood(l, alpha, trainY, n) + return out, nil +} + +// MarginalLogLikelihood returns the log marginal likelihood of the +// observations y under the prior the kernel defines over the training +// design: the evidence of the hyperparameters, integrating the +// training responses against their multivariate normal prior. It is +// the objective a hyperparameter fit maximises; the dependency graph +// gives stats no edge to optim, so the house minimiser drives this +// function from outside the package, over the kernel's scale (and +// period) and the noise variance. +func MarginalLogLikelihood(kernel Kernel, trainX, trainY *core.Array, noiseVariance float64) (float64, error) { + const name = "MarginalLogLikelihood" + l, alpha, _, n, _, err := gpFit(name, kernel, trainX, trainY, noiseVariance) + if err != nil { + return 0, err + } + return gpLogLikelihood(l, alpha, trainY, n), nil +} + +// gpFit validates the training inputs and returns the Cholesky factor +// of the regularised Gram matrix, the solved weights alpha and the +// flattened training design. +func gpFit(name string, kernel Kernel, trainX, trainY *core.Array, noiseVariance float64) ([][]float64, []float64, []float64, int, int, error) { + if kernel == nil { + return nil, nil, nil, 0, 0, base.Errf("%s: the kernel is nil", name) + } + if math.IsNaN(noiseVariance) || math.IsInf(noiseVariance, 0) || noiseVariance < 0 { + return nil, nil, nil, 0, 0, base.Errf("%s: the noise variance must be finite and non-negative, got %g", name, noiseVariance) + } + train, n, d, err := gpReadDesign(name, "the training design", trainX) + if err != nil { + return nil, nil, nil, 0, 0, err + } + if err := gpReadResponse(name, trainY, n); err != nil { + return nil, nil, nil, 0, 0, err + } + // The Gram matrix, computed on the lower triangle and mirrored + // exactly, so the factorisation's symmetry check reads two + // bit-identical halves. It is written straight into one flat + // row-major slice, the layout the factorisation reads, instead of + // through a row-of-rows form another pass would flatten. + gram := make([]float64, n*n) + for i := range n { + for j := range i + 1 { + v := kernel.Covariance(train[i*d:i*d+d], train[j*d:j*d+d]) + if i == j { + v += noiseVariance + } + gram[i*n+j] = v + gram[j*n+i] = v + } + } + l, err := mvnCholeskyFlat(name, gram, n) + if err != nil { + return nil, nil, nil, 0, 0, base.Errf("%s: %w", name, err) + } + y := make([]float64, n) + fy := rawFloats(trainY) + for i := range n { + if fy != nil { + y[i] = fy[i] + } else { + y[i] = trainY.FloatAt(i) + } + } + alpha := gpBackSolve(l, gpForwardSolve(l, y)) + return l, alpha, train, n, d, nil +} + +// gpLogLikelihood assembles the log marginal likelihood from the +// factor and the solved weights: +// -½·yᵀ·alpha - Σ ln Lᵢᵢ - (n/2)·ln 2π. +func gpLogLikelihood(l [][]float64, alpha []float64, trainY *core.Array, n int) float64 { + fy := rawFloats(trainY) + total := 0.0 + for i := range n { + var yv float64 + if fy != nil { + yv = fy[i] + } else { + yv = trainY.FloatAt(i) + } + total -= 0.5 * yv * alpha[i] + total -= math.Log(l[i][i]) + } + return total - 0.5*float64(n)*math.Log(2*math.Pi) +} + +// gpReadDesign validates a finite real rank-2 design and returns it +// flattened row-major. +func gpReadDesign(name, label string, x *core.Array) ([]float64, int, int, error) { + if x.NDim() != 2 { + return nil, 0, 0, base.Errf("%s: %s must be rank 2, got shape %s", name, label, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, 0, 0, base.Errf("%s: complex inputs are not supported", name) + } + if err := checkFinite(name, label, x); err != nil { + return nil, 0, 0, err + } + n, d := x.Shape()[0], x.Shape()[1] + if d < 1 { + return nil, 0, 0, base.Errf("%s: %s needs at least one column", name, label) + } + data := make([]float64, n*d) + if fs := rawFloats(x); fs != nil { + copy(data, fs) + } else { + for i := range data { + data[i] = x.FloatAt(i) + } + } + return data, n, d, nil +} + +// gpReadResponse validates the response vector against the training +// length. +func gpReadResponse(name string, y *core.Array, n int) error { + if y.NDim() != 1 { + return base.Errf("%s: the response must be rank 1", name) + } + if y.Dtype() == core.Complex { + return base.Errf("%s: complex inputs are not supported", name) + } + if y.Len() != n { + return base.Errf("%s: the design has %d rows but the response %d", name, n, y.Len()) + } + return checkFinite(name, "the response", y) +} + +// gpForwardSolve solves L·v = b for the lower triangular L. +func gpForwardSolve(l [][]float64, b []float64) []float64 { + v := make([]float64, len(b)) + for i := range b { + total := b[i] + for j := range i { + total -= l[i][j] * v[j] + } + v[i] = total / l[i][i] + } + return v +} + +// gpBackSolve solves Lᵀ·u = b for the lower triangular L. +func gpBackSolve(l [][]float64, b []float64) []float64 { + n := len(b) + u := make([]float64, n) + for i := n - 1; i >= 0; i-- { + total := b[i] + for j := i + 1; j < n; j++ { + total -= l[j][i] * u[j] + } + u[i] = total / l[i][i] + } + return u +} diff --git a/stats/gp_test.go b/stats/gp_test.go new file mode 100644 index 0000000..11b88ad --- /dev/null +++ b/stats/gp_test.go @@ -0,0 +1,424 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestGPInterpolatesNoiselessTrainingPoints fits a noiseless GP on its +// own training points: the posterior must reproduce every observation +// to rounding and carry no variance there. +func TestGPInterpolatesNoiselessTrainingPoints(t *testing.T) { + g := core.NewGenerator(41) + const n = 9 + train := core.New(core.Float, n, 1) + y := core.New(core.Float, n) + for i := range n { + x := -2 + 0.5*float64(i) + train.RawFloats()[i] = x + y.RawFloats()[i] = math.Sin(x) + 0.1*x*g.Unit() + } + kernel, err := SquaredExponentialKernel(0.7) + if err != nil { + t.Fatalf("SquaredExponentialKernel: %v", err) + } + res, err := GaussianProcessRegression(kernel, train, y, 0, train) + if err != nil { + t.Fatalf("GaussianProcessRegression: %v", err) + } + for i := range n { + if math.Abs(res.Mean[i]-y.FloatAt(i)) > 1e-9 { + t.Fatalf("point %d: posterior mean %.12g against the observation %.12g", + i, res.Mean[i], y.FloatAt(i)) + } + if res.Variance[i] > 1e-12 { + t.Fatalf("point %d: posterior variance %.3g at a noiseless training point", i, res.Variance[i]) + } + } + // The full posterior covariance vanishes on the training set too: + // the conditioning has removed everything the prior had there. + for i := range n { + for j := range n { + if math.Abs(res.Covariance[i*n+j]) > 1e-10 { + t.Fatalf("posterior covariance (%d, %d) = %.3g at the training points, want zero", + i, j, res.Covariance[i*n+j]) + } + } + } + // The reported marginal likelihood must match the standalone + // function bit for bit on the same inputs. + mll, err := MarginalLogLikelihood(kernel, train, y, 0) + if err != nil { + t.Fatalf("MarginalLogLikelihood: %v", err) + } + if mll != res.LogLikelihood { + t.Fatalf("the fit reported %.17g but the standalone %.17g", res.LogLikelihood, mll) + } +} + +// TestGPHugeLengthScaleFollowsGlobalMean pins the limiting behaviour +// of a giant length scale: when the prior cannot tell neighbouring +// inputs apart, the posterior mean collapses onto the constant that +// the likelihood alone supports, the sample mean of the training +// responses. A small noise keeps the Gram matrix well conditioned; +// with n = 30 and noise 1e-6 the constant fit sits within 1e-7 of the +// mean, far inside the 1e-4 tolerance. +func TestGPHugeLengthScaleFollowsGlobalMean(t *testing.T) { + g := core.NewGenerator(43) + const n = 30 + train := core.New(core.Float, n, 1) + y := core.New(core.Float, n) + total := 0.0 + for i := range n { + train.RawFloats()[i] = float64(i) / float64(n-1) + y.RawFloats()[i] = 3 + 0.4*g.NormalUnit() + total += y.FloatAt(i) + } + meanY := total / n + kernel, err := SquaredExponentialKernel(1e6) + if err != nil { + t.Fatalf("SquaredExponentialKernel: %v", err) + } + test := mustFloats(t, []float64{-5, 0.37, 12}, 3, 1) + res, err := GaussianProcessRegression(kernel, train, y, 1e-6, test) + if err != nil { + t.Fatalf("GaussianProcessRegression: %v", err) + } + for i, m := range res.Mean { + if math.Abs(m-meanY) > 1e-4 { + t.Fatalf("test point %d: posterior mean %.8g against the global mean %.8g", i, m, meanY) + } + } +} + +// TestGPVarianceFarFromDataEqualsPrior pins the uncertainty +// behaviour: a test point many length scales from every observation +// has learned nothing, so its posterior variance returns to the prior +// variance of 1 and its posterior covariance with the data region +// vanishes. +func TestGPVarianceFarFromDataEqualsPrior(t *testing.T) { + const n = 6 + train := mustFloats(t, []float64{0, 0.2, 0.4, 0.6, 0.8, 1.0}, n, 1) + y := mustFloats(t, []float64{1, -1, 0.5, -0.5, 2, -2}, n) + kernel, err := SquaredExponentialKernel(0.3) + if err != nil { + t.Fatalf("SquaredExponentialKernel: %v", err) + } + // 30 is one hundred length scales from the nearest observation. + test := mustFloats(t, []float64{0.5, 30}, 2, 1) + res, err := GaussianProcessRegression(kernel, train, y, 0, test) + if err != nil { + t.Fatalf("GaussianProcessRegression: %v", err) + } + if v := res.Variance[1]; math.Abs(v-1) > 1e-10 { + t.Fatalf("far posterior variance %.15g, want the prior 1", v) + } + // The far point's covariance with the data region: the cross + // covariances are underflows, so the off-diagonal entry must be a + // rounding dust. + if c := res.Covariance[0*2+1]; math.Abs(c) > 1e-10 { + t.Fatalf("far covariance with the data region %.3g, want a rounding dust", c) + } +} + +// TestMatern32MatchesSpectralForm holds the Matérn 3/2 closed form +// against its defining spectral equation: the density +// S(ω) = 4a³/(a² + ω²)², a = √3/L, inverts to the kernel through +// k(r) = (1/π)∫₀^∞ S(ω)·cos(ω·r) dω. The integral is evaluated by +// trapezoid out to five hundred decay widths, where the integrand's +// ω⁻⁴ tail has less than 1e-7 left to give. +func TestMatern32MatchesSpectralForm(t *testing.T) { + const lengthScale = 1.3 + kernel, err := Matern32Kernel(lengthScale) + if err != nil { + t.Fatalf("Matern32Kernel: %v", err) + } + a := math.Sqrt(3) / lengthScale + const omegaMax = 500 + const steps = 400000 + h := omegaMax / float64(steps) + spectral := func(r float64) float64 { + total := 0.0 + for s := range steps + 1 { + omega := float64(s) * h + w := 1.0 + if s == 0 || s == steps { + w = 0.5 + } + den := a*a + omega*omega + total += w * 4 * a * a * a / (den * den) * math.Cos(omega*r) + } + return total * h / math.Pi + } + for _, r := range []float64{0, 0.7, lengthScale, 2.9} { + want := kernel.Covariance([]float64{r}, []float64{0}) + got := spectral(r) + if math.Abs(got-want) > 1e-6 { + t.Fatalf("r = %g: closed form %.9g against the spectral inversion %.9g", r, want, got) + } + } + // And the spot values of both Matérns: k(0) = 1 and the documented + // closed forms at one length scale, r = L, where the shape + // parameter a = √3 (respectively √5) exactly. + m52, err := Matern52Kernel(lengthScale) + if err != nil { + t.Fatalf("Matern52Kernel: %v", err) + } + if kernel.Covariance([]float64{0}, []float64{0}) != 1 { + t.Fatal("the Matern 3/2 kernel is not of unit amplitude") + } + if m52.Covariance([]float64{0}, []float64{0}) != 1 { + t.Fatal("the Matern 5/2 kernel is not of unit amplitude") + } + a32 := math.Sqrt(3) + want32 := (1 + a32) * math.Exp(-a32) + if math.Abs(kernel.Covariance([]float64{lengthScale}, []float64{0})-want32) > 1e-12 { + t.Fatalf("Matern 3/2 at r = L: %.15g, want %.15g", + kernel.Covariance([]float64{lengthScale}, []float64{0}), want32) + } + a52 := math.Sqrt(5) + want52 := (1 + a52 + a52*a52/3) * math.Exp(-a52) + if math.Abs(m52.Covariance([]float64{lengthScale}, []float64{0})-want52) > 1e-12 { + t.Fatalf("Matern 5/2 at r = L: %.15g, want %.15g", + m52.Covariance([]float64{lengthScale}, []float64{0}), want52) + } +} + +// TestPeriodicKernelRepeats pins the period: a full period apart the +// kernel returns to 1, half a period apart with a short length scale +// it has decorrelated to rounding dust. +func TestPeriodicKernelRepeats(t *testing.T) { + kernel, err := PeriodicKernel(0.2, 3) + if err != nil { + t.Fatalf("PeriodicKernel: %v", err) + } + if got := kernel.Covariance([]float64{7}, []float64{7 + 3}); math.Abs(got-1) > 1e-12 { + t.Fatalf("one period apart the covariance is %.15g, want 1", got) + } + if got := kernel.Covariance([]float64{0}, []float64{1.5}); got > 1e-6 { + t.Fatalf("half a period apart with a short scale the covariance is %.3g", got) + } +} + +// TestMarginalLogLikelihoodPrefersTrueHyperparameters draws one path +// of a Gaussian process with known hyperparameters over the house +// generator and requires the marginal likelihood to rank the true pair +// above badly mismatched ones, with a margin far beyond the rounding. +func TestMarginalLogLikelihoodPrefersTrueHyperparameters(t *testing.T) { + g := core.NewGenerator(53) + const n = 30 + const trueScale = 1.0 + const trueNoise = 0.05 + train := core.New(core.Float, n, 1) + for i := range n { + train.RawFloats()[i] = 5 * float64(i) / float64(n-1) + } + // The path: one multivariate normal draw over the grid under the + // true kernel, plus the observation noise. + trueKernel, err := SquaredExponentialKernel(trueScale) + if err != nil { + t.Fatalf("SquaredExponentialKernel: %v", err) + } + cov := core.New(core.Float, n, n) + for i := range n { + for j := range n { + cov.RawFloats()[i*n+j] = trueKernel.Covariance(train.RawFloats()[i:i+1], train.RawFloats()[j:j+1]) + } + // A hair of jitter on the diagonal: the Gram matrix of a long + // length scale is positive definite by a margin that float64 + // rounding can eat on the way down, and the draw is fixture + // construction, not the model, whose noise variance is fitted + // separately below. + cov.RawFloats()[i*n+i] += 1e-9 + } + paths, err := MultivariateNormalDraws(g, 1, core.New(core.Float, n), cov) + if err != nil { + t.Fatalf("MultivariateNormalDraws: %v", err) + } + y := core.New(core.Float, n) + for i := range n { + y.RawFloats()[i] = paths.FloatAt(i) + trueNoise*g.NormalUnit() + } + truth, err := MarginalLogLikelihood(trueKernel, train, y, trueNoise) + if err != nil { + t.Fatalf("MarginalLogLikelihood: %v", err) + } + mismatched := []struct { + scale float64 + noise float64 + }{ + {0.05, 5}, + {10, 1e-4}, + } + for _, pair := range mismatched { + wrongKernel, kerr := SquaredExponentialKernel(pair.scale) + if kerr != nil { + t.Fatalf("SquaredExponentialKernel(%g): %v", pair.scale, kerr) + } + wrong, err := MarginalLogLikelihood(wrongKernel, train, y, pair.noise) + if err != nil { + t.Fatalf("MarginalLogLikelihood at scale %g: %v", pair.scale, err) + } + if truth < wrong+5 { + t.Fatalf("the true pair scored %.4g against (%g, %g) at %.4g, the margin is under 5", + truth, pair.scale, pair.noise, wrong) + } + } +} + +// TestGPRegressionMultiDimensional fits a three-dimensional input and +// requires the same exact interpolation, the path the Euclidean +// distance of every kernel takes over columns. +func TestGPRegressionMultiDimensional(t *testing.T) { + g := core.NewGenerator(59) + const n = 8 + train := core.New(core.Float, n, 3) + y := core.New(core.Float, n) + for i := range n { + for j := range 3 { + train.RawFloats()[i*3+j] = g.Unit() + } + y.RawFloats()[i] = g.NormalUnit() + } + kernel, err := Matern52Kernel(1.5) + if err != nil { + t.Fatalf("Matern52Kernel: %v", err) + } + res, err := GaussianProcessRegression(kernel, train, y, 0, train) + if err != nil { + t.Fatalf("GaussianProcessRegression: %v", err) + } + for i := range n { + if math.Abs(res.Mean[i]-y.FloatAt(i)) > 1e-8 { + t.Fatalf("point %d: mean %.10g against the observation %.10g", i, res.Mean[i], y.FloatAt(i)) + } + if res.Variance[i] > 1e-10 { + t.Fatalf("point %d: variance %.3g at a training point", i, res.Variance[i]) + } + } + // The posterior covariance is symmetric by construction; the + // mirror must hold. + for i := range n { + for j := range n { + if res.Covariance[i*n+j] != res.Covariance[j*n+i] { + t.Fatalf("the posterior covariance is not symmetric at (%d, %d)", i, j) + } + } + } +} + +// TestGPValidationAndKernelRefusals checks the constructors and the +// entry points refuse what they must, including the singular Gram +// matrix that duplicated noiseless rows name. +func TestGPValidationAndKernelRefusals(t *testing.T) { + kernel, kerr := SquaredExponentialKernel(1) + if kerr != nil { + t.Fatalf("SquaredExponentialKernel: %v", kerr) + } + train := mustFloats(t, []float64{0, 1, 2}, 3, 1) + y := mustFloats(t, []float64{1, -1, 0.5}, 3) + test := mustFloats(t, []float64{0.5}, 1, 1) + + for _, scale := range []float64{0, -1, math.Inf(1), math.NaN()} { + if _, err := SquaredExponentialKernel(scale); err == nil { + t.Fatalf("squared exponential accepted the scale %v", scale) + } + if _, err := Matern32Kernel(scale); err == nil { + t.Fatalf("Matern 3/2 accepted the scale %v", scale) + } + if _, err := Matern52Kernel(scale); err == nil { + t.Fatalf("Matern 5/2 accepted the scale %v", scale) + } + if _, err := PeriodicKernel(scale, 1); err == nil { + t.Fatalf("periodic kernel accepted the length scale %v", scale) + } + if _, err := PeriodicKernel(1, scale); err == nil { + t.Fatalf("periodic kernel accepted the period %v", scale) + } + } + if _, err := GaussianProcessRegression(nil, train, y, 0, test); err == nil || !strings.Contains(err.Error(), "kernel is nil") { + t.Fatalf("nil kernel: got %v, want the nil-kernel refusal", err) + } + for _, noise := range []float64{-0.1, math.Inf(-1), math.NaN()} { + if _, err := GaussianProcessRegression(kernel, train, y, noise, test); err == nil || !strings.Contains(err.Error(), "noise variance") { + t.Fatalf("noise variance %v: got %v, want the noise-variance refusal", noise, err) + } + if _, err := MarginalLogLikelihood(kernel, train, y, noise); err == nil || !strings.Contains(err.Error(), "noise variance") { + t.Fatalf("marginal likelihood, noise variance %v: got %v, want the noise-variance refusal", noise, err) + } + } + if _, err := GaussianProcessRegression(kernel, core.New(core.Float, 3), y, 0, test); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("rank-1 training design: got %v, want the rank refusal", err) + } + if _, err := GaussianProcessRegression(kernel, train, core.New(core.Float, 3, 1), 0, test); err == nil || !strings.Contains(err.Error(), "must be rank 1") { + t.Fatalf("rank-2 response: got %v, want the rank refusal", err) + } + if _, err := GaussianProcessRegression(kernel, train, mustFloats(t, []float64{1, -1}, 2), 0, test); err == nil || !strings.Contains(err.Error(), "rows but the response") { + t.Fatalf("short response: got %v, want the length refusal", err) + } + if _, err := GaussianProcessRegression(kernel, train, y, 0, core.New(core.Float, 2)); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("rank-1 test design: got %v, want the rank refusal", err) + } + wide := mustFloats(t, []float64{0.1, 0.2, 0.3, 0.4}, 2, 2) + if _, err := GaussianProcessRegression(kernel, train, y, 0, wide); err == nil || !strings.Contains(err.Error(), "wide but the test design") { + t.Fatalf("test design of the wrong width: got %v, want the width refusal", err) + } + sick := core.New(core.Float, 3, 1) + sick.RawFloats()[2] = math.NaN() + if _, err := GaussianProcessRegression(kernel, sick, y, 0, test); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("non-finite training design: got %v, want the non-finite refusal", err) + } + if _, err := GaussianProcessRegression(kernel, train, mustFloats(t, []float64{1, -1, math.Inf(1)}, 3), 0, test); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("non-finite response: got %v, want the non-finite refusal", err) + } + // Duplicated rows with no noise: the Gram matrix is singular and + // the refusal names the row. + dup := mustFloats(t, []float64{1, 1, 2}, 3, 1) + _, err := GaussianProcessRegression(kernel, dup, y, 0, test) + if err == nil || !strings.Contains(err.Error(), "positive definite") { + t.Fatalf("duplicated noiseless rows: got %v, want a positive-definiteness refusal", err) + } + // A zero-column design has nothing to evaluate. + if _, err := GaussianProcessRegression(kernel, core.New(core.Float, 3, 0), y, 0, test); err == nil || !strings.Contains(err.Error(), "at least one column") { + t.Fatalf("zero-column design: got %v, want the column refusal", err) + } + if _, err := GaussianProcessRegression(kernel, core.New(core.Complex, 3, 1), y, 0, test); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("complex training design: got %v, want the complex refusal", err) + } + if _, err := GaussianProcessRegression(kernel, train, core.New(core.Complex, 3), 0, test); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("complex response: got %v, want the complex refusal", err) + } + // Integer inputs take the widening accessor path end to end: + // design, response and the likelihood assembly. + intTrain, ierr := core.FromInts([]int64{0, 10, 20}, 3, 1) + if ierr != nil { + t.Fatalf("FromInts: %v", ierr) + } + intY, ierr := core.FromInts([]int64{1, -2, 4}, 3) + if ierr != nil { + t.Fatalf("FromInts: %v", ierr) + } + intRes, err := GaussianProcessRegression(kernel, intTrain, intY, 0, intTrain) + if err != nil { + t.Fatalf("GaussianProcessRegression over int inputs: %v", err) + } + for i := range 3 { + if math.Abs(intRes.Mean[i]-intY.FloatAt(i)) > 1e-8 { + t.Fatalf("int input point %d: mean %.10g against the observation %.10g", + i, intRes.Mean[i], intY.FloatAt(i)) + } + } + intMLL, err := MarginalLogLikelihood(kernel, intTrain, intY, 0) + if err != nil { + t.Fatalf("MarginalLogLikelihood over int inputs: %v", err) + } + if intMLL != intRes.LogLikelihood { + t.Fatalf("the int-input likelihoods disagree: %.17g against %.17g", intRes.LogLikelihood, intMLL) + } +} diff --git a/stats/helpers_test.go b/stats/helpers_test.go new file mode 100644 index 0000000..7cc4c5e --- /dev/null +++ b/stats/helpers_test.go @@ -0,0 +1,64 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +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 +} + +// mustFromFloats builds a float array, failing the test on a bad shape. +func mustFromFloats(t *testing.T, vals []float64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, shape...) + if err != nil { + t.Fatalf("FromFloats(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustFromComplexes builds a complex array, failing the test on a bad shape. +func mustFromComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array { + t.Helper() + a, err := core.FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustComplexes builds a complex array, failing the test on a bad shape. +func mustComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array { + t.Helper() + a, err := core.FromComplexes(vals, shape...) + if err != nil { + t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err) + } + return a +} + +// mustFromInts builds an int array, failing the test on a bad shape. +func mustFromInts(t *testing.T, vals []int64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromInts(vals, shape...) + if err != nil { + t.Fatalf("FromInts(%v, %v): %v", vals, shape, err) + } + return a +} diff --git a/stats/hierarchy.go b/stats/hierarchy.go new file mode 100644 index 0000000..db33634 --- /dev/null +++ b/stats/hierarchy.go @@ -0,0 +1,349 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "slices" + "strconv" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Agglomerative hierarchical clustering: every observation starts in +// its own cluster and the closest pair merges repeatedly until one +// cluster holds the sample, with the rule that names "closest" +// carried by the linkage. The distances are Euclidean over the +// sample's rows, the merge distances update through the +// Lance-Williams recurrences, and ties resolve toward the lowest +// indices, so the dendrogram is a deterministic function of the +// input. The result is the merge record itself, the dendrogram, from +// which any number of flat partitions can be cut afterwards. + +// Linkage names the rule a merge's distance follows. +type Linkage int + +const ( + // SingleLinkage joins on the smallest pair distance across the two + // clusters: the nearest-neighbour rule that follows chains. + SingleLinkage Linkage = iota + // CompleteLinkage joins on the largest pair distance, the + // farthest-neighbour rule that keeps clusters tight. + CompleteLinkage + // AverageLinkage joins on the mean pair distance over all pairs + // across the two clusters, the UPGMA compromise. + AverageLinkage + // CentroidLinkage joins on the distance between the clusters' + // centroids. The rule is not monotone: a later merge can sit + // below an earlier one in the dendrogram. + CentroidLinkage + // WardLinkage joins the pair whose merge raises the within-cluster + // sum of squares the least. The heights are the square roots of + // the recorded updates, the scale the two-singleton merges share + // with the plain distances. + WardLinkage +) + +// String returns the linkage's name, the spelling the error messages +// carry. +func (l Linkage) String() string { + switch l { + case SingleLinkage: + return "single" + case CompleteLinkage: + return "complete" + case AverageLinkage: + return "average" + case CentroidLinkage: + return "centroid" + case WardLinkage: + return "ward" + } + return "linkage(" + strconv.Itoa(int(l)) + ")" +} + +// HierarchicalMaxObservations is the sample cap of +// HierarchicalClustering: the working distance matrix is quadratic in +// the sample, and past the cap the refusal names the cost rather than +// handing gigabytes to the allocator. A larger sample wants an +// approximate construction the library does not pretend to carry. +const HierarchicalMaxObservations = 4096 + +// Dendrogram records an agglomerative clustering as its sequence of +// merges. Leaf i is the sample's row i; the merge at step t creates +// cluster n+t from the clusters Left[t] and Right[t], at +// Heights[t], holding Sizes[t] rows. Left always holds the smaller +// cluster id. +type Dendrogram struct { + // Left and Right name the two clusters each merge joins. + Left []int + Right []int + // Heights are the merge distances, in merge order. For the + // centroid and Ward linkages an inversion is possible and + // legitimate: the sequence need not be non-decreasing. + Heights []float64 + // Sizes counts the rows each merge's cluster carries. + Sizes []int +} + +// Cut partitions the sample into k flat clusters by undoing the last +// k−1 merges, and labels every row 0 to k−1. The labels are ordered +// by each cluster's smallest row, so the first row of the sample +// always lands in cluster 0. Refuses a nil dendrogram and a k outside +// [1, n]. +func (d *Dendrogram) Cut(k int) ([]int, error) { + const name = "Dendrogram.Cut" + if d == nil { + return nil, base.Errf("%s: the dendrogram is nil", name) + } + n := len(d.Sizes) + 1 + if k < 1 || k > n { + return nil, base.Errf("%s: k must lie in [1, %d], got %d", name, n, k) + } + return d.cutMerges(n - k), nil +} + +// CutHeight partitions the sample by applying merges in order while +// they sit at or below the given height, and labels every row 0 +// upward by each cluster's smallest row, as Cut does. A height below +// the first merge returns the singletons; one at or above the last +// returns the whole sample as one cluster. Refuses a nil dendrogram +// and a non-finite height. +func (d *Dendrogram) CutHeight(height float64) ([]int, error) { + const name = "Dendrogram.CutHeight" + if d == nil { + return nil, base.Errf("%s: the dendrogram is nil", name) + } + if math.IsNaN(height) || math.IsInf(height, 0) { + return nil, base.Errf("%s: the height must be finite, got %g", name, height) + } + m := 0 + for m < len(d.Heights) && d.Heights[m] <= height { + m++ + } + return d.cutMerges(m), nil +} + +// cutMerges applies the first m merges and labels the rows by +// cluster, in order of each cluster's smallest row. +func (d *Dendrogram) cutMerges(m int) []int { + n := len(d.Sizes) + 1 + // The merge tree as a parent table: the merge at step t lifts both + // children under the cluster n+t. Every id exceeds its parents, so + // finding a row's root is a plain upward walk. + parent := make([]int, 2*n-1) + for i := range parent { + parent[i] = i + } + // The smallest row each applied merge's cluster holds. + smallest := make([]int, 2*n-1) + for i := range n { + smallest[i] = i + } + for t := range m { + a, b := d.Left[t], d.Right[t] + parent[a] = n + t + parent[b] = n + t + c := n + t + smallest[c] = min(smallest[a], smallest[b]) + } + root := make([]int, n) + for i := range n { + c := i + for parent[c] != c { + c = parent[c] + } + root[i] = c + } + // The distinct smallest-rows, sorted, are the label indices. + reps := make([]int, 0, n) + seen := make(map[int]bool) + for i := range n { + r := smallest[root[i]] + if !seen[r] { + seen[r] = true + reps = append(reps, r) + } + } + slices.Sort(reps) + labels := make([]int, n) + for i := range n { + labels[i] = slices.Index(reps, smallest[root[i]]) + } + return labels +} + +// HierarchicalClustering builds the dendrogram of the sample's rows +// under the given linkage, over Euclidean distances. The algorithm is +// the naive agglomerative sweep: each step scans the working distance +// matrix for the closest pair, merges it, and rewrites the merged +// cluster's distances through the linkage's Lance-Williams update, so +// the run costs O(n³) time and O(n²) memory and is deterministic for +// the input, ties included. +// +// Refuses a sample that is not rank 2, complex input, a non-finite +// entry, fewer than two rows, more than HierarchicalMaxObservations +// rows, and an unknown linkage. +func HierarchicalClustering(x *core.Array, method Linkage) (*Dendrogram, error) { + const name = "HierarchicalClustering" + switch method { + case SingleLinkage, CompleteLinkage, AverageLinkage, CentroidLinkage, WardLinkage: + default: + return nil, base.Errf("%s: unknown linkage %s", name, method) + } + if x.NDim() != 2 { + return nil, base.Errf("%s: the sample must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, base.Errf("%s: complex samples are not supported", name) + } + if err := checkFinite(name, "the sample", x); err != nil { + return nil, err + } + n := x.Shape()[0] + d := x.Shape()[1] + if n < 2 { + return nil, base.Errf("%s: at least two observations are needed, got %d", name, n) + } + if n > HierarchicalMaxObservations { + return nil, base.Errf("%s: %d observations exceed the %d cap; the working matrix alone would hold %.1f GiB", + name, n, HierarchicalMaxObservations, float64(n)*float64(n)*8/(1<<30)) + } + // The sample is read once and the squared Euclidean distances are + // built from the flat rows: centroid and Ward update squared + // distances, and the plain linkages take their square root once at + // the start. + data := make([]float64, n*d) + if fs := rawFloats(x); fs != nil { + copy(data, fs[:n*d]) + } else { + for i := range data { + data[i] = x.FloatAt(i) + } + } + squared := method == CentroidLinkage || method == WardLinkage + dist := make([]float64, n*n) + for i := range n { + for j := range i { + total := 0.0 + for c := range d { + diff := data[i*d+c] - data[j*d+c] + total += diff * diff + } + if squared { + dist[i*n+j] = total + dist[j*n+i] = total + } else { + plain := math.Sqrt(total) + dist[i*n+j] = plain + dist[j*n+i] = plain + } + } + } + // The active clusters live in positions 0..k−1 of the working + // matrix; actID names each position's cluster id and actSize its + // row count. A merge rewrites the matrix in place: the merged + // cluster takes position i, the last position's rows move into j's, + // and the working width drops by one. + actID := make([]int, n) + actSize := make([]int, n) + for i := range n { + actID[i] = i + actSize[i] = 1 + } + dendrogram := &Dendrogram{ + Left: make([]int, n-1), + Right: make([]int, n-1), + Heights: make([]float64, n-1), + Sizes: make([]int, n-1), + } + k := n + for step := range n - 1 { + // The closest pair, ties toward the lowest positions. The + // matrix keeps its full stride n for its whole life: the + // active clusters occupy positions 0..k−1 and only the loop + // bounds shrink. + bestI, bestJ := 0, 1 + bestD := math.Inf(1) + for i := range k { + for j := i + 1; j < k; j++ { + if dist[i*n+j] < bestD { + bestD = dist[i*n+j] + bestI, bestJ = i, j + } + } + } + ni, nj := actSize[bestI], actSize[bestJ] + idI, idJ := actID[bestI], actID[bestJ] + height := bestD + if squared { + // The centroid and Ward updates carry squared quantities; + // the recorded heights take their root. A rounding descent + // below zero is clamped, the square root refusing it + // otherwise. + height = math.Sqrt(math.Max(0, bestD)) + } + dendrogram.Left[step] = min(idI, idJ) + dendrogram.Right[step] = max(idI, idJ) + dendrogram.Heights[step] = height + dendrogram.Sizes[step] = ni + nj + // The merged distances to every surviving cluster, then the + // compaction: the last position's cluster moves into bestJ's + // slot before the working width drops by one. + last := k - 1 + for q := range k { + if q == bestI || q == bestJ { + continue + } + dik := dist[bestI*n+q] + djk := dist[bestJ*n+q] + var updated float64 + switch method { + case SingleLinkage: + updated = min(dik, djk) + case CompleteLinkage: + updated = max(dik, djk) + case AverageLinkage: + updated = (float64(ni)*dik + float64(nj)*djk) / float64(ni+nj) + case CentroidLinkage: + total := float64(ni + nj) + updated = (float64(ni)*dik+float64(nj)*djk)/total - + float64(ni)*float64(nj)*bestD/(total*total) + case WardLinkage: + total := float64(ni + nj + actSize[q]) + updated = ((float64(ni+actSize[q]))*dik + (float64(nj+actSize[q]))*djk - + float64(actSize[q])*bestD) / total + } + if squared { + updated = math.Max(0, updated) + } + dist[bestI*n+q] = updated + dist[q*n+bestI] = updated + } + if bestJ != last { + // The last position's cluster moves into bestJ's slot: its + // row and column transfer whole, the merged cluster's own + // entries against it included, and only the diagonals are + // left alone, zero on both sides. + for q := range k { + switch { + case q == bestJ: + case q == bestI: + dist[bestI*n+bestJ] = dist[bestI*n+last] + dist[bestJ*n+bestI] = dist[bestI*n+last] + default: + dist[bestJ*n+q] = dist[last*n+q] + dist[q*n+bestJ] = dist[q*n+last] + } + } + actID[bestJ] = actID[last] + actSize[bestJ] = actSize[last] + } + actID[bestI] = n + step + actSize[bestI] = ni + nj + k = last + } + return dendrogram, nil +} diff --git a/stats/hierarchy_test.go b/stats/hierarchy_test.go new file mode 100644 index 0000000..da473ea --- /dev/null +++ b/stats/hierarchy_test.go @@ -0,0 +1,242 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The hierarchical clustering against hand-worked referents. The +// examples are small enough to read the merges off the distances, and +// the 1-D samples make every height checkable by arithmetic. + +// hierSample builds an (n × 1) sample from one row per value. +func hierSample(t *testing.T, vals []float64) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, len(vals), 1) + if err != nil { + t.Fatal(err) + } + return a +} + +func hierCluster(t *testing.T, vals []float64, method Linkage) *Dendrogram { + t.Helper() + d, err := HierarchicalClustering(hierSample(t, vals), method) + if err != nil { + t.Fatal(err) + } + return d +} + +func TestHierarchicalClusteringOneDimensionalMerges(t *testing.T) { + // Points 0, 1, 10, 12 on a line: the close pairs merge at 1 and + // 2, and the final merge's height names the linkage. + sample := []float64{0, 1, 10, 12} + cases := []struct { + method Linkage + heights []float64 + }{ + {SingleLinkage, []float64{1, 2, 9}}, // nearest pair across the halves + {CompleteLinkage, []float64{1, 2, 12}}, // farthest pair: 12 down to 0 + {AverageLinkage, []float64{1, 2, 10.5}}, // the mean of the four cross pairs + {CentroidLinkage, []float64{1, 2, 10.5}}, // centroids 0.5 and 11 are 10.5 apart + // Ward's updates carry twice the within-cluster sum of squares' + // increase: ESS of the four points is 112.75 against 2.5 in the + // two pairs, so the height is √(2·110.25). + {WardLinkage, []float64{1, 2, math.Sqrt(220.5)}}, + } + for _, c := range cases { + d := hierCluster(t, sample, c.method) + for step, h := range c.heights { + if math.Abs(d.Heights[step]-h) > 1e-12 { + t.Fatalf("%s merge %d height = %.16f, want %.16f", c.method, step, d.Heights[step], h) + } + } + if d.Sizes[2] != 4 || d.Sizes[1] != 2 || d.Sizes[0] != 2 { + t.Fatalf("%s sizes %v, want the close pairs first", c.method, d.Sizes) + } + } + // The close pairs merge first under every linkage, and Left holds + // the smaller id. + d := hierCluster(t, sample, SingleLinkage) + if d.Left[0] != 0 || d.Right[0] != 1 { + t.Fatalf("the first merge joined %d and %d, want rows 0 and 1", d.Left[0], d.Right[0]) + } +} + +func TestHierarchicalClusteringCut(t *testing.T) { + sample := []float64{0, 1, 10, 12, 40} + d := hierCluster(t, sample, SingleLinkage) + // k = 2 splits the 40 outlier from the rest. + labels, err := d.Cut(2) + if err != nil { + t.Fatal(err) + } + want := []int{0, 0, 0, 0, 1} + for i := range labels { + if labels[i] != want[i] { + t.Fatalf("Cut(2) labels %v, want %v", labels, want) + } + } + // k = 3 splits the two close pairs apart. + labels, err = d.Cut(3) + if err != nil { + t.Fatal(err) + } + want = []int{0, 0, 1, 1, 2} + for i := range labels { + if labels[i] != want[i] { + t.Fatalf("Cut(3) labels %v, want %v", labels, want) + } + } + // The extremes: one cluster, and the singletons in row order. + labels, err = d.Cut(1) + if err != nil { + t.Fatal(err) + } + for i := range labels { + if labels[i] != 0 { + t.Fatalf("Cut(1) label %d is %d, want 0", i, labels[i]) + } + } + labels, err = d.Cut(5) + if err != nil { + t.Fatal(err) + } + for i := range labels { + if labels[i] != i { + t.Fatalf("Cut(5) label %d is %d, want %d", i, labels[i], i) + } + } + // The height cut at 1.5 keeps the pair {0, 1} and leaves the rest + // singletons: four clusters, labels by smallest row. + labels, err = d.CutHeight(1.5) + if err != nil { + t.Fatal(err) + } + want = []int{0, 0, 1, 2, 3} + for i := range labels { + if labels[i] != want[i] { + t.Fatalf("CutHeight(1.5) labels %v, want %v", labels, want) + } + } + // A height past the last merge is the one-cluster answer. + labels, err = d.CutHeight(100) + if err != nil { + t.Fatal(err) + } + for i := range labels { + if labels[i] != 0 { + t.Fatalf("CutHeight(100) label %d is %d, want 0", i, labels[i]) + } + } +} + +func TestHierarchicalClusteringWardRecoversBlobs(t *testing.T) { + // Two tight clusters far apart: the Ward cut at two must match + // the blobs, and the heights must be non-decreasing (Ward is a + // monotone linkage). + blob := func(centre, spread float64, count int, offset int) []float64 { + out := make([]float64, count) + for i := range count { + out[i] = centre + spread*math.Sin(float64(i+offset)) + } + return out + } + vals := append(blob(0, 0.1, 8, 0), blob(50, 0.1, 8, 3)...) + d := hierCluster(t, vals, WardLinkage) + for step := 1; step < len(d.Heights); step++ { + if d.Heights[step] < d.Heights[step-1] { + t.Fatalf("Ward heights descend at %d: %g after %g", step, d.Heights[step], d.Heights[step-1]) + } + } + labels, err := d.Cut(2) + if err != nil { + t.Fatal(err) + } + for i := range labels { + want := 0 + if i >= 8 { + want = 1 + } + if labels[i] != want { + t.Fatalf("row %d landed in cluster %d, want %d", i, labels[i], want) + } + } + // Single and complete linkage recover the same partition, and + // their heights are monotone too. + for _, method := range []Linkage{SingleLinkage, CompleteLinkage, AverageLinkage} { + d := hierCluster(t, vals, method) + for step := 1; step < len(d.Heights); step++ { + if d.Heights[step] < d.Heights[step-1] { + t.Fatalf("%s heights descend at %d", method, step) + } + } + labels, err := d.Cut(2) + if err != nil { + t.Fatal(err) + } + if labels[0] != labels[7] || labels[8] != labels[15] || labels[0] == labels[8] { + t.Fatalf("%s split the blobs apart wrongly: %v", method, labels) + } + } +} + +func TestHierarchicalClusteringDeterministic(t *testing.T) { + vals := make([]float64, 16) + for i := range vals { + vals[i] = float64((i*37)%23) / 3 + } + run := func() *Dendrogram { + return hierCluster(t, vals, AverageLinkage) + } + a, b := run(), run() + for step := range a.Heights { + if a.Heights[step] != b.Heights[step] || a.Left[step] != b.Left[step] || + a.Right[step] != b.Right[step] || a.Sizes[step] != b.Sizes[step] { + t.Fatalf("two identical runs disagreed at merge %d", step) + } + } +} + +func TestHierarchicalClusteringRefusals(t *testing.T) { + if _, err := HierarchicalClustering(hierSample(t, []float64{1}), WardLinkage); err == nil { + t.Fatal("a one-row sample was accepted") + } + if _, err := HierarchicalClustering(mustStatArray(t, []float64{1, 2, 3, 4}, 4), WardLinkage); err == nil { + t.Fatal("a rank-1 sample was accepted") + } + if _, err := HierarchicalClustering(hierSample(t, []float64{1, math.NaN()}), WardLinkage); err == nil { + t.Fatal("a non-finite sample was accepted") + } + if _, err := HierarchicalClustering(hierSample(t, []float64{1, 2}), Linkage(9)); err == nil { + t.Fatal("an unknown linkage was accepted") + } + // One row over the cap, refused before the quadratic matrix. + big := make([]float64, HierarchicalMaxObservations+1) + for i := range big { + big[i] = float64(i) + } + if _, err := HierarchicalClustering(hierSample(t, big), WardLinkage); err == nil { + t.Fatal("an over-cap sample was accepted") + } + d := hierCluster(t, []float64{0, 1, 10}, SingleLinkage) + if _, err := d.Cut(0); err == nil { + t.Fatal("Cut(0) was accepted") + } + if _, err := d.Cut(4); err == nil { + t.Fatal("Cut(4) on a three-row sample was accepted") + } + if _, err := d.CutHeight(math.NaN()); err == nil { + t.Fatal("a NaN height was accepted") + } + var nilD *Dendrogram + if _, err := nilD.Cut(1); err == nil { + t.Fatal("a nil dendrogram was accepted") + } +} diff --git a/stats/histogram2d.go b/stats/histogram2d.go new file mode 100644 index 0000000..6e820bf --- /dev/null +++ b/stats/histogram2d.go @@ -0,0 +1,214 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Histogram2D bins the paired samples (x[i], y[i]) into a +// xBins × yBins count matrix over equal-width bins spanning each +// axis's own data range, mirroring Histogram's conventions: bins are +// closed on the left and the last bin absorbs the maximum. The count +// matrix comes back as a (xBins × yBins) array, the edges as the +// bin boundaries on each axis. Real inputs only, and the bin counts +// are bounded by maxHistBins in total, exactly as Histogram's are. +func Histogram2D(x, y *core.Array, xBins, yBins int) (*core.Array, []float64, []float64, error) { + const name = "Histogram2D" + if x.Len() != y.Len() { + return nil, nil, nil, base.Errf("%s: x and y must share their length, got %d and %d", name, x.Len(), y.Len()) + } + if x.Len() == 0 { + return nil, nil, nil, base.Errf("%s: empty samples have no histogram", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, nil, nil, base.Errf("%s: complex inputs are not supported", name) + } + if xBins < 1 || yBins < 1 { + return nil, nil, nil, base.Errf("%s: needs at least one bin per axis, got %d and %d", name, xBins, yBins) + } + // The feasibility check runs before any allocation: the count matrix + // costs xBins*yBins cells, and refused up front it can neither + // overflow the product nor ask the allocator for an unbounded + // buffer. + if xBins > maxHistBins || yBins > maxHistBins || xBins > maxHistBins/yBins { + return nil, nil, nil, base.Errf("%s: the %d x %d bin request exceeds the %d-bin limit", + name, xBins, yBins, maxHistBins) + } + xEdges, err := edgesFor(x, xBins, name, "x") + if err != nil { + return nil, nil, nil, err + } + yEdges, err := edgesFor(y, yBins, name, "y") + if err != nil { + return nil, nil, nil, err + } + counts := core.New(core.Int, xBins, yBins) + ints := counts.RawInts() + // The paired samples are read once: a dense float64 payload is walked + // in place and any other layout is widened into a private slice, so + // the sweep below reads no accessor and every value is the one + // FloatAt returns. + xVals, yVals := histogramAxis(x), histogramAxis(y) + // Counting a cell is exact integer arithmetic and every sample + // carries exactly one increment, so the sweep splits over disjoint + // slices of the sample and the private counters merge in any order + // into the totals a single pass produces. + bins := xBins * yBins + minPerWorker := max(hist2DSerialSamples, hist2DBinsPerWorker*bins) + // The first non-finite pair is the one the serial sweep would name: + // the chunks report their own first, and the lowest index wins. + bad := -1 + var mu sync.Mutex + note := func(i int) { + if i < 0 { + return + } + mu.Lock() + if bad < 0 || i < bad { + bad = i + } + mu.Unlock() + } + engine.ParallelMin(x.Len(), minPerWorker, func(start, end int) { + if start == 0 && end == x.Len() { + // The whole sample runs inline, the worker policy having + // found the split uneconomical: count straight into the + // result. + note(histogram2DCount(xVals, yVals, 0, xEdges, yEdges, xBins, yBins, ints)) + return + } + local := make([]int64, bins) + first := histogram2DCount(xVals[start:end], yVals[start:end], start, xEdges, yEdges, xBins, yBins, local) + mu.Lock() + for i, c := range local { + ints[i] += c + } + mu.Unlock() + note(first) + }) + if bad >= 0 { + return nil, nil, nil, base.Errf("%s: sample %d (%g, %g) is not finite", name, bad, xVals[bad], yVals[bad]) + } + return counts, xEdges, yEdges, nil +} + +// hist2DBinsPerWorker is the paired-sample count one counting worker +// must carry per cell before the binning splits across goroutines, and +// hist2DSerialSamples the count it must carry whatever the cell count: +// a worker counts into a private matrix of xBins·yBins cells, so the +// split only pays where the samples drained into it outweigh the cells +// it costs. +const ( + hist2DSerialSamples = 1 << 12 + hist2DBinsPerWorker = 8 +) + +// histogramAxis returns the axis's values as a plain slice, walking a +// dense float64 payload in place. Element i sits at payload index i, +// and a rebased view's payload may run past its own count, so the view +// is cut to the visible elements. +func histogramAxis(a *core.Array) []float64 { + if fs := rawFloats(a); fs != nil { + return fs[:a.Len()] + } + vals := make([]float64, a.Len()) + for i := range vals { + vals[i] = a.FloatAt(i) + } + return vals +} + +// histogram2DCount adds one slice of the paired sample to counts and +// returns the index of the first non-finite pair it met, or -1. The bin +// arithmetic is binOf's own, so a split of the sample changes no count. +func histogram2DCount(xv, yv []float64, base int, xEdges, yEdges []float64, xBins, yBins int, counts []int64) int { + xBase, xStep := xEdges[0], xEdges[1]-xEdges[0] + yBase, yStep := yEdges[0], yEdges[1]-yEdges[0] + for i := range xv { + xi := binOf(xv[i], xBase, xStep, xBins) + yi := binOf(yv[i], yBase, yStep, yBins) + if xi < 0 || yi < 0 { + return base + i + } + counts[xi*yBins+yi]++ + } + return -1 +} + +// edgesFor builds one axis's bin edges over the data's own range. +func edgesFor(a *core.Array, bins int, name, axis string) ([]float64, error) { + lo, hi := math.Inf(1), math.Inf(-1) + // The finiteness scan and the range share one pass over a dense + // float64 payload, the values being the ones FloatAt returns; every + // other layout keeps the accessor scan and the range pass after it. + // The payload walk is bounded by Len: a rebased view's payload may + // run past its own count. + if fs := rawFloats(a); fs != nil { + for i, v := range fs[:a.Len()] { + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: %s sample %d is not finite (%g)", name, axis, i, v) + } + if v < lo { + lo = v + } + if v > hi { + hi = v + } + } + } else { + for i := range a.Len() { + if v := a.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: %s sample %d is not finite (%g)", name, axis, i, v) + } + } + for i := range a.Len() { + v := a.FloatAt(i) + if v < lo { + lo = v + } + if v > hi { + hi = v + } + } + } + if lo == hi { + lo -= 0.5 + hi += 0.5 + } + width := (hi - lo) / float64(bins) + if math.IsInf(width, 0) || math.IsNaN(width) { + // Both float extremes in one axis span more than the float64 + // range: no finite edges exist, and the bin arithmetic would + // clamp every sample into bin 0 over ±Inf edges. + return nil, base.Errf("%s: the %s samples span more than the float64 range (%g to %g)", name, axis, lo, hi) + } + edges := make([]float64, bins+1) + for i := range bins + 1 { + edges[i] = lo + float64(i)*width + } + return edges, nil +} + +// binOf locates one sample's bin from the axis's base edge and the step +// between its first two edges, closing the last bin on the right; a +// non-finite sample returns −1. +func binOf(v, base, step float64, bins int) int { + if math.IsNaN(v) || math.IsInf(v, 0) { + return -1 + } + i := int((v - base) / step) + if i < 0 { + return 0 // cannot happen over the data's own range, kept guarded + } + if i >= bins { + i = bins - 1 + } + return i +} diff --git a/stats/hmm.go b/stats/hmm.go new file mode 100644 index 0000000..f0cfb39 --- /dev/null +++ b/stats/hmm.go @@ -0,0 +1,565 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Hidden Markov models over discrete observations: the filtered and +// smoothed state posteriors, the most likely state path and the +// Baum-Welch fit of the parameters. Every computation runs on the +// scaled recursions, where each sweep's normaliser absorbs the +// likelihood of what has been seen, so a long sequence cannot +// underflow the way the raw forward probabilities would; the log +// likelihood itself is the sum of the normalisers' logarithms. The +// decoding walks in log space with the lowest index winning ties, and +// the fit draws its starting point from the house generator, so the +// whole surface is deterministic for a given generator state and +// sequence. + +// HiddenMarkovModel is a hidden Markov model over a set of hidden +// states and a set of observation symbols. The fields hold the +// distribution parameters in row-major order: the rows of Transition +// are the conditionals P(next = j | current = k) and the rows of +// Emission the conditionals P(symbol = m | state = k). The +// constructors validate every row as a distribution, so a model value +// is always safe to evaluate. +type HiddenMarkovModel struct { + // Initial is the distribution over states at the first step. + Initial []float64 + // Transition is the states×states matrix of one-step + // probabilities, row-major. + Transition []float64 + // Emission is the states×symbols matrix of observation + // probabilities, row-major. + Emission []float64 +} + +// HiddenMarkovFitResult carries the Baum-Welch fit of a hidden Markov +// model. +type HiddenMarkovFitResult struct { + // Model is the fitted model. + Model *HiddenMarkovModel + // LogLikelihood is the log likelihood of the training sequence + // under the fitted model. + LogLikelihood float64 + // Iterations counts the Baum-Welch sweeps taken; Converged reports + // whether the log likelihood settled under the tolerance before + // the budget ran out. + Iterations int + Converged bool +} + +// hmmMaxIterations caps the Baum-Welch sweeps, with the same +// reporting convention the other fits carry. +const hmmMaxIterations = 500 + +// hmmTolerance is the convergence tolerance on the log likelihood: +// the sweeps stop once it moves by less than hmmTolerance scaled by +// 1 + |log likelihood|. Baum-Welch's ascent crawls along a flat ridge +// long before its parameters settle to the last digit, and the +// tolerance stops the sweep where the objective's remaining climb is +// immaterial rather than where the parameters stop moving. +const hmmTolerance = 1e-6 + +// hmmFloor is the probability floor the re-estimation applies before +// renormalising: a parameter driven to exactly zero would silence a +// state or a symbol forever, and the floor keeps every path alive at +// a width no plausible parameter sits under. +const hmmFloor = 1e-12 + +// NewHiddenMarkovModel validates and copies the parameters of a +// hidden Markov model: Initial of the length of the state set, +// Transition with one row per state over the states, Emission with +// one row per state over the symbols. Every entry must be finite and +// in [0, 1] and every row must sum to 1 within 1e-9; a model that +// failed any of these would silently weight impossible paths later. +func NewHiddenMarkovModel(initial, transition, emission []float64) (*HiddenMarkovModel, error) { + const name = "NewHiddenMarkovModel" + states := len(initial) + if states < 1 { + return nil, base.Errf("%s: the model needs at least one state", name) + } + if len(transition) != states*states { + return nil, base.Errf("%s: the transition matrix holds %d entries, want %d for %d states", + name, len(transition), states*states, states) + } + if len(emission) == 0 || len(emission)%states != 0 { + return nil, base.Errf("%s: the emission matrix holds %d entries, not a whole number of rows for %d states", + name, len(emission), states) + } + if err := hmmCheckRow(name, "the initial distribution", initial); err != nil { + return nil, err + } + for k := range states { + if err := hmmCheckRow(name, "the transition rows", transition[k*states:(k+1)*states]); err != nil { + return nil, err + } + } + symbols := len(emission) / states + for k := range states { + if err := hmmCheckRow(name, "the emission rows", emission[k*symbols:(k+1)*symbols]); err != nil { + return nil, err + } + } + return &HiddenMarkovModel{ + Initial: append([]float64(nil), initial...), + Transition: append([]float64(nil), transition...), + Emission: append([]float64(nil), emission...), + }, nil +} + +// hmmCheckRow validates one distribution row: finite, inside [0, 1], +// summing to 1 within 1e-9. +func hmmCheckRow(name, what string, row []float64) error { + total := 0.0 + for _, v := range row { + if math.IsNaN(v) || math.IsInf(v, 0) { + return base.Errf("%s: %s hold the non-finite value %g", name, what, v) + } + if v < 0 || v > 1 { + return base.Errf("%s: %s hold the probability %g outside [0, 1]", name, what, v) + } + total += v + } + if math.Abs(total-1) > 1e-9 { + return base.Errf("%s: %s sum to %g, not 1", name, what, total) + } + return nil +} + +// hmmStates returns the model's state count. +func (m *HiddenMarkovModel) hmmStates() int { return len(m.Initial) } + +// hmmSymbols returns the model's symbol count. +func (m *HiddenMarkovModel) hmmSymbols() int { return len(m.Emission) / max(1, len(m.Initial)) } + +// hmmWidth returns the model's symbol count from its own fields, the +// shape the recursions read the emission rows with. +func (m *HiddenMarkovModel) hmmWidth() int { return len(m.Emission) / len(m.Initial) } + +// hmmCheckSequence validates an observation sequence against the +// model's symbol set. +func hmmCheckSequence(name string, observations []int, symbols int) error { + if len(observations) == 0 { + return base.Errf("%s: the observation sequence is empty", name) + } + for t, o := range observations { + if o < 0 || o >= symbols { + return base.Errf("%s: observation %d is symbol %d, outside the model's %d symbols", name, t, o, symbols) + } + } + return nil +} + +// hmmForward runs the scaled forward recursion: alpha[t] is the +// P(state, observations 0..t) vector normalised by its own total, and +// the totals accumulate into the log likelihood. The returned rows +// are fresh. A normaliser of exactly zero means the sequence has +// probability zero under the model, the structural zeros the +// parameters carry having silenced every state at some step: the +// posteriors are undefined there and the recursion refuses rather +// than divide the zero into NaNs. +func (m *HiddenMarkovModel) hmmForward(observations []int) ([][]float64, []float64, error) { + const name = "HiddenMarkovModel.Forward" + states := m.hmmStates() + width := m.hmmWidth() + alpha := make([][]float64, len(observations)) + scales := make([]float64, len(observations)) + row := make([]float64, states) + for k := range states { + row[k] = m.Initial[k] * m.Emission[k*width+observations[0]] + } + total := 0.0 + for _, v := range row { + total += v + } + if total == 0 { + return nil, nil, base.Errf("%s: the observation sequence has probability zero under the model", name) + } + scales[0] = total + for k := range states { + row[k] /= total + } + alpha[0] = row + for t := 1; t < len(observations); t++ { + next := make([]float64, states) + sym := observations[t] + for j := range states { + total := 0.0 + for k := range states { + total += row[k] * m.Transition[k*states+j] + } + next[j] = total * m.Emission[j*width+sym] + } + total := 0.0 + for _, v := range next { + total += v + } + if total == 0 { + return nil, nil, base.Errf("%s: the observation sequence has probability zero under the model at step %d", name, t) + } + scales[t] = total + for j := range states { + next[j] /= total + } + alpha[t] = next + row = next + } + return alpha, scales, nil +} + +// hmmBackward runs the scaled backward recursion against the forward +// scales, so the product alpha·beta normalises directly into the +// smoothed posteriors. +func (m *HiddenMarkovModel) hmmBackward(observations []int, scales []float64) [][]float64 { + states := m.hmmStates() + width := m.hmmWidth() + beta := make([][]float64, len(observations)) + last := make([]float64, states) + for k := range states { + last[k] = 1 + } + beta[len(observations)-1] = last + for t := len(observations) - 2; t >= 0; t-- { + row := make([]float64, states) + sym := observations[t+1] + for k := range states { + total := 0.0 + for j := range states { + total += m.Transition[k*states+j] * m.Emission[j*width+sym] * beta[t+1][j] + } + row[k] = total / scales[t+1] + } + beta[t] = row + } + return beta +} + +// Forward runs the scaled forward recursion over the observation +// sequence and returns the filtered state posteriors, one row of +// state probabilities per step, with the sequence's log likelihood +// P(observations | model). Refuses a nil model, an empty sequence, +// any observation outside the model's symbol set, and a sequence of +// probability zero under the model, whose structural zero an emission +// or transition carries leaves the filtered posteriors undefined. +func (m *HiddenMarkovModel) Forward(observations []int) (filtered [][]float64, logLikelihood float64, err error) { + const name = "HiddenMarkovModel.Forward" + if m == nil { + return nil, 0, base.Errf("%s: the model is nil", name) + } + if err := hmmCheckSequence(name, observations, m.hmmSymbols()); err != nil { + return nil, 0, err + } + alpha, scales, ferr := m.hmmForward(observations) + if ferr != nil { + return nil, 0, ferr + } + return alpha, hmmLogLikelihood(scales), nil +} + +// Smooth runs the forward-backward recursions over the observation +// sequence and returns the smoothed state posteriors P(state at t | +// all observations), one row per step, with the sequence's log +// likelihood. The refusals are Forward's own. +func (m *HiddenMarkovModel) Smooth(observations []int) (smoothed [][]float64, logLikelihood float64, err error) { + const name = "HiddenMarkovModel.Smooth" + if m == nil { + return nil, 0, base.Errf("%s: the model is nil", name) + } + if err := hmmCheckSequence(name, observations, m.hmmSymbols()); err != nil { + return nil, 0, err + } + alpha, scales, ferr := m.hmmForward(observations) + if ferr != nil { + return nil, 0, ferr + } + beta := m.hmmBackward(observations, scales) + gamma := make([][]float64, len(observations)) + for t := range observations { + row := make([]float64, len(m.Initial)) + total := 0.0 + for k := range row { + row[k] = alpha[t][k] * beta[t][k] + total += row[k] + } + for k := range row { + row[k] /= total + } + gamma[t] = row + } + return gamma, hmmLogLikelihood(scales), nil +} + +// hmmLogLikelihood folds the forward scales into the sequence's log +// likelihood. +func hmmLogLikelihood(scales []float64) float64 { + total := 0.0 + for _, s := range scales { + total += math.Log(s) + } + return total +} + +// Viterbi decodes the most likely state path through the observation +// sequence in log space, breaking ties toward the lowest state index, +// and returns the path with its log probability +// log P(path, observations | model). The refusals are Forward's own. +func (m *HiddenMarkovModel) Viterbi(observations []int) (states []int, logProbability float64, err error) { + const name = "HiddenMarkovModel.Viterbi" + if m == nil { + return nil, 0, base.Errf("%s: the model is nil", name) + } + if err := hmmCheckSequence(name, observations, m.hmmSymbols()); err != nil { + return nil, 0, err + } + n := m.hmmStates() + width := m.hmmWidth() + scores := make([]float64, n) + for k := range n { + scores[k] = math.Log(m.Initial[k]) + math.Log(m.Emission[k*width+observations[0]]) + } + back := make([][]int, len(observations)) + back[0] = make([]int, n) + for t := 1; t < len(observations); t++ { + next := make([]float64, n) + prev := make([]int, n) + sym := observations[t] + for j := range n { + best := math.Inf(-1) + arg := 0 + for k := range n { + candidate := scores[k] + math.Log(m.Transition[k*n+j]) + if candidate > best { + best = candidate + arg = k + } + } + next[j] = best + math.Log(m.Emission[j*width+sym]) + prev[j] = arg + } + scores = next + back[t] = prev + } + best := math.Inf(-1) + arg := 0 + for k := range n { + if scores[k] > best { + best = scores[k] + arg = k + } + } + path := make([]int, len(observations)) + path[len(observations)-1] = arg + for t := len(observations) - 1; t > 0; t-- { + path[t-1] = back[t][path[t]] + } + return path, best, nil +} + +// FitHiddenMarkovModel fits a hidden Markov model of the given state +// and symbol counts to the observation sequence by Baum-Welch +// expectation maximisation. The starting point draws every +// distribution row from the flat Dirichlet through the house +// generator, so the fit is deterministic for the generator state; +// expectation maximisation climbs to a local maximum of the +// likelihood, and a different generator state may land on a different +// one. +// +// The sweeps stop once the log likelihood moves by less than +// hmmTolerance relative, within hmmMaxIterations sweeps; Converged +// names which happened. The re-estimation floors every parameter at +// hmmFloor and renormalises the rows, so no symbol or transition is +// silenced outright by one sweep. A nil generator, a state or symbol +// count below one, an empty sequence or an observation outside the +// symbol set is an error. +func FitHiddenMarkovModel(g *core.Generator, observations []int, states, symbols int) (*HiddenMarkovFitResult, error) { + const name = "FitHiddenMarkovModel" + if g == nil { + return nil, base.Errf("%s: the generator is nil", name) + } + if states < 1 { + return nil, base.Errf("%s: the model needs at least one state, got %d", name, states) + } + if symbols < 1 { + return nil, base.Errf("%s: the model needs at least one symbol, got %d", name, symbols) + } + if err := hmmCheckSequence(name, observations, symbols); err != nil { + return nil, err + } + initial := hmmDrawRow(g, states) + transition := make([]float64, 0, states*states) + for range states { + transition = append(transition, hmmDrawRow(g, states)...) + } + emission := make([]float64, 0, states*symbols) + for range states { + emission = append(emission, hmmDrawRow(g, symbols)...) + } + model, err := NewHiddenMarkovModel(initial, transition, emission) + if err != nil { + return nil, err + } + prev := math.Inf(-1) + converged := false + iterations := hmmMaxIterations + for iter := 1; iter <= hmmMaxIterations; iter++ { + logLik, gamma, xi, serr := model.hmmSweep(observations) + if serr != nil { + return nil, serr + } + move := logLik - prev + prev = logLik + if iter > 1 && math.Abs(move) <= hmmTolerance*(1+math.Abs(logLik)) { + // The model has not moved since the sweep above, so the + // log likelihood in hand is the fitted model's own. + converged = true + iterations = iter + break + } + model = hmmReestimate(model, observations, gamma, xi) + } + if !converged { + // The budget ran out after a re-estimation: the reported + // likelihood is the final model's own, not its predecessor's. + logLik, _, _, serr := model.hmmSweep(observations) + if serr != nil { + return nil, serr + } + prev = logLik + } + return &HiddenMarkovFitResult{ + Model: model, + LogLikelihood: prev, + Iterations: iterations, + Converged: converged, + }, nil +} + +// hmmDrawRow draws one distribution row of the given width from the +// flat Dirichlet through the generator: the normalised exponentials +// of uniform draws, the classic construction of the simplex's uniform +// distribution. +func hmmDrawRow(g *core.Generator, width int) []float64 { + row := make([]float64, width) + total := 0.0 + for i := range width { + weight := -math.Log(max(g.Unit(), 1e-300)) + row[i] = weight + total += weight + } + for i := range row { + row[i] /= total + } + return row +} + +// hmmSweep runs the forward-backward pass at the model's current +// parameters: the log likelihood, the smoothed posteriors gamma and +// the pairwise posteriors xi the M step reads. xi[t] is the +// states×states row-major table P(state at t, state at t+1 | +// observations) for the step from t to t+1. +func (m *HiddenMarkovModel) hmmSweep(observations []int) (float64, [][]float64, [][]float64, error) { + states := m.hmmStates() + width := m.hmmWidth() + alpha, scales, ferr := m.hmmForward(observations) + if ferr != nil { + return 0, nil, nil, ferr + } + beta := m.hmmBackward(observations, scales) + logLik := hmmLogLikelihood(scales) + gamma := make([][]float64, len(observations)) + for t := range observations { + row := make([]float64, states) + total := 0.0 + for k := range states { + row[k] = alpha[t][k] * beta[t][k] + total += row[k] + } + for k := range states { + row[k] /= total + } + gamma[t] = row + } + xi := make([][]float64, len(observations)-1) + for t := range xi { + table := make([]float64, states*states) + total := 0.0 + sym := observations[t+1] + for k := range states { + for j := range states { + v := alpha[t][k] * m.Transition[k*states+j] * m.Emission[j*width+sym] * beta[t+1][j] + table[k*states+j] = v + total += v + } + } + for i := range table { + table[i] /= total + } + xi[t] = table + } + return logLik, gamma, xi, nil +} + +// hmmReestimate applies one Baum-Welch M step: the counts the +// posteriors carry become the new rows, floored and renormalised. +func hmmReestimate(model *HiddenMarkovModel, observations []int, gamma, xi [][]float64) *HiddenMarkovModel { + states := model.hmmStates() + symbols := model.hmmSymbols() + initial := make([]float64, states) + copy(initial, gamma[0]) + transition := make([]float64, states*states) + for t := range xi { + for i, v := range xi[t] { + transition[i] += v + } + } + for k := range states { + den := 0.0 + for j := range states { + den += transition[k*states+j] + } + for j := range states { + transition[k*states+j] = math.Max(transition[k*states+j]/den, hmmFloor) + } + hmmRenormalise(transition[k*states : (k+1)*states]) + } + emission := make([]float64, states*symbols) + for t, row := range gamma { + sym := observations[t] + for k := range states { + emission[k*symbols+sym] += row[k] + } + } + for k := range states { + den := 0.0 + for s := range symbols { + den += emission[k*symbols+s] + } + for s := range symbols { + emission[k*symbols+s] = math.Max(emission[k*symbols+s]/den, hmmFloor) + } + hmmRenormalise(emission[k*symbols : (k+1)*symbols]) + } + // The floors and renormalisations above keep every row a valid + // distribution, so the constructor's refusals cannot fire here. + fitted, _ := NewHiddenMarkovModel(initial, transition, emission) + return fitted +} + +// hmmRenormalise scales one distribution row back to sum 1. +func hmmRenormalise(row []float64) []float64 { + total := 0.0 + for _, v := range row { + total += v + } + for i := range row { + row[i] /= total + } + return row +} diff --git a/stats/hmm_test.go b/stats/hmm_test.go new file mode 100644 index 0000000..3efcce8 --- /dev/null +++ b/stats/hmm_test.go @@ -0,0 +1,345 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "math/big" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The hidden Markov machinery against exact referents. The model +// probabilities are rational, so a short sequence's likelihood and +// its most likely path are computed exactly by enumerating every +// path, and the recursions are compared against that enumeration, not +// against quoted figures. + +// exactModel holds one small model's parameters as rationals. +type exactModel struct { + initial []float64 + transition []float64 // 2×2 + emission []float64 // 2×2 +} + +// ratRow converts a float row to rationals. +func ratRow(vals []float64) []*big.Rat { + out := make([]*big.Rat, len(vals)) + for i, v := range vals { + out[i] = big.NewRat(int64(v*1e15), 1e15) + } + return out +} + +// exactPathProbability evaluates one path's joint probability with +// the observations it explains. +func exactPathProbability(m *exactModel, observations, path []int) *big.Rat { + initial := ratRow(m.initial) + trans := make([][]*big.Rat, 2) + for k := range 2 { + trans[k] = ratRow(m.transition[k*2 : (k+1)*2]) + } + emis := make([][]*big.Rat, 2) + for k := range 2 { + emis[k] = ratRow(m.emission[k*2 : (k+1)*2]) + } + total := new(big.Rat).Set(initial[path[0]]) + total.Mul(total, emis[path[0]][observations[0]]) + for t := 1; t < len(path); t++ { + total.Mul(total, trans[path[t-1]][path[t]]) + total.Mul(total, emis[path[t]][observations[t]]) + } + return total +} + +func TestHiddenMarkovAgainstExactEnumeration(t *testing.T) { + m := &exactModel{ + initial: []float64{0.6, 0.4}, + transition: []float64{0.7, 0.3, 0.2, 0.8}, + emission: []float64{0.9, 0.1, 0.25, 0.75}, + } + model, err := NewHiddenMarkovModel(m.initial, m.transition, m.emission) + if err != nil { + t.Fatal(err) + } + observations := []int{0, 0, 1, 0, 1, 1, 0, 1} + // Every path over the sequence, exactly. + bestProb := new(big.Rat) + totalProb := new(big.Rat) + var bestPath []int + path := make([]int, len(observations)) + var walk func(t int) + walk = func(t int) { + if t == len(path) { + p := exactPathProbability(m, observations, path) + totalProb.Add(totalProb, p) + if p.Cmp(bestProb) > 0 { + bestProb.Set(p) + bestPath = append([]int(nil), path...) + } + return + } + for _, s := range []int{0, 1} { + path[t] = s + walk(t + 1) + } + } + walk(0) + // The forward recursion's likelihood against the exact total. + _, ll, err := model.Forward(observations) + if err != nil { + t.Fatal(err) + } + wantLL, _ := new(big.Float).SetRat(totalProb).Float64() + if math.Abs(ll-math.Log(wantLL)) > 1e-12 { + t.Fatalf("the forward likelihood = %.16f, want the exact %.16f", ll, math.Log(wantLL)) + } + // Viterbi against the exact best path. + states, pathProb, err := model.Viterbi(observations) + if err != nil { + t.Fatal(err) + } + wantBest, _ := new(big.Float).SetRat(bestProb).Float64() + if math.Abs(pathProb-math.Log(wantBest)) > 1e-12 { + t.Fatalf("Viterbi's log probability = %.16f, want the exact %.16f", pathProb, math.Log(wantBest)) + } + for step := range states { + if states[step] != bestPath[step] { + t.Fatalf("Viterbi's path %v leaves the exact best path %v at step %d", states, bestPath, step) + } + } + // The smoothing's last row must coincide with filtering's: both + // see the whole sequence by then. + filtered, _, err := model.Forward(observations) + if err != nil { + t.Fatal(err) + } + smoothed, _, err := model.Smooth(observations) + if err != nil { + t.Fatal(err) + } + for k := range filtered[len(observations)-1] { + if filtered[len(observations)-1][k] != smoothed[len(observations)-1][k] { + t.Fatalf("the last filtered posterior %v disagrees with the smoothed one %v", + filtered[len(observations)-1], smoothed[len(observations)-1]) + } + if smoothed[len(observations)-1][k] < 0 || smoothed[len(observations)-1][k] > 1 { + t.Fatalf("the smoothed posterior %g is not a probability", smoothed[len(observations)-1][k]) + } + } +} + +func TestHiddenMarkovSingleStateLikelihood(t *testing.T) { + // One state: the model is an independent-symbol law and the + // likelihood is the product of the emissions, by hand. + model, err := NewHiddenMarkovModel([]float64{1}, []float64{1}, []float64{0.25, 0.75}) + if err != nil { + t.Fatal(err) + } + observations := []int{1, 1, 0, 1, 1, 1, 0, 1, 1} + _, ll, err := model.Forward(observations) + if err != nil { + t.Fatal(err) + } + want := 0.0 + for _, o := range observations { + if o == 0 { + want += math.Log(0.25) + } else { + want += math.Log(0.75) + } + } + if math.Abs(ll-want) > 1e-12 { + t.Fatalf("the single-state likelihood = %.16f, want %.16f", ll, want) + } + path, _, err := model.Viterbi(observations) + if err != nil { + t.Fatal(err) + } + for _, s := range path { + if s != 0 { + t.Fatalf("the single-state decode left the only state: %v", path) + } + } +} + +func TestHiddenMarkovScalingSurvivesLongSequences(t *testing.T) { + // A thousand-step sequence under a sticky model: the raw forward + // probabilities underflow long before the end, and the scaled + // recursions must answer a finite likelihood and rows that stay + // distributions. + model, err := NewHiddenMarkovModel( + []float64{0.5, 0.5}, + []float64{0.99, 0.01, 0.02, 0.98}, + []float64{0.8, 0.2, 0.3, 0.7}) + if err != nil { + t.Fatal(err) + } + observations := make([]int, 1000) + for t := range observations { + observations[t] = (t * 7 / 3) % 2 + } + filtered, ll, err := model.Forward(observations) + if err != nil { + t.Fatal(err) + } + if math.IsInf(ll, 0) || math.IsNaN(ll) { + t.Fatalf("the long sequence's likelihood = %g", ll) + } + for step, row := range filtered { + total := 0.0 + for _, v := range row { + total += v + } + if math.Abs(total-1) > 1e-9 { + t.Fatalf("the filtered row at %d sums to %g", step, total) + } + } +} + +func TestHiddenMarkovFitRecovery(t *testing.T) { + // Baum-Welch on a sequence drawn from a known model: the fitted + // likelihood must clear the known model's own, the fit must + // converge, and a repeat run must be bit-identical. + truth, err := NewHiddenMarkovModel( + []float64{0.65, 0.35}, + []float64{0.8, 0.2, 0.15, 0.85}, + []float64{0.85, 0.15, 0.3, 0.7}) + if err != nil { + t.Fatal(err) + } + observations := hmmSimulate(t, truth, 400, 20260925) + _, truthLL, err := truth.Forward(observations) + if err != nil { + t.Fatal(err) + } + fit, err := FitHiddenMarkovModel(core.NewGenerator(7), observations, 2, 2) + if err != nil { + t.Fatal(err) + } + if !fit.Converged { + t.Fatalf("the fit did not converge (%d iterations)", fit.Iterations) + } + if fit.LogLikelihood < truthLL { + t.Fatalf("the fitted likelihood %.6f sits below the truth's %.6f, expectation maximisation failed to climb", + fit.LogLikelihood, truthLL) + } + again, err := FitHiddenMarkovModel(core.NewGenerator(7), observations, 2, 2) + if err != nil { + t.Fatal(err) + } + if again.LogLikelihood != fit.LogLikelihood || again.Iterations != fit.Iterations { + t.Fatal("two identical fits disagreed") + } + for i := range fit.Model.Initial { + if fit.Model.Initial[i] != again.Model.Initial[i] || + fit.Model.Transition[i] != again.Model.Transition[i] || + fit.Model.Emission[i] != again.Model.Emission[i] { + t.Fatal("two identical fits returned different parameters") + } + } + // The fitted model answers the same likelihood through Forward as + // the fit reported. + _, ll, err := fit.Model.Forward(observations) + if err != nil { + t.Fatal(err) + } + if ll != fit.LogLikelihood { + t.Fatalf("the fitted model's likelihood %.16f differs from the fit's report %.16f", ll, fit.LogLikelihood) + } +} + +// hmmSimulate draws an observation sequence from a model through the +// house generator, by inverse transform on the cumulative rows. +func hmmSimulate(t *testing.T, model *HiddenMarkovModel, length int, seed int64) []int { + t.Helper() + g := core.NewGenerator(seed) + states := len(model.Initial) + symbols := len(model.Emission) / states + draw := func(row []float64) int { + u := g.Unit() + total := 0.0 + for i, v := range row { + total += v + if u < total { + return i + } + } + return len(row) - 1 + } + state := draw(model.Initial) + out := make([]int, length) + for t := range length { + out[t] = draw(model.Emission[state*symbols : (state+1)*symbols]) + state = draw(model.Transition[state*states : (state+1)*states]) + } + return out +} + +func TestHiddenMarkovRefusals(t *testing.T) { + model, err := NewHiddenMarkovModel([]float64{0.6, 0.4}, []float64{0.7, 0.3, 0.2, 0.8}, []float64{0.9, 0.1, 0.25, 0.75}) + if err != nil { + t.Fatal(err) + } + if _, _, err := model.Forward(nil); err == nil { + t.Fatal("an empty sequence was accepted") + } + if _, _, err := model.Forward([]int{0, 2, 1}); err == nil { + t.Fatal("an out-of-range symbol was accepted") + } + if _, err := NewHiddenMarkovModel([]float64{0.6, 0.5}, []float64{0.7, 0.3, 0.2, 0.8}, []float64{0.9, 0.1, 0.25, 0.75}); err == nil { + t.Fatal("an initial distribution summing over 1 was accepted") + } + if _, err := NewHiddenMarkovModel([]float64{0.6, 0.4}, []float64{0.7, 0.3, 0.2}, []float64{0.9, 0.1, 0.25, 0.75}); err == nil { + t.Fatal("a short transition matrix was accepted") + } + if _, err := NewHiddenMarkovModel([]float64{0.6, 0.4}, []float64{0.7, 0.3, 0.2, 0.8}, []float64{0.9, 0.1, 0.25}); err == nil { + t.Fatal("a ragged emission matrix was accepted") + } + if _, err := NewHiddenMarkovModel([]float64{1.2, -0.2}, []float64{0.7, 0.3, 0.2, 0.8}, []float64{0.9, 0.1, 0.25, 0.75}); err == nil { + t.Fatal("a probability outside [0, 1] was accepted") + } + if _, err := FitHiddenMarkovModel(nil, []int{0, 1}, 2, 2); err == nil { + t.Fatal("a nil generator was accepted") + } + if _, err := FitHiddenMarkovModel(core.NewGenerator(1), []int{0, 1, 5}, 2, 2); err == nil { + t.Fatal("an out-of-range training symbol was accepted") + } + var nilModel *HiddenMarkovModel + if _, _, err := nilModel.Forward([]int{0}); err == nil { + t.Fatal("a nil model was accepted") + } +} + +func TestHiddenMarkovZeroProbabilitySequence(t *testing.T) { + // A structural zero in the emissions: the only state never emits + // symbol 1, so a sequence holding it has probability zero and no + // posterior exists. The recursions must refuse rather than divide + // the zero normaliser into NaN posteriors. + model, err := NewHiddenMarkovModel([]float64{1}, []float64{1}, []float64{1, 0}) + if err != nil { + t.Fatal(err) + } + if _, _, err := model.Forward([]int{1}); err == nil { + t.Fatal("a zero-probability sequence was accepted by Forward") + } + if _, _, err := model.Smooth([]int{1}); err == nil { + t.Fatal("a zero-probability sequence was accepted by Smooth") + } + // A sequence that dies one step in: the first symbol is live, the + // second is not, and the refusal is the same. + if _, _, err := model.Forward([]int{0, 1}); err == nil { + t.Fatal("a sequence that dies at step 1 was accepted by Forward") + } + // The live half of the same model still answers: the zero emission + // silences nothing a legal sequence needs. + filtered, ll, err := model.Forward([]int{0, 0}) + if err != nil { + t.Fatal(err) + } + if ll != 0 || filtered[0][0] != 1 { + t.Fatalf("the certain sequence answered (%g, %v), want (0, [1])", ll, filtered[0]) + } +} diff --git a/stats/inference.go b/stats/inference.go new file mode 100644 index 0000000..c04ae18 --- /dev/null +++ b/stats/inference.go @@ -0,0 +1,432 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import ( + "math" + "slices" +) + +// Statistical inference on samples: association matrices, group +// comparison, goodness of fit, distribution comparison and the +// percentile bootstrap. Everything here is classical frequentist +// tooling built on the distribution functions of cdf.go, so a +// p-value travels no further than the library's own incomplete +// gamma and beta. + +// CovarianceMatrix returns the sample covariance matrix of the +// observations in a, an (n, p) array whose rows are observations and +// columns variables: the (p, p) result has cov(i, j) the sample +// covariance of columns i and j with the 1/(n−1) normalisation. At +// least two observations are needed. +func CovarianceMatrix(a *core.Array) (*core.Array, error) { + return scatterMatrix("CovarianceMatrix", a, false) +} + +// CorrelationMatrix returns the Pearson correlation matrix of the +// observations in a, the covariance matrix normalised by each +// column's sample standard deviation; the diagonal is exactly 1. +func CorrelationMatrix(a *core.Array) (*core.Array, error) { + return scatterMatrix("CorrelationMatrix", a, true) +} + +// scatterMatrix builds the covariance or correlation matrix of an +// (n, p) observation array. +func scatterMatrix(name string, a *core.Array, correlate bool) (*core.Array, error) { + if a.NDim() != 2 { + return nil, base.Errf("%s: needs a 2-D array of observations, got shape %s", + name, base.ShapeText(a.Shape())) + } + n, p := a.Shape()[0], a.Shape()[1] + if n < 2 { + return nil, base.Errf("%s: at least two observations are needed, got %d", name, n) + } + if a.Dtype() == core.Complex { + return nil, base.Errf("%s: complex observations are not supported", name) + } + if err := checkFinite(name, "the observations", a); err != nil { + return nil, err + } + // Column means, then centred column copies. A dense float64 payload + // is read straight from the raw slice, one dispatch-free pass. + rows := rawFloats(a) + means := make([]float64, p) + for j := range p { + s := 0.0 + if rows != nil { + for i := range n { + s += rows[i*p+j] + } + } else { + for i := range n { + s += a.FloatAt(i*p + j) + } + } + means[j] = s / float64(n) + } + cols := make([][]float64, p) + for j := range p { + cols[j] = make([]float64, n) + if rows != nil { + for i := range n { + cols[j][i] = rows[i*p+j] - means[j] + } + } else { + for i := range n { + cols[j][i] = a.FloatAt(i*p+j) - means[j] + } + } + } + scale := float64(n - 1) + cov := make([]float64, p*p) + stds := make([]float64, p) + for i := range p { + for j := i; j < p; j++ { + s := 0.0 + for k := range n { + s += cols[i][k] * cols[j][k] + } + cov[i*p+j] = s / scale + cov[j*p+i] = cov[i*p+j] + } + if correlate { + stds[i] = math.Sqrt(cov[i*p+i]) + if stds[i] == 0 { + return nil, base.Errf("CorrelationMatrix: column %d has zero variance", i) + } + } + } + if correlate { + for i := range p { + for j := i; j < p; j++ { + c := cov[i*p+j] / (stds[i] * stds[j]) + cov[i*p+j] = c + cov[j*p+i] = c + } + cov[i*p+i] = 1 + } + } + return floatsToArray(cov, []int{p, p}), nil +} + +// WelchTTest compares two independent samples with Welch's t-test, +// which asks no equal-variance assumption: the statistic is +// (mean(a) − mean(b)) over the pooled standard error of the two +// means, and the degrees of freedom follow the Welch-Satterthwaite +// estimate. Returns the statistic, the (generally fractional) degrees +// of freedom and the two-sided p-value. Both samples need at least +// two observations. +func WelchTTest(a, b *core.Array) (t, df, pValue float64, err error) { + m1, v1, n1, err := sampleMeanVar(a, "WelchTTest") + if err != nil { + return 0, 0, 0, err + } + m2, v2, n2, err := sampleMeanVar(b, "WelchTTest") + if err != nil { + return 0, 0, 0, err + } + se1 := v1 / float64(n1) + se2 := v2 / float64(n2) + se := se1 + se2 + if se == 0 { + return 0, 0, 0, base.Errf("WelchTTest: both samples have zero variance") + } + t = (m1 - m2) / math.Sqrt(se) + df = se * se / (se1*se1/float64(n1-1) + se2*se2/float64(n2-1)) + // Two-sided p = P(T > |t|) + P(T < −|t|) = I_z(df/2, 1/2) with + // z = df/(df + t²), the Student-t tail in closed form. + z := df / (df + t*t) + pValue, err = BetaIncomplete(z, df/2, 0.5) + if err != nil { + return 0, 0, 0, base.Errf("WelchTTest: %w", err) + } + return t, df, pValue, nil +} + +// sampleMeanVar returns the mean and the 1/(n−1) variance of a sample. +func sampleMeanVar(a *core.Array, name string) (mean, variance float64, n int, err error) { + if a.Dtype() == core.Complex { + return 0, 0, 0, base.Errf("%s: complex samples are not supported", name) + } + n = a.Len() + if n < 2 { + return 0, 0, 0, base.Errf("%s: at least two observations are needed, got %d", name, n) + } + // The finiteness refusal belongs here, before the arithmetic: a NaN + // sample would otherwise surface only when the tail function + // rejects the NaN statistic it produced, under the tail function's + // own name. + if err := checkFinite(name, "the sample", a); err != nil { + return 0, 0, 0, err + } + mean = 0.0 + // The walk is bounded by the sample's own element count: a rebased + // view's payload may run past its visible elements, and those + // invisible tail slots are nobody's observations. + if fs := rawFloats(a); fs != nil { + fs = fs[:n] + for _, v := range fs { + mean += v + } + mean /= float64(n) + variance = sqDeviations(fs, mean) + } else { + for i := range n { + mean += a.FloatAt(i) + } + mean /= float64(n) + variance = sqDeviationsAt(a, mean) + } + variance /= float64(n - 1) + return mean, variance, n, nil +} + +// ChiSquareGoodnessOfFit runs Pearson's test of an observed frequency +// table against expected ones: the statistic is Σ(O−E)²/E over the +// bins, the degrees of freedom the bin count minus one, and the +// p-value the upper tail of that χ² distribution. Every expected +// entry must be positive. +func ChiSquareGoodnessOfFit(observed, expected *core.Array) (chi2 float64, df int, pValue float64, err error) { + if observed.Dtype() == core.Complex || expected.Dtype() == core.Complex { + return 0, 0, 0, base.Errf("ChiSquareGoodnessOfFit: complex arrays are not supported") + } + n := observed.Len() + if expected.Len() != n { + return 0, 0, 0, base.Errf("ChiSquareGoodnessOfFit: observed has %d bins, expected %d", + n, expected.Len()) + } + if n < 2 { + return 0, 0, 0, base.Errf("ChiSquareGoodnessOfFit: at least two bins are needed, got %d", n) + } + // A NaN observation would only surface later, rejected by the + // tail function under its own name; the scan keeps the error with + // the entry point the caller actually used. + if err := checkFinite("ChiSquareGoodnessOfFit", "the observed frequencies", observed); err != nil { + return 0, 0, 0, err + } + // A non-finite expectation passes the positivity gate below (+Inf + // is positive) and turns the statistic into a NaN the tail function + // would reject under its own, misleading, name. + if err := checkFinite("ChiSquareGoodnessOfFit", "the expected frequencies", expected); err != nil { + return 0, 0, 0, err + } + chi2 = 0.0 + // Both walks are bounded by the bin count: a rebased view's payload + // may run past its visible bins, and indexing one past the other's + // payload would be a panic before it was a wrong statistic. + obs := rawFloats(observed) + exp := rawFloats(expected) + if obs != nil { + obs = obs[:n] + } + if exp != nil { + exp = exp[:n] + } + if obs != nil && exp != nil { + for i, e := range exp { + if !(e > 0) { + return 0, 0, 0, base.Errf("ChiSquareGoodnessOfFit: expected[%d] = %g must be positive", i, e) + } + d := obs[i] - e + chi2 += d * d / e + } + } else { + for i := range n { + e := expected.FloatAt(i) + if !(e > 0) { + return 0, 0, 0, base.Errf("ChiSquareGoodnessOfFit: expected[%d] = %g must be positive", i, e) + } + d := observed.FloatAt(i) - e + chi2 += d * d / e + } + } + df = n - 1 + pValue, err = GammaUpper(float64(df)/2, chi2/2) + if err != nil { + return 0, 0, 0, base.Errf("ChiSquareGoodnessOfFit: %w", err) + } + return chi2, df, pValue, nil +} + +// KolmogorovSmirnovTest compares two samples by the largest vertical +// distance d between their empirical distribution functions. The +// p-value is the asymptotic Kolmogorov distribution evaluated at +// λ = d·√(nm/(n+m)), accurate for samples of a few dozen and upward; +// small samples only get an order-of-magnitude answer. Both samples +// must be finite: a NaN never compares true against the support point +// the merge walk advances on, so it would leave the walk stuck, and an +// ±Inf would distort the distance. +func KolmogorovSmirnovTest(a, b *core.Array) (d, pValue float64, err error) { + const name = "KolmogorovSmirnovTest" + if a.Len() == 0 || b.Len() == 0 { + return 0, 0, base.Errf("%s: both samples must be non-empty", name) + } + if a.Dtype() == core.Complex || b.Dtype() == core.Complex { + return 0, 0, base.Errf("%s: complex samples are not supported", name) + } + if err := checkFinite(name, "the first sample", a); err != nil { + return 0, 0, err + } + if err := checkFinite(name, "the second sample", b); err != nil { + return 0, 0, err + } + xs := make([]float64, a.Len()) + if fs := rawFloats(a); fs != nil { + copy(xs, fs) + } else { + for i := range a.Len() { + xs[i] = a.FloatAt(i) + } + } + ys := make([]float64, b.Len()) + if fs := rawFloats(b); fs != nil { + copy(ys, fs) + } else { + for i := range b.Len() { + ys[i] = b.FloatAt(i) + } + } + slices.Sort(xs) + slices.Sort(ys) + // Walk the merged support, tracking both empirical CDFs. + i, j := 0, 0 + d = 0.0 + for i < len(xs) || j < len(ys) { + var x float64 + switch { + case i < len(xs) && (j >= len(ys) || xs[i] <= ys[j]): + x = xs[i] + default: + x = ys[j] + } + for i < len(xs) && xs[i] <= x { + i++ + } + for j < len(ys) && ys[j] <= x { + j++ + } + if gap := math.Abs(float64(i)/float64(len(xs)) - float64(j)/float64(len(ys))); gap > d { + d = gap + } + } + en := math.Sqrt(float64(a.Len()) * float64(b.Len()) / float64(a.Len()+b.Len())) + pValue = kolmogorovTail(d * en) + return d, pValue, nil +} + +// kolmogorovEps is the magnitude at which a term of the Kolmogorov +// series is below the rounding level of the sum, so the alternating +// series can be truncated there: the first omitted term bounds the +// truncation error, and the terms decrease from the first one on. +const kolmogorovEps = 1e-18 + +// maxKolmogorovTerms caps the series for the smallest lambdas that +// reach it at all. +const maxKolmogorovTerms = 100 + +// kolmogorovTermCount returns how many terms of the Kolmogorov series +// matter at lambda: the first k whose term 2·exp(−2k²λ²) has fallen +// below the rounding level. The series itself sums exactly that many, +// so the count is the truncation rule in one place instead of a +// condition the loop cannot reach. +func kolmogorovTermCount(lambda float64) int { + for k := 1; k <= maxKolmogorovTerms; k++ { + if 2*math.Exp(-2*float64(k)*float64(k)*lambda*lambda) < kolmogorovEps { + return k + } + } + return maxKolmogorovTerms +} + +// kolmogorovTail evaluates the asymptotic Kolmogorov distribution +// Q(λ) = 2·Σ_{k≥1} (−1)^{k−1}·e^{−2k²λ²}, summing the terms that +// matter. +func kolmogorovTail(lambda float64) float64 { + if lambda < 0.2 { + return 1 + } + total := 0.0 + for k := 1; k <= kolmogorovTermCount(lambda); k++ { + term := 2 * math.Exp(-2*float64(k)*float64(k)*lambda*lambda) + if k%2 == 0 { + total -= term + } else { + total += term + } + } + return min(1, max(0, total)) +} + +// BootstrapCI estimates a confidence interval of a statistic by the +// percentile bootstrap: resamples the data with replacement +// resamples times from the seeded generator, evaluates the statistic +// on every resample and returns the alpha/2 and 1−alpha/2 quantiles +// of the bootstrap distribution, alpha = 1 − level. The statistic +// receives the resample as a fresh array and may return its own +// error. The data must be real-valued. +func BootstrapCI(data *core.Array, statistic func(*core.Array) (float64, error), level float64, + resamples int, seed int64) (lower, upper float64, err error) { + if data.Len() == 0 { + return 0, 0, base.Errf("BootstrapCI: the data must not be empty") + } + if data.Dtype() == core.Complex { + return 0, 0, base.Errf("BootstrapCI: complex data are not supported") + } + if !(level > 0 && level < 1) { + return 0, 0, base.Errf("BootstrapCI: level must lie in (0, 1), got %g", level) + } + if resamples < 2 { + return 0, 0, base.Errf("BootstrapCI: resamples must be ≥ 2, got %d", resamples) + } + g := core.NewGenerator(seed) + n := data.Len() + // The source values are read once and the index payload hoisted: + // the resample loop then indexes plain slices. + dataVals := rawFloats(data) + values := make([]float64, resamples) + for r := range resamples { + idx, ierr := core.Ints(g, n, 0, int64(n)) + if ierr != nil { + return 0, 0, base.Errf("BootstrapCI: %w", ierr) + } + idxVals := idx.RawInts() + // A fresh slice per resample honours the documented contract: a + // statistic that retains its argument must not observe the next + // resample's mutation through the alias, and the resample is + // written into that slice directly rather than into a scratch the + // caller then copies. + sample := make([]float64, n) + if dataVals != nil { + for i := range n { + sample[i] = dataVals[idxVals[i]] + } + } else { + for i := range n { + sample[i] = data.FloatAt(int(idxVals[i])) + } + } + v, serr := statistic(wrapVector(sample)) + if serr != nil { + return 0, 0, base.Errf("BootstrapCI: %w", serr) + } + values[r] = v + } + alpha := (1 - level) / 2 + qs, qerr := Quantile(floatsToArray(values, []int{resamples}), []float64{alpha, 1 - alpha}) + if qerr != nil { + return 0, 0, base.Errf("BootstrapCI: %w", qerr) + } + return qs.FloatAt(0), qs.FloatAt(1), nil +} + +// wrapVector views a float64 slice as a rank-1 Array without copying. +func wrapVector(v []float64) *core.Array { + a, _ := core.FloatsFromArray(v, len(v)) + return a +} diff --git a/stats/inference_test.go b/stats/inference_test.go new file mode 100644 index 0000000..95be0fd --- /dev/null +++ b/stats/inference_test.go @@ -0,0 +1,198 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "testing" +) + +// TestCovarianceMatrixExact checks the covariance of y = 2x + 1 over +// x = 1..4: every sample variance is 5/3 and the cross covariance +// 10/3, all exact fractions. +func TestCovarianceMatrixExact(t *testing.T) { + obs := mustFloats(t, []float64{ + 1, 3, + 2, 5, + 3, 7, + 4, 9, + }, 4, 2) + cov, err := CovarianceMatrix(obs) + if err != nil { + t.Fatalf("CovarianceMatrix: %v", err) + } + if cov.Shape()[0] != 2 || cov.Shape()[1] != 2 { + t.Fatalf("covariance shape %v, want (2, 2)", cov.Shape()) + } + for _, want := range []struct { + i, j int + v float64 + }{ + {0, 0, 5.0 / 3}, {1, 1, 20.0 / 3}, {0, 1, 10.0 / 3}, {1, 0, 10.0 / 3}, + } { + if math.Abs(cov.FloatAt(want.i*2+want.j)-want.v) > 1e-14 { + t.Fatalf("cov(%d, %d) = %.16g, want %.16g", want.i, want.j, + cov.FloatAt(want.i*2+want.j), want.v) + } + } + corr, err := CorrelationMatrix(obs) + if err != nil { + t.Fatalf("CorrelationMatrix: %v", err) + } + if math.Abs(corr.FloatAt(0)-1) > 1e-14 || math.Abs(corr.FloatAt(3)-1) > 1e-14 { + t.Fatalf("correlation diagonal = %g, %g, want exactly 1", + corr.FloatAt(0), corr.FloatAt(3)) + } + if math.Abs(corr.FloatAt(1)-1) > 1e-14 { + t.Fatalf("perfect linear relation has correlation %g, want 1", corr.FloatAt(1)) + } +} + +func TestCovarianceMatrixErrors(t *testing.T) { + if _, err := CovarianceMatrix(mustFloats(t, []float64{1, 2, 3, 4}, 4)); err == nil { + t.Fatal("1-D input: want an error") + } + single := mustFloats(t, []float64{1, 2}, 1, 2) + if _, err := CovarianceMatrix(single); err == nil { + t.Fatal("one observation: want an error") + } + constant := mustFloats(t, []float64{1, 2, 1, 2, 1, 2, 1, 2}, 4, 2) + if _, err := CorrelationMatrix(constant); err == nil { + t.Fatal("zero-variance column under correlation: want an error") + } +} + +// TestWelchTTestSeparated checks the statistic and degrees of freedom +// on samples with an exact analytic answer: a = [1, 2, 3], +// b = [4, 5, 6] gives t = −3√(3/2) and df = 2(n−1) = 4. +func TestWelchTTestSeparated(t *testing.T) { + a := mustFloats(t, []float64{1, 2, 3}) + b := mustFloats(t, []float64{4, 5, 6}) + stat, df, p, err := WelchTTest(a, b) + if err != nil { + t.Fatalf("WelchTTest: %v", err) + } + wantT := -3 / math.Sqrt(2.0/3) + if math.Abs(stat-wantT) > 1e-14 { + t.Fatalf("t = %.16g, want %.16g", stat, wantT) + } + if math.Abs(df-4) > 1e-14 { + t.Fatalf("df = %.16g, want 4", df) + } + // The p-value must equal the closed-form tail of the same + // distribution. + wantP, err := BetaIncomplete(4.0/(4.0+wantT*wantT), 2, 0.5) + if err != nil || math.Abs(p-wantP) > 1e-14 { + t.Fatalf("p = %v (%v), want %.16g", p, err, wantP) + } + if p > 0.2 { + t.Fatalf("shifted samples keep p = %g, want a small tail", p) + } + // Identical samples: statistic 0, p 1. + stat, _, p, err = WelchTTest(a, a) + if err != nil || stat != 0 || p < 1 { + t.Fatalf("identical samples: t = %v, p = %v, %v", stat, p, err) + } + if _, _, _, err := WelchTTest(mustFloats(t, []float64{1}), b); err == nil { + t.Fatal("one-observation sample: want an error") + } +} + +// TestChiSquareGoodnessOfFairDie pins the statistic of a near-fair +// die: (16, 18, 22, 21, 19, 24) against 20 each gives exactly 2.1. +func TestChiSquareGoodnessOfFairDie(t *testing.T) { + observed := mustFloats(t, []float64{16, 18, 22, 21, 19, 24}) + expected := mustFloats(t, []float64{20, 20, 20, 20, 20, 20}) + chi2, df, p, err := ChiSquareGoodnessOfFit(observed, expected) + if err != nil { + t.Fatalf("ChiSquareGoodnessOfFit: %v", err) + } + if math.Abs(chi2-2.1) > 1e-14 { + t.Fatalf("chi² = %.16g, want 2.1", chi2) + } + if df != 5 { + t.Fatalf("df = %d, want 5", df) + } + wantP, err := GammaUpper(2.5, 1.05) + if err != nil || math.Abs(p-wantP) > 1e-14 { + t.Fatalf("p = %v (%v), want %.16g", p, err, wantP) + } + if p < 0.5 || p > 0.95 { + t.Fatalf("near-fair die p = %g outside the plausible band", p) + } + if _, _, _, err := ChiSquareGoodnessOfFit(observed, mustFloats(t, []float64{20, 0, 20, 20, 20, 20})); err == nil { + t.Fatal("zero expected bin: want an error") + } +} + +// TestKolmogorovSmirnovShifted pins the statistic on a known shift: +// samples {1..4} and {2..5} have sup distance exactly 1/4. +func TestKolmogorovSmirnovShifted(t *testing.T) { + a := mustFloats(t, []float64{1, 2, 3, 4}) + b := mustFloats(t, []float64{2, 3, 4, 5}) + d, p, err := KolmogorovSmirnovTest(a, b) + if err != nil { + t.Fatalf("KolmogorovSmirnovTest: %v", err) + } + if math.Abs(d-0.25) > 1e-15 { + t.Fatalf("d = %.16g, want 0.25", d) + } + if p < 0.8 { + t.Fatalf("a one-step shift of four points keeps p = %g, want a large value", p) + } + // Identical samples: d = 0, p = 1. + d, p, err = KolmogorovSmirnovTest(b, b) + if err != nil || d != 0 || p < 1 { + t.Fatalf("identical samples: d = %v, p = %v, %v", d, p, err) + } + if _, _, err := KolmogorovSmirnovTest(mustFloats(t, []float64{}), b); err == nil { + t.Fatal("empty sample: want an error") + } +} + +// TestBootstrapCIMean resamples a sample whose mean is exactly 10: +// the interval must bracket 10 with a plausible width and stay +// deterministic under the seed. +func TestBootstrapCIMean(t *testing.T) { + data := mustFloats(t, []float64{ + 10.7, 9.3, 10.1, 8.9, 11.2, 9.8, 10.4, 9.1, 10.9, 9.6, + }) + mean := func(v *core.Array) (float64, error) { + return core.Mean(v) + } + lower, upper, err := BootstrapCI(data, mean, 0.9, 500, 42) + if err != nil { + t.Fatalf("BootstrapCI: %v", err) + } + m, merr := core.Mean(data) + if merr != nil { + t.Fatalf("Mean: %v", merr) + } + if !(lower <= m && m <= upper) { + t.Fatalf("interval [%g, %g] misses the sample mean %g", lower, upper, m) + } + if upper-lower <= 0 || upper-lower > 2 { + t.Fatalf("interval width %g is implausible", upper-lower) + } + l2, u2, err := BootstrapCI(data, mean, 0.9, 500, 42) + if err != nil || l2 != lower || u2 != upper { + t.Fatalf("seeded run not deterministic: [%g, %g] vs [%g, %g], %v", + l2, u2, lower, upper, err) + } + // A higher level must not be tighter. + l3, u3, err := BootstrapCI(data, mean, 0.99, 500, 42) + if err != nil || l3 > lower || u3 < upper { + t.Fatalf("99 %% interval [%g, %g] tighter than 90 %% [%g, %g], %v", + l3, u3, lower, upper, err) + } + boom := func(*core.Array) (float64, error) { return 0, base.Errf("statistic failed") } + if _, _, err := BootstrapCI(data, boom, 0.9, 10, 1); err == nil { + t.Fatal("statistic error: want an error") + } + if _, _, err := BootstrapCI(data, mean, 1.5, 10, 1); err == nil { + t.Fatal("level outside (0, 1): want an error") + } +} diff --git a/stats/kde.go b/stats/kde.go new file mode 100644 index 0000000..5dc337a --- /dev/null +++ b/stats/kde.go @@ -0,0 +1,125 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Kernel density estimation: the smoothed histogram that turns a +// finite sample into a continuous, everywhere-positive density +// estimate, with the bandwidth doing all the work. + +// kdeParallelMinPairs is the point-sample pair count one worker must +// carry before the density sweep splits across goroutines: below it the +// hand-off costs more than the kernel evaluations it would carry. +const kdeParallelMinPairs = 1 << 12 + +// KernelDensity evaluates the Gaussian-kernel density estimate of the +// sample at every point: each sample contributes a unit-variance +// normal of width bandwidth, and the estimate averages them. A +// non-positive bandwidth asks for Silverman's rule, +// 0.9·min(σ, IQR/1.34)·n^(−1/5), the plug-in width for Gaussian +// truth; a degenerate interquartile range falls back to σ alone. The +// sample and the points must be finite: a NaN or ±Inf entry is +// refused by name rather than spread through the estimate. +func KernelDensity(sample *core.Array, bandwidth float64, points *core.Array) (*core.Array, error) { + const name = "KernelDensity" + if sample.NDim() != 1 || points.NDim() != 1 { + return nil, base.Errf("%s: the sample and the points must be vectors", name) + } + if sample.Dtype() == core.Complex || points.Dtype() == core.Complex { + return nil, base.Errf("%s: complex samples are not supported", name) + } + if err := checkFinite(name, "the sample", sample); err != nil { + return nil, err + } + if err := checkFinite(name, "the points", points); err != nil { + return nil, err + } + n := sample.Len() + if n < 2 { + return nil, base.Errf("%s: the sample needs at least two points, got %d", name, n) + } + if bandwidth <= 0 { + bandwidth = silvermanBandwidth(sample) + } + if math.IsNaN(bandwidth) || bandwidth <= 0 { + return nil, base.Errf("%s: the bandwidth resolved to %g, want a positive width", name, bandwidth) + } + // The sample is read once into a plain slice: the O(n·points) sweep + // below then runs without per-pair accessor dispatch. The values are + // the ones FloatAt returns, so the estimate is unchanged bit for bit. + // A rebased view's payload may run past its own count, so the dense + // path is cut to the visible elements. + sampleVals := rawFloats(sample) + if sampleVals == nil { + sampleVals = make([]float64, n) + for i := range sampleVals { + sampleVals[i] = sample.FloatAt(i) + } + } else { + sampleVals = sampleVals[:n] + } + out := core.New(core.Float, points.Len()) + vals := out.RawFloats() + // The points are read once too: a dense float64 payload is walked in + // place, and any other layout is widened into a private slice, so the + // sweep below reads no accessor and every value is the one FloatAt + // returns. + pointVals := rawFloats(points) + if pointVals == nil { + pointVals = make([]float64, points.Len()) + for i := range pointVals { + pointVals[i] = points.FloatAt(i) + } + } else { + pointVals = pointVals[:points.Len()] + } + norm := 1 / (float64(n) * bandwidth * math.Sqrt(2*math.Pi)) + // An output point is written once by the worker that owns it and its + // sum keeps the sample's own order, so the split of the point range + // never moves a bit: the estimate is identical at any worker count. + density := func(start, end int) { + for i := start; i < end; i++ { + x := pointVals[i] + total := 0.0 + for _, s := range sampleVals { + z := (x - s) / bandwidth + total += math.Exp(-0.5 * z * z) + } + vals[i] = total * norm + } + } + engine.ParallelMin(points.Len(), max(1, (kdeParallelMinPairs+n-1)/n), density) + return out, nil +} + +// silvermanBandwidth computes the plug-in width from the sample's +// spread: 0.9·min(σ, IQR/1.34)·n^(−1/5), with the interquartile +// fallback to σ when the middle of the sample is degenerate. +func silvermanBandwidth(sample *core.Array) float64 { + n := float64(sample.Len()) + sigma, err := Std(sample) + if err != nil { + return 0 + } + qArr, err := Quantile(sample, []float64{0.25, 0.75}) + if err != nil { + return 0 + } + iqr := qArr.FloatAt(1) - qArr.FloatAt(0) + spread := sigma + if iqr > 0 { + spread = math.Min(spread, iqr/1.34) + } + if spread <= 0 { + return 0 + } + return 0.9 * spread * math.Pow(n, -0.2) +} diff --git a/stats/lasso.go b/stats/lasso.go new file mode 100644 index 0000000..9604aa6 --- /dev/null +++ b/stats/lasso.go @@ -0,0 +1,509 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regularised linear regression by coordinate descent: the lasso and +// its elastic net generalisation. The fit runs on the standardised +// design (every column centred and scaled to unit population standard +// deviation, the response centred), the penalty falls on the slopes +// and never on the intercept, and the standardisation is recorded on +// the result and inverted before the coefficients leave, so the +// numbers a caller receives apply to the design as supplied. The path +// over a documented lambda grid is warm-started: every fit after the +// first begins at the previous lambda's coefficients, which is why +// the path costs a fraction of the cold fits it replaces. + +// ElasticNetResult carries one regularised fit at a single lambda. +type ElasticNetResult struct { + // Intercept and Coefficients are the fit on the original scale of + // the design as supplied: the standardisation the solver worked on + // has been inverted. Coefficients holds one slope per design + // column, in order. + Intercept float64 + Coefficients []float64 + // Fitted and Residuals align with the rows of the design. + Fitted []float64 + Residuals []float64 + // ColumnMeans and ColumnScales record the standardisation the fit + // ran under: the solver saw (x − ColumnMeans)/ColumnScales, and + // the lambda is measured against the centred response in the + // response's own units. + ColumnMeans []float64 + ColumnScales []float64 + // Lambda and Alpha are the penalty and the L2 mixing actually used. + Lambda float64 + Alpha float64 + // Iterations counts full coordinate-descent cycles and Converged + // reports whether no slope moved more than the tolerance in the + // last cycle. An exhausted budget returns the fit found so far + // rather than an error: coordinate descent on this convex + // objective cannot diverge, so the iterate stays a usable answer + // and the flag says how much to trust it. + Iterations int + Converged bool +} + +// LassoPathResult carries a warm-started regularisation path. +type LassoPathResult struct { + // Alpha is the mixing the path ran with. + Alpha float64 + // Lambdas are the penalty values of the path in descending order. + // The documented grid is lassoGridSteps values log-spaced from + // lambdaMax down to lambdaMax/1000, where lambdaMax is the + // smallest penalty at which the lasso keeps every slope at zero, + // max_j |Σ_i x_ij (y_i − ȳ)|/n divided by alpha. For alpha = 0 + // the pure-lasso lambdaMax (alpha taken as 1) defines the same + // grid, so ridge walks exactly the path the lasso would. + Lambdas []float64 + // Intercepts and Coefficients hold one fit per lambda, on the + // original scale, in the order of Lambdas. + Intercepts []float64 + Coefficients [][]float64 + // Iterations counts the coordinate-descent cycles each lambda + // needed. The warm start is what keeps them low along the path: + // each fit but the first begins at the previous lambda's + // coefficients, far closer to the answer than the zero start a + // cold fit uses, and the counts fall accordingly. + Iterations []int + // Converged reports whether every fit on the path converged + // within the iteration budget. + Converged bool +} + +// The documented shape of the lambda grid and the documented stopping +// rule of the coordinate descent: lassoGridSteps log-spaced values +// covering three decades down from lambdaMax, and a fit is settled +// once no slope moves more than lassoTolerance in a full cycle, with +// lassoMaxIterations cycles the most any single fit may spend. +const ( + lassoGridSteps = 100 + lassoGridRatio = 1e-3 + lassoTolerance = 1e-8 + lassoMaxIter = 10000 +) + +// Lasso fits y by linear regression under the least absolute +// shrinkage penalty at a single lambda: the pure lasso, alpha = 1 of +// ElasticNet. The slopes and only the slopes are penalised; the +// intercept is never. See ElasticNet for the contract the two share. +func Lasso(x, y *core.Array, lambda float64) (*ElasticNetResult, error) { + return ElasticNet(x, y, lambda, 1) +} + +// ElasticNet fits y = X·β under the elastic net penalty +// +// (1/2n)·Σᵢ(yᵢ − β₀ − xᵢ·β)² + λ·(α·Σ|βⱼ| + (1−α)/2·Σβⱼ²), +// +// the convex combination of the lasso (alpha 1) and ridge (alpha 0) +// penalties, by coordinate descent on the standardised design. The +// design carries n rows and p columns exactly as LinearRegression's, +// but a rank-deficient design is not an error here: the penalty keeps +// the problem well posed, and more columns than rows is a legitimate +// use. n must be at least 2, the design must carry at least one +// non-constant column, alpha must lie in [0, 1] and lambda must not +// be negative; complex or non-finite input is an error. +// +// Two exact anchors hold. On an orthonormal design the lasso answer +// is the soft-thresholded projection, coefficient by coefficient, and +// the fit reproduces that closed form. At alpha = 0 the answer is the +// Tikhonov ridge system (ZᵀZ/n + λI)β = Zᵀ(y − ȳ)/n over the +// standardised design, and the fit agrees with the answer the shared +// LU solve delivers for that system. +func ElasticNet(x, y *core.Array, lambda, alpha float64) (*ElasticNetResult, error) { + const name = "ElasticNet" + if math.IsNaN(lambda) || lambda < 0 { + return nil, base.Errf("%s: lambda must not be negative, got %g", name, lambda) + } + if math.IsNaN(alpha) || alpha < 0 || alpha > 1 { + return nil, base.Errf("%s: alpha must lie in [0, 1], got %g", name, alpha) + } + n, p, fx, fy, err := lassoInputs(name, x, y) + if err != nil { + return nil, err + } + std, err := lassoStandardise(name, n, p, x, y, fx, fy) + if err != nil { + return nil, err + } + beta, iterations, converged := elasticNetFit(std, lambda, alpha, nil, &lassoWorkspace{}) + return lassoResult(name, n, p, x, y, fx, fy, std, beta, lambda, alpha, iterations, converged) +} + +// LassoPath fits the whole regularisation path over the documented +// lambda grid, warm-started: the first fit, at the largest lambda, +// starts from zero coefficients (which are the exact answer there for +// the lasso), and every later fit starts at the previous lambda's +// answer. alpha is the elastic net mixing as in ElasticNet, and the +// alpha = 0 ridge walks the same grid the pure lasso defines. The +// result records how many coordinate cycles each lambda spent, which +// is the evidence the warm start earns its keep. +func LassoPath(x, y *core.Array, alpha float64) (*LassoPathResult, error) { + const name = "LassoPath" + if math.IsNaN(alpha) || alpha < 0 || alpha > 1 { + return nil, base.Errf("%s: alpha must lie in [0, 1], got %g", name, alpha) + } + n, p, fx, fy, err := lassoInputs(name, x, y) + if err != nil { + return nil, err + } + std, err := lassoStandardise(name, n, p, x, y, fx, fy) + if err != nil { + return nil, err + } + // lambdaMax: the smallest penalty the lasso answers with an + // all-zero slope vector, divided by alpha because the elastic net + // subgradient carries only the alpha fraction as the L1 term. The + // ridge (alpha = 0) has no sparsity threshold to define a largest + // lambda, so it borrows the pure-lasso grid and walks the same + // path. A response orthogonal to every column gives lambdaMax 0: + // no penalty does anything, and the grid value is immaterial, so a + // unit grid keeps the documented shape. + gridAlpha := alpha + if gridAlpha == 0 { + gridAlpha = 1 + } + // The projection norms are computed with the same multiply by the + // reciprocal the coordinate sweep uses, so at the top of the grid + // the soft threshold compares identical bits against identical + // bits and the all-zero answer is exact, not nearly zero. + inv := 1 / float64(n) + lambdaMax := 0.0 + for j := range p { + col := std.xs[j] + s := 0.0 + for i := range n { + s += col[i] * std.yc[i] + } + if a := math.Abs(s*inv) / gridAlpha; a > lambdaMax { + lambdaMax = a + } + } + if lambdaMax == 0 { + lambdaMax = 1 + } + out := &LassoPathResult{ + Alpha: alpha, + Lambdas: make([]float64, lassoGridSteps), + Intercepts: make([]float64, lassoGridSteps), + Coefficients: make([][]float64, lassoGridSteps), + Iterations: make([]int, lassoGridSteps), + Converged: true, + } + warm := make([]float64, p) + ws := &lassoWorkspace{beta: make([]float64, p), r: make([]float64, n)} + for k := range lassoGridSteps { + lambda := lambdaMax * math.Pow(lassoGridRatio, float64(k)/float64(lassoGridSteps-1)) + beta, iterations, converged := elasticNetFit(std, lambda, alpha, warm, ws) + copy(warm, beta) + // The path records the fit itself and has no field for a + // residual: lifting the standardised slopes back to the design + // as supplied is the whole reading, so the per-lambda pass over + // the original design that Fitted and Residuals would need is + // not run. + intercept, coefficients := lassoCoefficients(std, beta) + out.Lambdas[k] = lambda + out.Intercepts[k] = intercept + out.Coefficients[k] = coefficients + out.Iterations[k] = iterations + out.Converged = out.Converged && converged + } + return out, nil +} + +// lassoStandardisation holds the standardised system the coordinate +// descent runs on: standardised columns stored column-wise for cache +// friendly sweeps, their squared norms, and the centred response. +type lassoStandardisation struct { + xs [][]float64 // the standardised columns, each of length n + zs []float64 // z_j = Σᵢ x²ᵢⱼ/n, the standardised column energies + yc []float64 // the centred response + yMean float64 // the response mean the intercept reabsorbs + means []float64 // the original column means + scales []float64 // the original column population standard deviations +} + +// lassoWorkspace holds the coordinate descent's running buffers: the +// slope vector and the residual the sweeps keep current. One +// workspace serves a whole path, so a hundred warm-started fits +// allocate none of them again. +type lassoWorkspace struct { + beta []float64 + r []float64 +} + +// lassoInputs validates a design and response exactly as the rest of +// the estimation entries do, and returns the row and column counts +// with the dense raw payloads when they exist. +func lassoInputs(name string, x, y *core.Array) (n, p int, fx, fy []float64, err error) { + if x.NDim() != 2 { + return 0, 0, nil, nil, base.Errf("%s: the design must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 { + return 0, 0, nil, nil, base.Errf("%s: the response must be rank 1", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return 0, 0, nil, nil, base.Errf("%s: complex inputs are not supported", name) + } + n, p = x.Shape()[0], x.Shape()[1] + if y.Len() != n { + return 0, 0, nil, nil, base.Errf("%s: the design has %d rows but the response %d", name, n, y.Len()) + } + if n < 2 { + return 0, 0, nil, nil, base.Errf("%s: at least two observations are needed, got %d", name, n) + } + if p == 0 { + return 0, 0, nil, nil, base.Errf("%s: the design must carry at least one column", name) + } + if err := checkFinite(name, "the design", x); err != nil { + return 0, 0, nil, nil, err + } + if err := checkFinite(name, "the response", y); err != nil { + return 0, 0, nil, nil, err + } + if rawFloats(x) != nil { + fx = rawFloats(x) + } + if rawFloats(y) != nil { + fy = rawFloats(y) + } + return n, p, fx, fy, nil +} + +// lassoStandardise builds the standardised system: every column +// centred and divided by its population standard deviation, the +// response centred. A constant column has no standard deviation to +// divide by and is refused: its slope is unidentified at any penalty. +func lassoStandardise(name string, n, p int, x, y *core.Array, fx, fy []float64) (*lassoStandardisation, error) { + std := &lassoStandardisation{ + xs: make([][]float64, p), + zs: make([]float64, p), + yc: make([]float64, n), + means: make([]float64, p), + scales: make([]float64, p), + } + inv := 1 / float64(n) + for j := range p { + mean := 0.0 + if fx != nil { + for i := range n { + mean += fx[i*p+j] + } + } else { + for i := range n { + mean += x.FloatAt(i*p + j) + } + } + mean *= inv + variance := 0.0 + if fx != nil { + for i := range n { + d := fx[i*p+j] - mean + variance += d * d + } + } else { + for i := range n { + d := x.FloatAt(i*p+j) - mean + variance += d * d + } + } + scale := math.Sqrt(variance * inv) + if scale == 0 { + return nil, base.Errf("%s: design column %d is constant and cannot be standardised", name, j) + } + col := make([]float64, n) + if fx != nil { + for i := range n { + col[i] = (fx[i*p+j] - mean) / scale + } + } else { + for i := range n { + col[i] = (x.FloatAt(i*p+j) - mean) / scale + } + } + std.xs[j] = col + std.means[j] = mean + std.scales[j] = scale + } + yMean := 0.0 + // The walk is bounded by the response's own element count: a rebased + // view's payload may run past its visible elements. + if fy != nil { + for _, v := range fy[:n] { + yMean += v + } + } else { + for i := range n { + yMean += y.FloatAt(i) + } + } + yMean *= inv + std.yMean = yMean + for i := range n { + var yv float64 + if fy != nil { + yv = fy[i] + } else { + yv = y.FloatAt(i) + } + std.yc[i] = yv - yMean + } + for j := range p { + s := 0.0 + for i := range n { + s += std.xs[j][i] * std.xs[j][i] + } + std.zs[j] = s * inv + } + return std, nil +} + +// elasticNetFit runs the coordinate descent on the standardised +// system. start, when given, is the warm start the path hands down; +// a nil start begins at zero, the exact answer at the top of the +// grid. The running buffers live in ws and are reused across calls: +// the path fits one lambda after another into the same workspace, and +// a single fit simply owns one it just allocated. Each cycle sweeps +// the coordinates in order, moving each slope to the exact minimiser +// of the objective along its own coordinate and folding the movement +// into the running residual, until a full cycle moves nothing by more +// than the tolerance. The per-coordinate minimiser is the +// soft-thresholded least-squares coordinate over the denominator the +// elastic net puts there: pure lasso thresholding when alpha is 1, +// plain ridge division when alpha is 0. +func elasticNetFit(std *lassoStandardisation, lambda, alpha float64, start []float64, ws *lassoWorkspace) (beta []float64, iterations int, converged bool) { + n := len(std.yc) + p := len(std.xs) + inv := 1 / float64(n) + if start == nil { + beta = append(ws.beta[:0], make([]float64, p)...) + } else { + beta = append(ws.beta[:0], start...) + } + ws.beta = beta + // The running residual r = yc − Xβ, kept current by folding each + // coordinate's movement in as it happens: the sweep then reads the + // partial residual every coordinate needs without a full product + // per coordinate. + r := append(ws.r[:0], std.yc...) + ws.r = r + if start != nil { + for j := range p { + bj := beta[j] + if bj == 0 { + continue + } + col := std.xs[j] + for i := range n { + r[i] -= col[i] * bj + } + } + } + for iter := 1; iter <= lassoMaxIter; iter++ { + worst := 0.0 + for j := range p { + col := std.xs[j] + zj := std.zs[j] + s := 0.0 + for i := range n { + s += col[i] * r[i] + } + rho := s*inv + zj*beta[j] + newBeta := softThreshold(rho, lambda*alpha) / (zj + lambda*(1-alpha)) + if newBeta != beta[j] { + d := newBeta - beta[j] + for i := range n { + r[i] -= col[i] * d + } + beta[j] = newBeta + if a := math.Abs(d); a > worst { + worst = a + } + } + } + if worst < lassoTolerance { + return beta, iter, true + } + } + return beta, lassoMaxIter, false +} + +// softThreshold is the proximal operator of the absolute value: the +// identity past the threshold and zero inside it. It is what makes +// the lasso exact on an orthonormal design, where the coordinates +// decouple and every slope is this function of its own projection. +func softThreshold(v, t float64) float64 { + if v > t { + return v - t + } + if v < -t { + return v + t + } + return 0 +} + +// lassoCoefficients lifts the standardised slopes back to the design as +// supplied: the slope of a standardised column scales by the column's +// standard deviation, and the intercept reabsorbs the column means, in +// the columns' own order, so the accumulation is the one the fit's +// reporting pass runs. +func lassoCoefficients(std *lassoStandardisation, betaStd []float64) (intercept float64, coefficients []float64) { + p := len(betaStd) + coefficients = make([]float64, p) + intercept = std.yMean + for j := range p { + coefficients[j] = betaStd[j] / std.scales[j] + intercept -= coefficients[j] * std.means[j] + } + return intercept, coefficients +} + +// lassoResult lifts the standardised slopes back to the design as +// supplied: the slope of a standardised column scales by the column's +// standard deviation, and the intercept reabsorbs the column means. +// Fitted and Residuals are then computed on the original design, the +// same final sweep WeightedLinearRegression runs. +func lassoResult(name string, n, p int, x, y *core.Array, fx, fy []float64, std *lassoStandardisation, betaStd []float64, lambda, alpha float64, iterations int, converged bool) (*ElasticNetResult, error) { + intercept, coefficients := lassoCoefficients(std, betaStd) + out := &ElasticNetResult{ + Intercept: intercept, + Coefficients: coefficients, + ColumnMeans: std.means, + ColumnScales: std.scales, + Lambda: lambda, + Alpha: alpha, + Iterations: iterations, + Converged: converged, + } + out.Fitted = make([]float64, n) + out.Residuals = make([]float64, n) + for i := range n { + f := out.Intercept + if fx != nil { + row := fx[i*p : i*p+p] + for j, xj := range row { + f += coefficients[j] * xj + } + } else { + for j := range p { + f += coefficients[j] * x.FloatAt(i*p+j) + } + } + var yv float64 + if fy != nil { + yv = fy[i] + } else { + yv = y.FloatAt(i) + } + out.Fitted[i] = f + out.Residuals[i] = yv - f + } + return out, nil +} diff --git a/stats/lasso_test.go b/stats/lasso_test.go new file mode 100644 index 0000000..878376b --- /dev/null +++ b/stats/lasso_test.go @@ -0,0 +1,581 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// orthonormalDesign builds an (n, p) design whose columns are +// orthogonal in the (1/n)·XᵀX = I sense: every column has mean zero +// and population standard deviation one. The Gram-Schmidt walk starts +// from the constant vector and throws that direction away again: it +// stands for the intercept the design's columns must be orthogonal +// to, and what is left is exactly the space the slopes live in. +func orthonormalDesign(t *testing.T, n, p int, seed int64) *core.Array { + t.Helper() + g := core.NewGenerator(seed) + cols := make([][]float64, 0, p+1) + for j := range p + 1 { + v := make([]float64, n) + if j == 0 { + for i := range n { + v[i] = 1 + } + } else { + for i := range n { + v[i] = g.NormalUnit() + } + } + for _, c := range cols { + dot := 0.0 + for i := range n { + dot += v[i] * c[i] + } + for i := range n { + v[i] -= dot / float64(n) * c[i] + } + } + norm := 0.0 + for i := range n { + norm += v[i] * v[i] + } + factor := math.Sqrt(float64(n) / norm) + for i := range n { + v[i] *= factor + } + cols = append(cols, v) + } + cols = cols[1:] + vals := make([]float64, 0, n*p) + for i := range n { + for j := range p { + vals = append(vals, cols[j][i]) + } + } + return mustFromFloats(t, vals, n, p) +} + +// TestLassoOrthonormalClosedForm pins the coordinate descent against +// the closed form the orthonormal case admits: the columns decouple, +// and every slope is the soft-thresholded projection of the centred +// response on its own column, sign(ρ)·max(|ρ| − λ, 0), independent of +// the other coordinates and of the iteration. +func TestLassoOrthonormalClosedForm(t *testing.T) { + n, p := 8, 4 + design := orthonormalDesign(t, n, p, 5) + // The precondition itself, checked rather than assumed: (1/n)XᵀX + // is the identity and the column means are zero. + for j := range p { + mean := 0.0 + for i := range n { + mean += design.FloatAt(i*p + j) + } + if math.Abs(mean/float64(n)) > 1e-12 { + t.Fatalf("column %d has mean %.3g, want 0", j, mean/float64(n)) + } + for k := j; k < p; k++ { + dot := 0.0 + for i := range n { + dot += design.FloatAt(i*p+j) * design.FloatAt(i*p+k) + } + want := 0.0 + if j == k { + want = 1 + } + if math.Abs(dot/float64(n)-want) > 1e-12 { + t.Fatalf("columns %d and %d have inner product %.12f, want %.12f", j, k, dot/float64(n), want) + } + } + } + y := mustFromFloats(t, []float64{3, -1, 4, 1, 5, -2, 0, 2}, n) + res, err := Lasso(design, y, 0.6) + if err != nil { + t.Fatalf("Lasso: %v", err) + } + // The closed form, evaluated on the standardised system the fit + // documents in its own result. + yMean := 0.0 + for i := range n { + yMean += y.FloatAt(i) + } + yMean /= float64(n) + for j := range p { + mean := res.ColumnMeans[j] + scale := res.ColumnScales[j] + rho := 0.0 + for i := range n { + z := (design.FloatAt(i*p+j) - mean) / scale + rho += z * (y.FloatAt(i) - yMean) + } + rho /= float64(n) + // The threshold is spelled inline, not borrowed from + // softThreshold: a helper on both sides of the comparison would + // be wrong together. sign(ρ)·max(|ρ|−λ, 0) is the closed form. + want := 0.0 + if rho > 0.6 { + want = rho - 0.6 + } else if rho < -0.6 { + want = rho + 0.6 + } + if math.Abs(res.Coefficients[j]-want) > 1e-9 { + t.Fatalf("coefficient %d = %.12f, want the soft threshold %.12f", j, res.Coefficients[j], want) + } + } + // The threshold itself, on literals: the shrinkage keeps the sign, + // shrinks by exactly λ and floors at zero. + for _, c := range []struct { + v, lambda, want float64 + }{ + {0.0, 0.5, 0.0}, + {0.25, 0.5, 0.0}, + {0.5, 0.5, 0.0}, + {1.5, 0.5, 1.0}, + {-0.5, 0.5, 0.0}, + {-1.5, 0.5, -1.0}, + {2.5, 0.5, 2.0}, + {-2.5, 0.5, -2.0}, + } { + if got := softThreshold(c.v, c.lambda); got != c.want { + t.Fatalf("softThreshold(%g, %g) = %g, want %g", c.v, c.lambda, got, c.want) + } + } + // The intercept keeps the unpenalised identity: the fit passes + // through the column means. + predicted := res.Intercept + for j := range p { + predicted += res.Coefficients[j] * res.ColumnMeans[j] + } + if math.Abs(predicted-yMean) > 1e-12 { + t.Fatalf("the intercept identity broke: %.12g at the column means, want ȳ = %.12g", predicted, yMean) + } + if !res.Converged { + t.Fatalf("the orthonormal fit did not report convergence") + } +} + +// TestElasticNetRidgeAgreement pins the alpha = 0 degeneration: with +// the L1 term gone the coordinate descent is solving the Tikhonov +// ridge system (ZᵀZ/n + λI)β = Zᵀ(y − ȳ)/n on the standardised +// design, and the fit must agree with the answer the shared LU solve +// delivers for that system, the same solve LinearRegression runs. +func TestElasticNetRidgeAgreement(t *testing.T) { + const ( + n = 40 + p = 5 + lambda = 0.7 + ) + g := core.NewGenerator(11) + x := make([]float64, 0, n*p) + yv := make([]float64, 0, n) + betaTrue := []float64{2, -1, 0.5, 0, 3} + for range n { + factor := g.NormalUnit() + row := make([]float64, p) + fitted := 1.0 + for j := range p { + row[j] = factor + 0.3*g.NormalUnit() + fitted += betaTrue[j] * row[j] + x = append(x, row[j]) + } + yv = append(yv, fitted+0.25*g.NormalUnit()) + } + design := mustFromFloats(t, x, n, p) + y := mustFromFloats(t, yv, n) + res, err := ElasticNet(design, y, lambda, 0) + if err != nil { + t.Fatalf("ElasticNet: %v", err) + } + // The standardised system, rebuilt in the test from the + // standardisation the result records. + yMean := 0.0 + for _, v := range yv { + yMean += v + } + yMean /= float64(n) + z := make([][]float64, p) + zty := make([]float64, p) + for j := range p { + z[j] = make([]float64, n) + for i := range n { + z[j][i] = (x[i*p+j] - res.ColumnMeans[j]) / res.ColumnScales[j] + zty[j] += z[j][i] * (yv[i] - yMean) + } + zty[j] /= float64(n) + } + normal := make([][]float64, p) + for j := range p { + normal[j] = make([]float64, p) + for k := range p { + for i := range n { + normal[j][k] += z[j][i] * z[k][i] + } + normal[j][k] /= float64(n) + } + normal[j][j] += lambda + } + solved, err := base.SolveSystem("ridgeReference", normal, [][]float64{zty}) + if err != nil { + t.Fatalf("the reference ridge solve failed: %v", err) + } + worst := 0.0 + for j := range p { + want := solved[0][j] / res.ColumnScales[j] + if math.Abs(res.Coefficients[j]-want) > 1e-7 { + t.Fatalf("ridge coefficient %d = %.10f, want the Tikhonov answer %.10f", j, res.Coefficients[j], want) + } + if d := math.Abs(res.Coefficients[j] - want); d > worst { + worst = d + } + } + wantIntercept := yMean + for j := range p { + wantIntercept -= solved[0][j] / res.ColumnScales[j] * res.ColumnMeans[j] + } + if math.Abs(res.Intercept-wantIntercept) > 1e-7 { + t.Fatalf("ridge intercept = %.10f, want %.10f", res.Intercept, wantIntercept) + } + t.Logf("alpha = 0 agrees with the LU ridge solve to %.3g", worst) +} + +// TestLassoPathSameGrid pins the shared path: the alpha = 0 ridge +// walks exactly the lambda grid the pure lasso defines, so the two +// fits are comparable coefficient for coefficient along it. +func TestLassoPathSameGrid(t *testing.T) { + g := core.NewGenerator(3) + const n, p = 30, 4 + x := make([]float64, 0, n*p) + yv := make([]float64, 0, n) + for range n { + fitted := 2.0 + for j := range p { + v := g.NormalUnit() + fitted += float64(p-j) * v + x = append(x, v) + } + yv = append(yv, fitted+0.5*g.NormalUnit()) + } + design := mustFromFloats(t, x, n, p) + y := mustFromFloats(t, yv, n) + lasso, err := LassoPath(design, y, 1) + if err != nil { + t.Fatalf("LassoPath: %v", err) + } + ridge, err := LassoPath(design, y, 0) + if err != nil { + t.Fatalf("LassoPath: %v", err) + } + if !slices.Equal(lasso.Lambdas, ridge.Lambdas) { + t.Fatalf("the ridge path left the lasso grid") + } + // Grid shape: descending, log-spaced, spanning the documented + // three decades. + for k := 1; k < len(lasso.Lambdas); k++ { + if lasso.Lambdas[k] >= lasso.Lambdas[k-1] { + t.Fatalf("the grid is not descending at %d: %g then %g", k, lasso.Lambdas[k-1], lasso.Lambdas[k]) + } + } + ratio := lasso.Lambdas[1] / lasso.Lambdas[0] + if math.Abs(ratio-math.Pow(lassoGridRatio, 1.0/float64(lassoGridSteps-1))) > 1e-12 { + t.Fatalf("the grid is not log-spaced: consecutive ratio %.12g", ratio) + } + if math.Abs(lasso.Lambdas[len(lasso.Lambdas)-1]/lasso.Lambdas[0]-lassoGridRatio) > 1e-9 { + t.Fatalf("the grid spans %.6g decades of ratio, want %.6g", + lasso.Lambdas[len(lasso.Lambdas)-1]/lasso.Lambdas[0], lassoGridRatio) + } +} + +// TestLassoSparseSupportRecovery recovers the support of a sparse +// true model from seeded generated data: along the path there is a +// lambda interval where the nonzero set is exactly the true one and +// the estimates sit close to the truth, while the top of the grid is +// exactly the all-zero answer it promises. +func TestLassoSparseSupportRecovery(t *testing.T) { + const n, p = 300, 15 + g := core.NewGenerator(7) + betaTrue := make([]float64, p) + betaTrue[2], betaTrue[7], betaTrue[11] = 1.5, -2.0, 0.8 + x := make([]float64, 0, n*p) + yv := make([]float64, 0, n) + for range n { + fitted := 3.0 + for j := range p { + v := g.NormalUnit() + fitted += betaTrue[j] * v + x = append(x, v) + } + yv = append(yv, fitted+0.5*g.NormalUnit()) + } + design := mustFromFloats(t, x, n, p) + y := mustFromFloats(t, yv, n) + path, err := LassoPath(design, y, 1) + if err != nil { + t.Fatalf("LassoPath: %v", err) + } + if !path.Converged { + t.Fatalf("the path did not converge within its budget") + } + // The top of the grid: every slope exactly zero, the documented + // meaning of lambdaMax. + for j := range p { + if path.Coefficients[0][j] != 0 { + t.Fatalf("coefficient %d = %g at the top of the grid, want exactly 0", j, path.Coefficients[0][j]) + } + } + // Somewhere along the path the support is the true one. + trueSet := []int{2, 7, 11} + recovered := false + for k := range path.Lambdas { + support := []int{} + for j := range p { + if path.Coefficients[k][j] != 0 { + support = append(support, j) + } + } + if !slices.Equal(support, trueSet) { + continue + } + recovered = true + closeEnough := true + for _, j := range trueSet { + if math.Abs(path.Coefficients[k][j]-betaTrue[j]) > 0.2 { + closeEnough = false + } + } + if closeEnough { + t.Logf("support recovered at lambda[%d] = %.4g, coefficients within %.3f of the truth", + k, path.Lambdas[k], 0.2) + break + } + recovered = false + } + if !recovered { + t.Fatalf("no lambda on the path recovered the true support {2, 7, 11}") + } +} + +// TestLassoPathWarmStart measures the warm start: the same grid fitted +// cold, one ElasticNet call per lambda from zero every time, must +// spend more coordinate cycles than the warm path, and by the bottom +// of the grid the difference is the whole point of the path. +func TestLassoPathWarmStart(t *testing.T) { + const n, p = 100, 10 + g := core.NewGenerator(13) + betaTrue := make([]float64, p) + betaTrue[1], betaTrue[4], betaTrue[8] = 1.2, -1.6, 0.9 + x := make([]float64, 0, n*p) + yv := make([]float64, 0, n) + for range n { + factor := g.NormalUnit() + fitted := 1.0 + for j := range p { + v := factor + 0.5*g.NormalUnit() + fitted += betaTrue[j] * v + x = append(x, v) + } + yv = append(yv, fitted+0.4*g.NormalUnit()) + } + design := mustFromFloats(t, x, n, p) + y := mustFromFloats(t, yv, n) + warm, err := LassoPath(design, y, 1) + if err != nil { + t.Fatalf("LassoPath: %v", err) + } + cold := make([]int, len(warm.Lambdas)) + totalWarm, totalCold := 0, 0 + for k, lambda := range warm.Lambdas { + fit, err := ElasticNet(design, y, lambda, 1) + if err != nil { + t.Fatalf("ElasticNet: %v", err) + } + cold[k] = fit.Iterations + totalWarm += warm.Iterations[k] + totalCold += cold[k] + if warm.Iterations[k] > cold[k] { + t.Fatalf("the warm start lost to the cold fit at lambda[%d]: %d cycles against %d", + k, warm.Iterations[k], cold[k]) + } + } + t.Logf("warm path %d cycles against cold %d; at the smallest lambda %d against %d", + totalWarm, totalCold, warm.Iterations[len(cold)-1], cold[len(cold)-1]) + if totalWarm >= totalCold { + t.Fatalf("the warm path spent %d cycles, the cold path only %d", totalWarm, totalCold) + } + if warm.Iterations[len(cold)-1] >= cold[len(cold)-1] { + t.Fatalf("at the smallest lambda the warm start spent %d cycles against the cold %d", + warm.Iterations[len(cold)-1], cold[len(cold)-1]) + } +} + +// TestLassoDuplicateColumnResolve pins the collinear resolution: two +// identical columns make the optimum non-unique, the effect free to +// sit anywhere on the face the pair spans. The coordinate descent +// settles on that face deterministically, so identical inputs give +// bit-identical coefficients, and the fit is the single-copy answer +// in everything a consumer can measure: the fitted values, the +// summed effect of the pair and the residual structure. +func TestLassoDuplicateColumnResolve(t *testing.T) { + const n = 60 + g := core.NewGenerator(17) + x0 := make([]float64, n) + x2 := make([]float64, n) + yv := make([]float64, n) + for i := range n { + x0[i] = g.NormalUnit() + x2[i] = g.NormalUnit() + yv[i] = 1 + 2*x0[i] + 0.5*x2[i] + 0.3*g.NormalUnit() + } + dupVals := make([]float64, 0, 3*n) + singleVals := make([]float64, 0, 2*n) + for i := range n { + dupVals = append(dupVals, x0[i], x0[i], x2[i]) + singleVals = append(singleVals, x0[i], x2[i]) + } + dup := mustFromFloats(t, dupVals, n, 3) + single := mustFromFloats(t, singleVals, n, 2) + y := mustFromFloats(t, yv, n) + res, err := ElasticNet(dup, y, 0.02, 1) + if err != nil { + t.Fatalf("ElasticNet: %v", err) + } + if !res.Converged { + t.Fatalf("the duplicated design did not converge") + } + for j := range 3 { + if math.IsNaN(res.Coefficients[j]) || math.IsInf(res.Coefficients[j], 0) { + t.Fatalf("coefficient %d diverged to %g", j, res.Coefficients[j]) + } + } + // Determinism: a second, identical fit lands on the same bits. + again, err := ElasticNet(dup, y, 0.02, 1) + if err != nil { + t.Fatalf("the repeated ElasticNet failed: %v", err) + } + if !slices.Equal(res.Coefficients, again.Coefficients) { + t.Fatalf("the duplicated design settled differently on a repeat: %v against %v", + res.Coefficients, again.Coefficients) + } + // The degenerate face carries the single-copy effect: the pair + // sums to it, and the third column agrees with its own fit. + ref, err := ElasticNet(single, y, 0.02, 1) + if err != nil { + t.Fatalf("ElasticNet on the single-copy design: %v", err) + } + if math.Abs(res.Coefficients[0]+res.Coefficients[1]-ref.Coefficients[0]) > 1e-6 { + t.Fatalf("the duplicate pair summed to %g, the single copy took %g", + res.Coefficients[0]+res.Coefficients[1], ref.Coefficients[0]) + } + if math.Abs(res.Coefficients[2]-ref.Coefficients[1]) > 1e-6 { + t.Fatalf("the independent column moved from %g to %g under the duplicate", + ref.Coefficients[1], res.Coefficients[2]) + } + for i := range n { + if math.Abs(res.Fitted[i]-ref.Fitted[i]) > 1e-6 { + t.Fatalf("the duplicate changed fitted value %d: %.10g against %.10g", i, res.Fitted[i], ref.Fitted[i]) + } + } + t.Logf("duplicate pair resolved deterministically as (%g, %g), the single-copy effect being %g", + res.Coefficients[0], res.Coefficients[1], ref.Coefficients[0]) +} + +// TestLassoRefitsSmoke walks the mixed alphas through one small fit +// each, so the whole alpha range shares one code path and none of it +// is only exercised by the pins above. +func TestLassoAlphaRangeSmoke(t *testing.T) { + design := orthonormalDesign(t, 12, 3, 21) + y := mustFromFloats(t, []float64{2, -1, 3, 0, 1, -2, 4, 1, 0, -1, 2, 3}, 12) + for _, alpha := range []float64{0, 0.25, 0.5, 0.75, 1} { + res, err := ElasticNet(design, y, 0.3, alpha) + if err != nil { + t.Fatalf("ElasticNet alpha %g: %v", alpha, err) + } + if !res.Converged { + t.Fatalf("ElasticNet alpha %g did not converge", alpha) + } + for j := range 3 { + if math.IsNaN(res.Coefficients[j]) { + t.Fatalf("ElasticNet alpha %g produced a NaN at %d", alpha, j) + } + } + } +} + +// TestLassoOnIntegerArrays exercises the widening accessor's fallback +// paths in the standardisation and the final sweep: an integer design +// and response reach the fit through FloatAt rather than a raw float +// payload, and the answer matches the widened floats exactly. +func TestLassoOnIntegerArrays(t *testing.T) { + design := mustFromInts(t, []int64{0, 0, 1, 1, 2, 0, 3, 1, 4, 0, 5, 1, 6, 0}, 7, 2) + y := mustFromInts(t, []int64{2, 4, 5, 7, 8, 10, 11}, 7) + res, err := Lasso(design, y, 1e-5) + if err != nil { + t.Fatalf("Lasso on integer input: %v", err) + } + if !res.Converged { + t.Fatalf("the integer-input fit did not converge") + } + if math.Abs(res.Coefficients[0]-1.5) > 1e-4 || math.Abs(res.Coefficients[1]-0.5) > 1e-4 { + t.Fatalf("the integer-input fit is (%.6f, %.6f), want (1.5, 0.5)", + res.Coefficients[0], res.Coefficients[1]) + } + widened := mustFromFloats(t, []float64{0, 0, 1, 1, 2, 0, 3, 1, 4, 0, 5, 1, 6, 0}, 7, 2) + yw := mustFromFloats(t, []float64{2, 4, 5, 7, 8, 10, 11}, 7) + reference, err := Lasso(widened, yw, 1e-5) + if err != nil { + t.Fatalf("Lasso on the widened floats: %v", err) + } + if !slices.Equal(res.Coefficients, reference.Coefficients) { + t.Fatalf("the integer input fit %v against the float input %v", + res.Coefficients, reference.Coefficients) + } +} + +// TestLassoInputValidation refuses every malformed input the fit +// cannot answer, naming the condition in each case. +func TestLassoInputValidation(t *testing.T) { + good := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 4, 2) + resp := mustFromFloats(t, []float64{1, 2, 3, 4}, 4) + constant := mustFromFloats(t, []float64{1, 1, 1, 1, 1, 2, 1, 3}, 4, 2) + nanY := mustFromFloats(t, []float64{1, 2, math.NaN(), 4}, 4) + if _, err := ElasticNet(mustFromFloats(t, []float64{1, 2, 3, 4}, 4), resp, 0.1, 1); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("a rank 1 design: got %v, want the rank refusal", err) + } + if _, err := ElasticNet(good, mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2), 0.1, 1); err == nil || !strings.Contains(err.Error(), "must be rank 1") { + t.Fatalf("a rank 2 response: got %v, want the rank refusal", err) + } + if _, err := ElasticNet(good, mustFromFloats(t, []float64{1, 2, 3}, 3), 0.1, 1); err == nil || !strings.Contains(err.Error(), "rows but the response") { + t.Fatalf("a row mismatch: got %v, want the length refusal", err) + } + if _, err := ElasticNet(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 4, 2), nanY, 0.1, 1); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("a non-finite response: got %v, want the non-finite refusal", err) + } + if _, err := ElasticNet(constant, resp, 0.1, 1); err == nil || !strings.Contains(err.Error(), "cannot be standardised") { + t.Fatalf("a constant design column: got %v, want the standardisation refusal", err) + } + if _, err := ElasticNet(good, resp, -0.5, 1); err == nil || !strings.Contains(err.Error(), "lambda") { + t.Fatalf("a negative lambda: got %v, want the lambda refusal", err) + } + if _, err := ElasticNet(good, resp, 0.1, 1.5); err == nil || !strings.Contains(err.Error(), "alpha") { + t.Fatalf("an alpha above 1: got %v, want the alpha refusal", err) + } + if _, err := ElasticNet(good, resp, 0.1, -0.1); err == nil || !strings.Contains(err.Error(), "alpha") { + t.Fatalf("a negative alpha: got %v, want the alpha refusal", err) + } + if _, err := ElasticNet(mustFromFloats(t, []float64{1, 2}, 1, 2), mustFromFloats(t, []float64{1}, 1), 0.1, 1); err == nil || !strings.Contains(err.Error(), "at least two observations") { + t.Fatalf("a single observation: got %v, want the observation floor refusal", err) + } + complexDesign := core.New(core.Complex, 4, 2) + if _, err := ElasticNet(complexDesign, resp, 0.1, 1); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("complex input: got %v, want the complex refusal", err) + } + if _, err := LassoPath(good, resp, 2); err == nil || !strings.Contains(err.Error(), "alpha") { + t.Fatalf("LassoPath, an alpha above 1: got %v, want the alpha refusal", err) + } +} diff --git a/stats/mixed.go b/stats/mixed.go new file mode 100644 index 0000000..7805575 --- /dev/null +++ b/stats/mixed.go @@ -0,0 +1,826 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The linear mixed model: a fixed-effect design shared by every +// observation plus a random-effect design whose coefficients vary by +// group, y = X·β + Z·b + ε with b_g ~ N(0, Σ) per group and +// ε ~ N(0, σ²I). The variance components (Σ and σ²) are estimated by +// residual maximum likelihood, the REML criterion that compares only +// the error contrasts and so does not let the fixed effects soak up +// degrees of freedom the variance estimate needs; the fixed effects +// follow by general least squares at the fitted components. +// +// The maximisation is the expectation-maximisation sweep over the +// random-effect posterior: every update is a conditional expectation, +// so the sweep climbs the log likelihood monotonically until the +// tolerance stops it, and the whole fit is deterministic for a given +// input. No generator enters: the starting point comes from the +// data's own ordinary least squares, and the sweep walks from there. + +// mixedMaxIterations caps the EM sweeps; a fit that has not settled by +// then is reported with Converged false rather than pretended to have +// converged. +const mixedMaxIterations = 500 + +// mixedTolerance is the convergence tolerance on the REML log +// likelihood: the sweeps stop once it moves by less than +// mixedTolerance scaled by 1 + |log likelihood|, past which no +// parameter the M step can produce moves the objective materially. +const mixedTolerance = 1e-10 + +// LinearMixedModelResult carries the fit of a linear mixed model. +type LinearMixedModelResult struct { + // Coefficients are the fitted fixed effects β̂, one per column of + // the fixed design, and StandardErrors their estimated standard + // deviations, the square roots of the diagonal of the GLS + // covariance (XᵀV⁻¹X)⁻¹ at the fitted variance components. + Coefficients []float64 + StandardErrors []float64 + // RandomEffects holds one coefficient vector per group, in the + // order GroupLabels names, each of the length of a row of the + // random design: the posterior means b̂_g at the fitted components. + RandomEffects [][]float64 + // GroupLabels are the distinct group labels in the order the fit + // met them, the order RandomEffects keeps. + GroupLabels []int + // RandomCovariance is the fitted between-group covariance Σ̂ of the + // random effects, row-major, and ResidualVariance the fitted σ̂². + RandomCovariance []float64 + ResidualVariance float64 + // LogLikelihood is the maximised REML log likelihood, the + // criterion the sweep climbed. + LogLikelihood float64 + // Fitted and Residuals align with the rows of the design: the + // conditional fit X·β̂ + Z·b̂ per row and its remainder. + Fitted []float64 + Residuals []float64 + // Iterations counts the EM sweeps taken; Converged reports whether + // the log likelihood settled under mixedTolerance before the + // budget ran out. + Iterations int + Converged bool +} + +// mixedGroup holds one group's fixed pieces, read once before the +// sweep: the row indices, both designs restricted to the group, the +// response restricted to it, and the random design's own scatter +// ZᵀZ, which no iteration moves. +type mixedGroup struct { + label int + rows []int + yg []float64 + xg []float64 + zg []float64 + ztz []float64 + xtz []float64 + ng int +} + +// LinearMixedModel fits the linear mixed model over the response y, +// the fixed design x (n rows, p columns, the intercept included by +// the caller as a constant column when one is wanted), the random +// design z (n rows, q columns) and the group label of every row. The +// labels need no particular order or contiguity: the fit meets them +// in row order and reports the distinct ones in GroupLabels in the +// order it met them. The random effects share one unstructured q×q +// covariance Σ across the groups. +// +// The fit runs expectation maximisation over the random-effect +// posterior to the documented tolerance, within mixedMaxIterations +// sweeps, and reports the REML log likelihood of the fitted state. A +// group whose marginal covariance fails to factor under the fitted +// components, or a fixed design that is singular under them, is an +// error naming the group or the design. +// +// Refuses complex input and any non-finite entry, a response shorter +// than two rows, a design of fewer than one column, a row count +// mismatch, more fixed coefficients than observations, and a +// negative or missing group label. +func LinearMixedModel(y, x, z *core.Array, groups []int) (*LinearMixedModelResult, error) { + const name = "LinearMixedModel" + if y.NDim() != 1 { + return nil, base.Errf("%s: the response must be rank 1, got shape %s", name, base.ShapeText(y.Shape())) + } + if x.NDim() != 2 || z.NDim() != 2 { + return nil, base.Errf("%s: both designs must be rank 2", name) + } + if y.Dtype() == core.Complex || x.Dtype() == core.Complex || z.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + n := y.Len() + p := x.Shape()[1] + q := z.Shape()[1] + if x.Shape()[0] != n || z.Shape()[0] != n { + return nil, base.Errf("%s: the response holds %d rows, the designs %d and %d", name, n, x.Shape()[0], z.Shape()[0]) + } + if n < 2 { + return nil, base.Errf("%s: at least two observations are needed, got %d", name, n) + } + if p < 1 || q < 1 { + return nil, base.Errf("%s: both designs need at least one column, got %d and %d", name, p, q) + } + if n <= p { + return nil, base.Errf("%s: %d observations cannot carry %d fixed coefficients", name, n, p) + } + if len(groups) != n { + return nil, base.Errf("%s: %d group labels for %d observations", name, len(groups), n) + } + if err := checkFinite(name, "the response", y); err != nil { + return nil, err + } + if err := checkFinite(name, "the fixed design", x); err != nil { + return nil, err + } + if err := checkFinite(name, "the random design", z); err != nil { + return nil, err + } + // The response and the designs are read once into plain slices: + // every sweep below indexes them flat. + yVals := make([]float64, n) + if fs := rawFloats(y); fs != nil { + copy(yVals, fs) + } else { + for i := range n { + yVals[i] = y.FloatAt(i) + } + } + xVals := make([]float64, n*p) + if fs := rawFloats(x); fs != nil { + copy(xVals, fs[:n*p]) + } else { + for i := range xVals { + xVals[i] = x.FloatAt(i) + } + } + zVals := make([]float64, n*q) + if fs := rawFloats(z); fs != nil { + copy(zVals, fs[:n*q]) + } else { + for i := range zVals { + zVals[i] = z.FloatAt(i) + } + } + // The groups are canonicalised by first appearance: the labels + // stay the caller's own, the fit only needs them distinct and + // stable. + labelIndex := make(map[int]int) + labels := make([]int, 0, 8) + member := make([][]int, 0, 8) + for i, label := range groups { + if label < 0 { + return nil, base.Errf("%s: group label %d is negative", name, label) + } + gi, ok := labelIndex[label] + if !ok { + gi = len(labels) + labelIndex[label] = gi + labels = append(labels, label) + member = append(member, nil) + } + member[gi] = append(member[gi], i) + } + gs := make([]mixedGroup, len(labels)) + for gi, rows := range member { + g := &gs[gi] + g.label = labels[gi] + g.rows = rows + g.ng = len(rows) + g.yg = make([]float64, g.ng) + g.xg = make([]float64, g.ng*p) + g.zg = make([]float64, g.ng*q) + g.ztz = make([]float64, q*q) + g.xtz = make([]float64, p*q) + for li, row := range rows { + g.yg[li] = yVals[row] + copy(g.xg[li*p:(li+1)*p], xVals[row*p:(row+1)*p]) + copy(g.zg[li*q:(li+1)*q], zVals[row*q:(row+1)*q]) + } + for a := range q { + for b := range q { + s := 0.0 + for li := range g.ng { + s += g.zg[li*q+a] * g.zg[li*q+b] + } + g.ztz[a*q+b] = s + } + } + for j := range p { + for b := range q { + s := 0.0 + for li := range g.ng { + s += g.xg[li*p+j] * g.zg[li*q+b] + } + g.xtz[j*q+b] = s + } + } + } + // The starting point: ordinary least squares on the fixed design, + // the residual variance around it, and a between-group scatter of + // its residuals floored above zero so the first sweep can see the + // random effects at all. + beta, err := mixedOLSStart(name, xVals, yVals, n, p) + if err != nil { + return nil, err + } + residual := make([]float64, n) + rss := 0.0 + totalSquare := 0.0 + for i := range n { + fit := 0.0 + for j := range p { + fit += xVals[i*p+j] * beta[j] + } + residual[i] = yVals[i] - fit + rss += residual[i] * residual[i] + totalSquare += yVals[i] * yVals[i] + } + scale := 1 + totalSquare/float64(n) + sigma2 := math.Max(rss/float64(n), 1e-12*scale) + between := 0.0 + if len(gs) > 1 { + overall := 0.0 + means := make([]float64, len(gs)) + for gi, g := range gs { + s := 0.0 + for _, row := range g.rows { + s += residual[row] + } + means[gi] = s / float64(g.ng) + overall += means[gi] + } + overall /= float64(len(gs)) + for _, m := range means { + d := m - overall + between += d * d + } + between /= float64(len(gs) - 1) + } + comp := math.Max(between, 0.1*sigma2) + sigma := make([]float64, q*q) + for a := range q { + sigma[a*q+a] = comp + } + // The sweep. The per-group buffers a pass touches are allocated + // once here and refilled in place; the fixed-effect solve keeps + // its own small p×p workspace, reallocated each pass. + chol := make([][][]float64, len(gs)) + scratch := make([]mixedSweep, len(gs)) + for gi := range gs { + ng := gs[gi].ng + chol[gi] = make([][]float64, ng) + for i := range ng { + chol[gi][i] = make([]float64, ng) + } + scratch[gi] = newMixedSweep(ng, p, q) + } + // The fixed scatter XᵀX, which no iteration moves: the REML M + // step's expectation of the squared residual carries its trace + // against the fixed effects' posterior covariance. + xtx := make([]float64, p*p) + for i := range n { + for j := range p { + xj := xVals[i*p+j] + for k := range p { + xtx[j*p+k] += xj * xVals[i*p+k] + } + } + } + xtvix := make([]float64, p*p) + xtviy := make([]float64, p) + sigmaNew := make([]float64, q*q) + logLik := math.Inf(-1) + converged := false + iterations := mixedMaxIterations + for iter := 1; iter <= mixedMaxIterations; iter++ { + pieces, perr := mixedPass(gs, chol, scratch, xtx, beta, sigma, sigma2, xtvix, xtviy, sigmaNew) + if perr != nil { + return nil, base.Errf("%s: %w", name, perr) + } + next := -0.5 * (pieces.logDetV + pieces.quad + pieces.logDetX + float64(n-p)*math.Log(2*math.Pi)) + if math.IsNaN(next) { + return nil, base.Errf("%s: the fit left the finite domain at iteration %d", name, iter) + } + move := next - logLik + logLik = next + if iter > 1 && math.Abs(move) <= mixedTolerance*(1+math.Abs(logLik)) { + converged = true + iterations = iter + break + } + // The M step: the averaged posterior scatter, the expected + // residual variance with its posterior-trace correction, and + // the GLS fixed effects at this pass's components. + copy(sigma, sigmaNew) + for a := range q { + for b := a + 1; b < q; b++ { + sigma[b*q+a] = sigma[a*q+b] + } + } + sigma2 = math.Max((pieces.sse+pieces.trace)/float64(n), 1e-12*scale) + // The divergence fuse: some designs give the REML criterion no + // interior optimum, the classical case being a random design + // that spans the fixed one under an unstructured covariance, + // where the surface climbs a ridge of singular Σ without a + // summit. The fit refuses with the condition named instead of + // publishing the climb. + if sigma2 > 1e12*scale { + return nil, base.Errf("%s: the variance components diverged past the data's scale; this design gives the REML criterion no interior optimum", name) + } + for _, v := range sigma { + if math.Abs(v) > 1e12*scale { + return nil, base.Errf("%s: the variance components diverged past the data's scale; this design gives the REML criterion no interior optimum", name) + } + } + if err := lltSolveSquare(name, xtvix, xtviy, p); err != nil { + return nil, base.Errf("%s: the fixed design is singular under the fitted covariance (%w)", name, err) + } + copy(beta, xtviy) + } + // The reporting pass at the fitted parameters: the same sweep once + // more at fixed state, then the quantities the caller holds. The + // sweep is deterministic, so its totals are the fitted state's + // own. + pieces, perr := mixedPass(gs, chol, scratch, xtx, beta, sigma, sigma2, xtvix, xtviy, sigmaNew) + if perr != nil { + return nil, base.Errf("%s: %w", name, perr) + } + logLik = -0.5*(pieces.logDetV+pieces.quad+pieces.logDetX) - 0.5*float64(n-p)*math.Log(2*math.Pi) + // The standard errors: the square roots of the diagonal of the GLS + // covariance, from the inverse of the same matrix the log + // determinant came from. + diag, serr := lltInverseDiagonal(name, xtvix, p) + if serr != nil { + return nil, base.Errf("%s: %w", name, serr) + } + se := make([]float64, p) + for i := range p { + se[i] = math.Sqrt(diag[i]) + } + fitted := make([]float64, n) + resids := make([]float64, n) + for gi, g := range gs { + bhat := scratch[gi].bhat + for li, row := range g.rows { + fit := 0.0 + for j := range p { + fit += xVals[row*p+j] * beta[j] + } + for a := range q { + fit += g.zg[li*q+a] * bhat[a] + } + fitted[row] = fit + resids[row] = yVals[row] - fit + } + } + effects := make([][]float64, len(gs)) + for gi := range gs { + effects[gi] = append([]float64(nil), scratch[gi].bhat...) + } + return &LinearMixedModelResult{ + Coefficients: append([]float64(nil), beta...), + StandardErrors: se, + RandomEffects: effects, + GroupLabels: labels, + RandomCovariance: append([]float64(nil), sigma...), + ResidualVariance: sigma2, + LogLikelihood: logLik, + Fitted: fitted, + Residuals: resids, + Iterations: iterations, + Converged: converged, + }, nil +} + +// mixedSweep holds one group's reusable pass buffers: the marginal +// covariance and its factor's solves against the response and both +// designs, the conditional residual, and the posterior mean and +// scatter of the group's random effects. +type mixedSweep struct { + u []float64 // V⁻¹ r + vy []float64 // V⁻¹ y + vz []float64 // V⁻¹ Z, ng×q + vx []float64 // V⁻¹ X, ng×p + slope []float64 // −Σ Zᵀ (V⁻¹ X), the posterior mean's slope in β, q×p + bhat []float64 // Σ Zᵀ u + post []float64 // P + b̂b̂ᵀ, the conditional scatter + v []float64 // the marginal covariance, ng×ng + resid []float64 // y − Xβ per row + mzx []float64 // Σ Zᵀ (V⁻¹ Z), q×q + sc []float64 // slope·C, refilled by the REML correction, q×p + bq []float64 // sc·slopeᵀ, the correction's scatter, q×q + colvec []float64 // one design column, refilled per solve +} + +// newMixedSweep allocates one group's buffers. +func newMixedSweep(ng, p, q int) mixedSweep { + return mixedSweep{ + u: make([]float64, ng), + vy: make([]float64, ng), + vz: make([]float64, ng*q), + vx: make([]float64, ng*p), + slope: make([]float64, q*p), + bhat: make([]float64, q), + post: make([]float64, q*q), + v: make([]float64, ng*ng), + resid: make([]float64, ng), + mzx: make([]float64, q*q), + sc: make([]float64, q*p), + bq: make([]float64, q*q), + colvec: make([]float64, ng), + } +} + +// mixedPassTotals gathers what one pass accumulates across the +// groups. +type mixedPassTotals struct { + logDetV float64 // Σ log|V_g|, the marginal covariances' own logarithms + quad float64 // Σ r_gᵀ V_g⁻¹ r_g + sse float64 // Σ ‖r_g − Z_g b̂_g‖² + trace float64 // Σ tr(Z_gᵀZ_g·P_g) + logDetX float64 // log|XᵀV⁻¹X| +} + +// mixedPass runs one expectation pass over the groups at the given +// parameters: it factors every group's marginal covariance, solves it +// against the response and both designs, forms the posterior means +// and scatters, and accumulates the totals the log likelihood and the +// M step read. The M step's expectations are taken under the flat +// prior on the fixed effects, so the fixed effects' own posterior +// covariance (XᵀV⁻¹X)⁻¹ joins both updates through the posterior +// mean's slope in β; that is the step that makes the fixed point the +// REML optimum rather than the ML one. The accumulators xtvix, xtviy +// and sigmaNew are overwritten; the posterior means stay in the +// group's scratch for the reporting pass to collect. +func mixedPass(gs []mixedGroup, chol [][][]float64, scratch []mixedSweep, xtx []float64, + beta []float64, sigma []float64, sigma2 float64, xtvix, xtviy, sigmaNew []float64) (*mixedPassTotals, error) { + const name = "LinearMixedModel" + p := len(beta) + q := len(scratch[0].bhat) + clear(xtvix) + clear(xtviy) + clear(sigmaNew) + totals := &mixedPassTotals{} + for gi := range gs { + g := &gs[gi] + s := &scratch[gi] + ng := g.ng + // The marginal covariance V = Z Σ Zᵀ + σ²I, both triangles + // written so the factorisation's symmetry check sees a mirror + // pair with identical bits. + for i := range ng { + for j := range i + 1 { + total := 0.0 + for a := range q { + za := g.zg[i*q+a] + if za != 0 { + for b := range q { + total += za * sigma[a*q+b] * g.zg[j*q+b] + } + } + } + if i == j { + total += sigma2 + } + s.v[i*ng+j] = total + s.v[j*ng+i] = total + } + } + if err := mvnCholeskyFlatInto(name, s.v, chol[gi], ng); err != nil { + return nil, base.Errf("group %d's marginal covariance failed to factor (%w)", g.label, err) + } + l := chol[gi] + for i := range ng { + // The factor's determinant answers |V_g| = |L|², so the + // diagonal logs enter doubled: the criterion the sweep reads + // below carries log|V_g| itself, not half of it. + totals.logDetV += 2 * math.Log(l[i][i]) + } + // The residual against the fixed part alone, then u = V⁻¹ r. + for li := range ng { + r := 0.0 + for j := range p { + r += g.xg[li*p+j] * beta[j] + } + s.resid[li] = g.yg[li] - r + } + copy(s.u, s.resid) + lltSolveInPlace(l, s.u) + totals.quad += dot(s.resid, s.u) + // V⁻¹ Z and V⁻¹ X, column by column through the same factor. + for a := range q { + for li := range ng { + s.colvec[li] = g.zg[li*q+a] + } + lltSolveInPlace(l, s.colvec) + copy(s.vz[a*ng:a*ng+ng], s.colvec) + } + for j := range p { + for li := range ng { + s.colvec[li] = g.xg[li*p+j] + } + lltSolveInPlace(l, s.colvec) + copy(s.vx[j*ng:j*ng+ng], s.colvec) + } + // V⁻¹ y, free of charge out of the solves already run: + // V⁻¹ y = V⁻¹ r + (V⁻¹ X)·β. + for li := range ng { + total := s.u[li] + for j := range p { + total += s.vx[j*ng+li] * beta[j] + } + s.vy[li] = total + } + // The posterior mean b̂ = Σ Zᵀ u. + clear(s.bhat) + for a := range q { + total := 0.0 + for li := range ng { + zt := 0.0 + for b := range q { + zt += sigma[a*q+b] * g.zg[li*q+b] + } + total += zt * s.u[li] + } + s.bhat[a] = total + } + // mzx = Σ Zᵀ (V⁻¹ Z), then the posterior covariance + // P = Σ − mzx·Σ on its upper triangle, joined by the mean's + // outer product into the scatter the M step averages. + clear(s.mzx) + for a := range q { + for b := range q { + total := 0.0 + for c := range q { + factor := sigma[a*q+c] + if factor != 0 { + for li := range ng { + total += factor * g.zg[li*q+c] * s.vz[li*q+b] + } + } + } + s.mzx[a*q+b] = total + } + } + clear(s.post) + for a := range q { + for b := a; b < q; b++ { + pv := sigma[a*q+b] + for c := range q { + pv -= s.mzx[a*q+c] * sigma[c*q+b] + } + s.post[a*q+b] = pv + s.bhat[a]*s.bhat[b] + } + } + for a := range q { + for b := a; b < q; b++ { + sigmaNew[a*q+b] += s.post[a*q+b] / float64(len(gs)) + } + } + // The conditional residual and its posterior-trace correction. + for li := range ng { + r := s.resid[li] + for a := range q { + r -= g.zg[li*q+a] * s.bhat[a] + } + s.resid[li] = r + totals.sse += r * r + } + for a := range q { + for b := a; b < q; b++ { + pv := sigma[a*q+b] + for c := range q { + pv -= s.mzx[a*q+c] * sigma[c*q+b] + } + totals.trace += g.ztz[a*q+b] * pv + if b != a { + totals.trace += g.ztz[b*q+a] * pv + } + } + } + // XᵀV⁻¹X and XᵀV⁻¹y accumulate across the groups. The + // information is symmetric by construction but its two rounding + // paths differ, so only the upper triangle is accumulated and + // mirrored: the factorisation's relative symmetry check would + // otherwise condemn a matrix whose one side rounded to an exact + // zero. + for j := range p { + for k := j; k < p; k++ { + total := 0.0 + for li := range ng { + total += g.xg[li*p+k] * s.vx[j*ng+li] + } + xtvix[j*p+k] += total + if k != j { + xtvix[k*p+j] += total + } + } + total := 0.0 + for li := range ng { + total += g.xg[li*p+j] * s.vy[li] + } + xtviy[j] += total + } + // The posterior mean's slope in β: E[b_g|β, y] moves as + // b̂ − (Σ ZᵀV⁻¹X)(β − β̂), and the slope carries the fixed + // effects' uncertainty into the M step's expectations. + for a := range q { + for j := range p { + total := 0.0 + for b := range q { + zt := sigma[a*q+b] + if zt != 0 { + for li := range ng { + total += zt * g.zg[li*q+b] * s.vx[j*ng+li] + } + } + } + s.slope[a*p+j] = -total + } + } + } + // The fixed information through its own Cholesky: the log + // determinant the REML criterion reads, and the posterior + // covariance C = (XᵀV⁻¹X)⁻¹ the corrections below walk through. + // A non-positive pivot means the design has collapsed under the + // fitted covariance, which is an error, not a fit. + cholX := make([][]float64, p) + for i := range p { + cholX[i] = make([]float64, p) + } + if err := mvnCholeskyFlatInto(name, xtvix, cholX, p); err != nil { + return nil, base.Errf("the fixed design is singular under the fitted covariance (%w)", err) + } + logDetX := 0.0 + for i := range p { + logDetX += math.Log(cholX[i][i]) + } + totals.logDetX = 2 * logDetX + inverse := make([]float64, p*p) + column := make([]float64, p) + for j := range p { + clear(column) + column[j] = 1 + lltSolveInPlace(cholX, column) + for i := range p { + inverse[i*p+j] = column[i] + } + } + // The flat-prior corrections. The expected squared residual gains + // the trace of XᵀX against C and, per group, twice the cross term + // of XᵀZ against C·slopeᵀ; the expected random-effect scatter + // gains slope·C·slopeᵀ. With no random effect the two terms leave + // σ² = RSS/(n−p), the REML answer, as the fixed point. + traceXX := 0.0 + for a := range p { + for b := range p { + traceXX += xtx[a*p+b] * inverse[a*p+b] + } + } + crossTotal := 0.0 + for gi := range gs { + g := &gs[gi] + s := &scratch[gi] + for a := range q { + for c := range p { + total := 0.0 + for j := range p { + total += s.slope[a*p+j] * inverse[j*p+c] + } + s.sc[a*p+c] = total + } + } + clear(s.bq) + for a := range q { + for b := a; b < q; b++ { + total := 0.0 + for c := range p { + total += s.sc[a*p+c] * s.slope[b*p+c] + } + s.bq[a*q+b] = total + } + } + for a := range q { + for b := a; b < q; b++ { + sigmaNew[a*q+b] += s.bq[a*q+b] / float64(len(gs)) + totals.trace += g.ztz[a*q+b] * s.bq[a*q+b] + if b != a { + totals.trace += g.ztz[b*q+a] * s.bq[a*q+b] + } + } + } + for j := range p { + for b := range q { + crossTotal += g.xtz[j*q+b] * s.sc[b*p+j] + } + } + } + totals.sse += traceXX + 2*crossTotal + return totals, nil +} + +// dot returns the plain dot product of two equal-length slices. +func dot(a, b []float64) float64 { + total := 0.0 + for i, v := range a { + total += v * b[i] + } + return total +} + +// lltSolveInPlace solves L·Lᵀ x = x in place: the forward sweep +// leaves L's solution in x, the backward sweep reads the entries the +// descending order has already updated. +func lltSolveInPlace(l [][]float64, x []float64) { + n := len(x) + for i := range n { + total := x[i] + for j := range i { + total -= l[i][j] * x[j] + } + x[i] = total / l[i][i] + } + for i := n - 1; i >= 0; i-- { + total := x[i] + for j := i + 1; j < n; j++ { + total -= l[j][i] * x[j] + } + x[i] = total / l[i][i] + } +} + +// lltSolveSquare solves the symmetric positive definite system a·x = +// b in place, factoring a through the house Cholesky and leaving b +// holding x. +func lltSolveSquare(name string, a, b []float64, d int) error { + l := make([][]float64, d) + for i := range d { + l[i] = make([]float64, d) + } + if err := mvnCholeskyFlatInto(name, a, l, d); err != nil { + return err + } + lltSolveInPlace(l, b) + return nil +} + +// lltInverseDiagonal returns the diagonal of the inverse of the +// symmetric positive definite a, one solve against each unit vector. +func lltInverseDiagonal(name string, a []float64, d int) ([]float64, error) { + l := make([][]float64, d) + for i := range d { + l[i] = make([]float64, d) + } + if err := mvnCholeskyFlatInto(name, a, l, d); err != nil { + return nil, err + } + diag := make([]float64, d) + column := make([]float64, d) + for i := range d { + clear(column) + column[i] = 1 + lltSolveInPlace(l, column) + diag[i] = column[i] + } + return diag, nil +} + +// mixedOLSStart returns the ordinary least squares coefficients of +// the fixed design, the fixed effects' starting point. +func mixedOLSStart(name string, xVals, yVals []float64, n, p int) ([]float64, error) { + xtx := make([]float64, p*p) + xty := make([]float64, p) + for i := range n { + for j := range p { + xj := xVals[i*p+j] + xty[j] += xj * yVals[i] + for k := range p { + xtx[j*p+k] += xj * xVals[i*p+k] + } + } + } + rhs := make([][]float64, 1) + rhs[0] = xty + solved, err := base.SolveSystem(name, toRows(xtx, p), rhs) + if err != nil { + return nil, base.Errf("%s: the fixed design is singular, the fixed effects cannot be initialised (%w)", name, err) + } + return append([]float64(nil), solved[0]...), nil +} + +// toRows views a flat row-major matrix as rows for the shared solve. +func toRows(flat []float64, d int) [][]float64 { + rows := make([][]float64, d) + for i := range d { + rows[i] = flat[i*d : i*d+d] + } + return rows +} diff --git a/stats/mixed_test.go b/stats/mixed_test.go new file mode 100644 index 0000000..afe01b0 --- /dev/null +++ b/stats/mixed_test.go @@ -0,0 +1,484 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The mixed model against referents. The balanced one-way random +// effects model has a closed-form REML answer, the analysis of +// variance estimators, so the sweep's optimum is compared against +// figures computed from the raw data rather than quoted; the rest +// pins the recovery of known effects, the refusal surface and the +// determinism of the whole pipeline. + +// mixedNoise is a deterministic stand-in for measurement noise: a +// bounded, aperiodic wiggle no sweep can mistake for structure. +func mixedNoise(i int) float64 { + return 0.3*math.Sin(7.3*float64(i)+1.1)*math.Cos(2.1*float64(i)) + + 0.1*math.Sin(0.7*float64(i)) +} + +// mixedJitter is deterministic white jitter on [−1, 1): the xorshift +// finaliser of the house generator's mixing constants, run on the row +// index. Unlike the smooth wiggle above it cannot be absorbed by a +// within-group linear span, which is what the random-slope fit needs +// its residual scale to be. +func mixedJitter(i int) float64 { + z := uint64(i)*2685821657736338717 + 1 + z ^= z >> 13 + z ^= z << 7 + z ^= z >> 17 + return float64(z>>11)/(1<<52)*2 - 1 +} + +func mixedVec(t *testing.T, vals []float64, shape ...int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, shape...) + if err != nil { + t.Fatal(err) + } + return a +} + +func TestMixedModelBalancedOneWayANOVAReferent(t *testing.T) { + // The balanced one-way random effects model: y_ig = μ + b_g + ε_ig + // with k observations in each of m groups. REML's optimum is the + // ANOVA answer: σ̂²_e = MSW and σ̂²_b = (MSB − MSW)/k, and the + // intercept's GLS variance is (k·σ²_b + σ²_e)/(mk). + const ( + mGroups = 8 + k = 6 + ) + y := make([]float64, 0, mGroups*k) + for g := range mGroups { + effect := 2.0 * math.Sin(1.7*float64(g)+0.4) // the drawn b_g + for i := range k { + y = append(y, 5+effect+mixedNoise(g*k+i)) + } + } + groups := make([]int, mGroups*k) + for g := range mGroups { + for i := range k { + groups[g*k+i] = g + } + } + res, err := LinearMixedModel( + mixedVec(t, y, len(y)), + mixedVec(t, ones(len(y)), len(y), 1), + mixedVec(t, ones(len(y)), len(y), 1), + groups) + if err != nil { + t.Fatal(err) + } + if !res.Converged { + t.Fatalf("the balanced fit did not converge (%d iterations)", res.Iterations) + } + // MSW and MSB from the raw data. + grand := 0.0 + groupMeans := make([]float64, mGroups) + for g := range mGroups { + s := 0.0 + for i := range k { + s += y[g*k+i] + } + groupMeans[g] = s / float64(k) + grand += groupMeans[g] + } + grand /= float64(mGroups) + msw := 0.0 + for g := range mGroups { + for i := range k { + d := y[g*k+i] - groupMeans[g] + msw += d * d + } + } + msw /= float64(mGroups * (k - 1)) + msb := 0.0 + for g := range mGroups { + d := groupMeans[g] - grand + msb += d * d + } + msb *= float64(k) / float64(mGroups-1) + wantWithin := math.Max((msb-msw)/float64(k), 0) + if math.Abs(res.ResidualVariance-msw) > 0.02*msw { + t.Fatalf("σ̂²_e = %.6f, want the MSW %.6f", res.ResidualVariance, msw) + } + if math.Abs(res.RandomCovariance[0]-wantWithin) > 0.05*wantWithin { + t.Fatalf("σ̂²_b = %.6f, want the ANOVA answer %.6f", res.RandomCovariance[0], wantWithin) + } + if math.Abs(res.Coefficients[0]-grand) > 1e-6 { + t.Fatalf("μ̂ = %.8f, want the grand mean %.8f", res.Coefficients[0], grand) + } + wantVar := (float64(k)*res.RandomCovariance[0] + res.ResidualVariance) / float64(mGroups*k) + if se := res.StandardErrors[0]; math.Abs(se*se-wantVar) > 1e-9*math.Max(1, wantVar) { + t.Fatalf("SE² = %.10f, want the GLS variance %.10f", se*se, wantVar) + } + // The conditional fitted values reproduce the group means plus the + // shrinkage the model applies; the residuals must complement them + // to the response. + for i := range len(y) { + if math.Abs(res.Fitted[i]+res.Residuals[i]-y[i]) > 1e-9 { + t.Fatalf("row %d: fitted + residuals = %g, want %g", i, res.Fitted[i]+res.Residuals[i], y[i]) + } + } +} + +// ones returns n constant 1 values, the intercept column. +func ones(n int) []float64 { + out := make([]float64, n) + for i := range out { + out[i] = 1 + } + return out +} + +func TestMixedModelRandomSlopeRecovery(t *testing.T) { + // Fixed effects of 1 and 2 with a per-group random slope, built by + // hand so the truth is known exactly. The random design carries + // the covariate alone: with the intercept column beside it the + // random span covers the fixed design and the REML surface loses + // its interior optimum to a ridge of singular covariance (pinned + // by the divergence test below). The slope values are centred, so + // the fixed part of the truth is exactly (1, 2). + slopeValues := []float64{0.9, -1.1, 1.9, -0.3, -1.8, 0.4} + groups := make([]int, 0, 36) + xs := make([]float64, 0, 36) + y := make([]float64, 0, 36) + jitter := make([]float64, 0, 36) + for g := range slopeValues { + for _, x := range []float64{-1, -0.6, -0.2, 0.2, 0.6, 1} { + groups = append(groups, g) + xs = append(xs, x) + jitter = append(jitter, mixedJitter(len(xs)-1)) + } + } + // The jitter is centred on its own sample: the fixed part of the + // truth must stay exactly (1, 2), and a noise vector with a mean + // would tilt the intercept instead of testing the recovery. + mean := 0.0 + for _, j := range jitter { + mean += j + } + mean /= float64(len(jitter)) + for g, s := range slopeValues { + for k := range 6 { + i := g*6 + k + x := xs[i] + y = append(y, 1+2*x+0.5*s*x+0.1*(jitter[i]-mean)) + } + } + design := make([]float64, 0, 2*len(xs)) + for _, x := range xs { + design = append(design, 1, x) + } + res, err := LinearMixedModel( + mixedVec(t, y, len(y)), + mixedVec(t, design, len(xs), 2), + mixedVec(t, xs, len(xs), 1), + groups) + if err != nil { + t.Fatal(err) + } + if !res.Converged { + t.Fatalf("the slope fit did not converge (%d iterations)", res.Iterations) + } + if math.Abs(res.Coefficients[0]-1) > 0.05 || math.Abs(res.Coefficients[1]-2) > 0.05 { + t.Fatalf("β̂ = (%.4f, %.4f), want (1, 2)", res.Coefficients[0], res.Coefficients[1]) + } + // The random slopes must rank with the true ones, and the group + // labels must come back in first-appearance order. + if len(res.GroupLabels) != 6 { + t.Fatalf("group labels %v, want six groups", res.GroupLabels) + } + strongest := 0 + weakest := 0 + for g, s := range slopeValues { + if s > slopeValues[strongest] { + strongest = g + } + if s < slopeValues[weakest] { + weakest = g + } + } + if !(res.RandomEffects[weakest][0] < res.RandomEffects[strongest][0]) { + t.Fatalf("the random slopes do not rank with the truth (%v against %v)", + res.RandomEffects[weakest][0], slopeValues[strongest]) + } + for _, se := range res.StandardErrors { + if !(se > 0) || math.IsInf(se, 0) { + t.Fatalf("the standard error %g is not finite and positive", se) + } + } +} + +func TestMixedModelDivergenceRefusal(t *testing.T) { + // A random design that spans the fixed one under an unstructured + // covariance: the intercept and slope of every group absorb what + // the fixed effects name, and the REML surface climbs a ridge of + // singular Σ without a summit. The fit refuses with the condition + // named instead of publishing the climb. + groups := []int{0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2} + xs := []float64{-1, -0.5, 0.5, 1, -1, -0.5, 0.5, 1, -1, -0.5, 0.5, 1} + y := make([]float64, len(groups)) + for i, g := range groups { + y[i] = 1 + 2*xs[i] + 0.4*float64(g)*xs[i] + float64(g) + 0.1*mixedNoise(i) + } + design := make([]float64, 0, 2*len(xs)) + for _, x := range xs { + design = append(design, 1, x) + } + _, err := LinearMixedModel( + mixedVec(t, y, len(y)), + mixedVec(t, design, len(xs), 2), + mixedVec(t, design, len(xs), 2), + groups) + if err == nil { + t.Fatal("a saturated random design was accepted") + } +} + +func TestMixedModelMatchesOLSWithoutRandomEffects(t *testing.T) { + // When the groups carry no shared signal the components collapse + // towards zero and the fit must land on the plain least squares + // answer. + xs := make([]float64, 18) + y := make([]float64, 18) + for i := range 18 { + x := -1 + 2*float64(i)/17 + xs[i] = x + y[i] = 3 - x + 0.4*mixedNoise(i) + } + groups := make([]int, 18) + for i := range 18 { + groups[i] = i % 6 + } + design := make([]float64, 0, 36) + for _, x := range xs { + design = append(design, 1, x) + } + res, err := LinearMixedModel( + mixedVec(t, y, len(y)), + mixedVec(t, design, len(xs), 2), + mixedVec(t, design, len(xs), 2), + groups) + if err != nil { + t.Fatal(err) + } + ref, rerr := LinearRegression(mixedVec(t, design, len(xs), 2), mixedVec(t, y, len(y))) + if rerr != nil { + t.Fatal(rerr) + } + for j := range 2 { + if math.Abs(res.Coefficients[j]-ref.Coefficients[j]) > 0.01 { + t.Fatalf("coefficient %d = %.6f, want the OLS %.6f", j, res.Coefficients[j], ref.Coefficients[j]) + } + } + if res.RandomCovariance[0] > 0.01 { + t.Fatalf("Σ̂ = %g on a group-free sample, want a collapsed component", res.RandomCovariance[0]) + } +} + +func TestMixedModelDeterministic(t *testing.T) { + y := make([]float64, 16) + groups := make([]int, 16) + for i := range 16 { + y[i] = 2 + 0.9*math.Sin(float64(i%4)) + mixedNoise(i) + groups[i] = i / 4 + } + run := func() *LinearMixedModelResult { + res, err := LinearMixedModel( + mixedVec(t, y, len(y)), + mixedVec(t, ones(len(y)), len(y), 1), + mixedVec(t, ones(len(y)), len(y), 1), + groups) + if err != nil { + t.Fatal(err) + } + return res + } + a, b := run(), run() + if a.LogLikelihood != b.LogLikelihood || a.ResidualVariance != b.ResidualVariance || + a.RandomCovariance[0] != b.RandomCovariance[0] || a.Iterations != b.Iterations { + t.Fatal("two identical fits disagreed") + } + for i := range a.Fitted { + if a.Fitted[i] != b.Fitted[i] { + t.Fatalf("row %d: fitted %g against %g", i, a.Fitted[i], b.Fitted[i]) + } + } +} + +func TestMixedModelShuffledLabels(t *testing.T) { + // Labels out of order and with gaps must canonicalise by first + // appearance, and permuting the rows with their labels must leave + // the fitted components where they started. + y := []float64{1.0, 1.2, 3.0, 3.1, 5.2, 5.1, 1.1, 3.2, 5.0, 1.3} + groups := []int{7, 7, 3, 3, 5, 5, 7, 3, 5, 7} + res, err := LinearMixedModel( + mixedVec(t, y, len(y)), + mixedVec(t, ones(len(y)), len(y), 1), + mixedVec(t, ones(len(y)), len(y), 1), + groups) + if err != nil { + t.Fatal(err) + } + want := []int{7, 3, 5} + for i, label := range res.GroupLabels { + if label != want[i] { + t.Fatalf("group labels %v, want %v", res.GroupLabels, want) + } + } + if !(res.RandomEffects[0][0] < res.RandomEffects[1][0] && res.RandomEffects[1][0] < res.RandomEffects[2][0]) { + t.Fatalf("the random effects %v do not rank with the group means", res.RandomEffects) + } +} + +func TestMixedModelRefusals(t *testing.T) { + y := mixedVec(t, []float64{1, 2, 3, 4}, 4) + x := mixedVec(t, []float64{1, 1, 1, 1}, 4, 1) + z := mixedVec(t, []float64{1, 1, 1, 1}, 4, 1) + groups := []int{0, 0, 1, 1} + if _, err := LinearMixedModel(mixedVec(t, []float64{1, 2}, 2, 1), x, z, groups); err == nil { + t.Fatal("a rank-2 response was accepted") + } + if _, err := LinearMixedModel(mixedVec(t, []float64{1, 2, 3, 4, 5}, 5), x, z, groups); err == nil { + t.Fatal("a row count mismatch was accepted") + } + if _, err := LinearMixedModel(y, mixedVec(t, []float64{1, 1, 1, 1}, 4, 1), mixedVec(t, []float64{1, 1, 1, 1}, 4, 1), []int{0, 1, 2}); err == nil { + t.Fatal("a short label vector was accepted") + } + if _, err := LinearMixedModel(y, x, z, []int{0, 0, 1, -2}); err == nil { + t.Fatal("a negative label was accepted") + } + // More coefficients than observations. + big := mixedVec(t, []float64{1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1, 1}, 4, 3) + if _, err := LinearMixedModel(y, big, z, groups); err == nil { + t.Fatal("a saturated design was accepted") + } + // A singular fixed design: two identical columns. + sing := mixedVec(t, []float64{1, 1, 1, 1, 2, 2, 2, 2}, 4, 2) + if _, err := LinearMixedModel(y, sing, z, groups); err == nil { + t.Fatal("a collinear fixed design was accepted") + } + // A single observation carries no fit. + if _, err := LinearMixedModel(mixedVec(t, []float64{1}, 1), mixedVec(t, []float64{1}, 1, 1), mixedVec(t, []float64{1}, 1, 1), []int{0}); err == nil { + t.Fatal("a one-row fit was accepted") + } +} + +// cholSolveInTest factors a symmetric positive definite matrix by its +// own Cholesky and solves against one right-hand side, an independent +// route the REML referent below evaluates its pieces through. +func cholSolveInTest(v []float64, rhs []float64, m int) (logDet float64, solution []float64) { + l := make([]float64, m*m) + for i := range m { + for j := range i + 1 { + s := v[i*m+j] + for k := range j { + s -= l[i*m+k] * l[j*m+k] + } + if i == j { + l[i*m+j] = math.Sqrt(s) + } else { + l[i*m+j] = s / l[j*m+j] + } + } + logDet += 2 * math.Log(l[i*m+i]) + } + x := append([]float64(nil), rhs...) + for i := range m { + s := x[i] + for k := range i { + s -= l[i*m+k] * x[k] + } + x[i] = s / l[i*m+i] + } + for i := m - 1; i >= 0; i-- { + s := x[i] + for k := i + 1; k < m; k++ { + s -= l[k*m+i] * x[k] + } + x[i] = s / l[i*m+i] + } + return logDet, x +} + +func TestMixedModelREMLLogLikelihoodReferent(t *testing.T) { + // The reported LogLikelihood is checked against an independent + // evaluation of the REML criterion at the fitted components, + // assembled from first principles in this test: + // -2·logL = Σ_g log|V_g| + rᵀV⁻¹r + log|XᵀV⁻¹X| + (n−p)·ln 2π, + // with V_g = Z_gΣZ_gᵀ + σ²I and r the fixed-part residual. The + // one-way balanced design keeps V_g compound symmetric, so the + // pieces are small and the route shares no arithmetic with the fit. + const ( + mGroups = 6 + k = 4 + ) + y := make([]float64, 0, mGroups*k) + groups := make([]int, mGroups*k) + for g := range mGroups { + effect := 1.5 * math.Sin(0.9*float64(g)+0.2) + for i := range k { + y = append(y, 3+effect+mixedNoise(g*k+i)) + groups[g*k+i] = g + } + } + res, err := LinearMixedModel( + mixedVec(t, y, len(y)), + mixedVec(t, ones(len(y)), len(y), 1), + mixedVec(t, ones(len(y)), len(y), 1), + groups) + if err != nil { + t.Fatal(err) + } + if !res.Converged { + t.Fatalf("the fit did not converge (%d iterations)", res.Iterations) + } + sigma2 := res.ResidualVariance + tau2 := res.RandomCovariance[0] + beta := res.Coefficients[0] + n := len(y) + p := 1 + logDetV := 0.0 + quad := 0.0 + xtvix := 0.0 + for g := range mGroups { + m := k + v := make([]float64, m*m) + for i := range m { + for j := range m { + v[i*m+j] = tau2 + } + v[i*m+i] += sigma2 + } + r := make([]float64, m) + for i := range m { + r[i] = y[g*k+i] - beta + } + ld, u := cholSolveInTest(v, r, m) + logDetV += ld + for i := range m { + quad += r[i] * u[i] + } + onesRHS := make([]float64, m) + for i := range onesRHS { + onesRHS[i] = 1 + } + _, w := cholSolveInTest(v, onesRHS, m) + for i := range m { + xtvix += w[i] + } + } + want := -0.5 * (logDetV + quad + math.Log(xtvix) + float64(n-p)*math.Log(2*math.Pi)) + if math.Abs(res.LogLikelihood-want) > 1e-8*(1+math.Abs(want)) { + t.Fatalf("LogLikelihood = %.10f, want the independent REML %.10f (difference %.3e)", + res.LogLikelihood, want, res.LogLikelihood-want) + } +} diff --git a/stats/multipletest.go b/stats/multipletest.go new file mode 100644 index 0000000..83c0e56 --- /dev/null +++ b/stats/multipletest.go @@ -0,0 +1,112 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +import ( + "cmp" + "math" + "slices" +) + +// Multiple-testing corrections over a vector of p-values: Bonferroni, +// Holm's step-down and Benjamini-Hochberg's step-up, each returning +// the adjusted p-value vector the caller can threshold directly. The +// Holm and Benjamini-Hochberg procedures are defined through a running +// extreme over the sorted p-values, which enforces the monotonicity +// their adjusted values must show: equal or larger raw p-values can +// never receive smaller adjustments, and the enforced running extreme +// maps every sorted position back to its own index. Every adjustment +// is at least the raw p-value it belongs to and never leaves [0, 1]. + +// Bonferroni returns the Bonferroni-adjusted p-values: each p scaled +// by the vector length m and clamped at 1, the family-wise error rate +// control that asks nothing of the dependence between the tests. +func Bonferroni(p []float64) ([]float64, error) { + const name = "Bonferroni" + if err := checkPValues(name, p); err != nil { + return nil, err + } + out := make([]float64, len(p)) + for i, v := range p { + out[i] = min(1, v*float64(len(p))) + } + return out, nil +} + +// Holm returns the Holm step-down adjusted p-values. Sorted ascending, +// the ith smallest receives the multiplier m−i and the running maximum +// over its predecessors enforces the non-decreasing order the step-down +// procedure implies; the adjusted vector is then mapped back through +// the original positions. +func Holm(p []float64) ([]float64, error) { + const name = "Holm" + if err := checkPValues(name, p); err != nil { + return nil, err + } + m := len(p) + order := make([]int, m) + for i := range order { + order[i] = i + } + slices.SortFunc(order, func(a, b int) int { return cmp.Compare(p[a], p[b]) }) + out := make([]float64, m) + running := 0.0 + for rank, idx := range order { + running = max(running, float64(m-rank)*p[idx]) + // The max against the raw p is arithmetic beltwork: the + // multiplier never drops below 1, and this pins the guarantee + // exactly rather than to rounding. + out[idx] = min(1, max(running, p[idx])) + } + return out, nil +} + +// BenjaminiHochberg returns the Benjamini-Hochberg step-up adjusted +// p-values, the q-values of the false discovery rate literature. +// Sorted ascending, the ith smallest receives the multiplier m/(i+1) +// and the running minimum over its successors enforces the +// non-decreasing order the step-up procedure implies; the adjusted +// vector is then mapped back through the original positions. +func BenjaminiHochberg(p []float64) ([]float64, error) { + const name = "BenjaminiHochberg" + if err := checkPValues(name, p); err != nil { + return nil, err + } + m := len(p) + order := make([]int, m) + for i := range order { + order[i] = i + } + slices.SortFunc(order, func(a, b int) int { return cmp.Compare(p[a], p[b]) }) + out := make([]float64, m) + running := 1.0 + for i := m - 1; i >= 0; i-- { + idx := order[i] + running = min(running, float64(m)/float64(i+1)*p[idx]) + out[idx] = min(1, max(running, p[idx])) + } + return out, nil +} + +// checkPValues validates a p-value vector for the corrections: it must +// be non-empty, every entry finite, and every entry inside [0, 1], the +// refusals reported with the offending index and value. +func checkPValues(name string, p []float64) error { + if len(p) == 0 { + return base.Errf("%s: the p-value vector must not be empty", name) + } + for i, v := range p { + if math.IsNaN(v) || math.IsInf(v, 0) { + return base.Errf("%s: p[%d] is not finite (%g)", name, i, v) + } + if v < 0 || v > 1 { + return base.Errf("%s: p[%d] = %g lies outside [0, 1]", name, i, v) + } + } + return nil +} diff --git a/stats/multipletest_test.go b/stats/multipletest_test.go new file mode 100644 index 0000000..47ff128 --- /dev/null +++ b/stats/multipletest_test.go @@ -0,0 +1,175 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "testing" +) + +// The worked vector for all three corrections: sorted it reads +// 0.005, 0.01, 0.03, 0.04 and every adjusted value below is computed +// by hand from that order. +var multipleTestP = []float64{0.01, 0.04, 0.03, 0.005} + +// TestBonferroniAdjusted pins the hand values m·p: (0.04, 0.16, 0.12, +// 0.02), and the clamp at 1. +func TestBonferroniAdjusted(t *testing.T) { + got, err := Bonferroni(multipleTestP) + if err != nil { + t.Fatalf("Bonferroni: %v", err) + } + want := []float64{0.04, 0.16, 0.12, 0.02} + for i := range want { + if math.Abs(got[i]-want[i]) > 1e-15 { + t.Fatalf("Bonferroni[%d] = %.17g, want %.17g", i, got[i], want[i]) + } + } + clamped, err := Bonferroni([]float64{0.5, 0.6}) + if err != nil || clamped[0] != 1 || clamped[1] != 1 { + t.Fatalf("clamped Bonferroni = %v, want [1, 1]", clamped) + } +} + +// TestHolmAdjusted pins the step-down walk. Sorted, the multipliers +// 4, 3, 2, 1 give 0.02, 0.03, 0.06, 0.06 after the running maximum, +// mapped back through the original order to (0.03, 0.06, 0.06, 0.02). +func TestHolmAdjusted(t *testing.T) { + got, err := Holm(multipleTestP) + if err != nil { + t.Fatalf("Holm: %v", err) + } + want := []float64{0.03, 0.06, 0.06, 0.02} + for i := range want { + if math.Abs(got[i]-want[i]) > 1e-15 { + t.Fatalf("Holm[%d] = %.17g, want %.17g", i, got[i], want[i]) + } + } +} + +// TestBenjaminiHochbergAdjusted pins the step-up walk. Sorted, from +// the top: 1·0.04 = 0.04, min(0.04, 4/3·0.03) = 0.04, +// min(0.04, 2·0.01) = 0.02, min(0.02, 4·0.005) = 0.02, so the original +// order carries (0.02, 0.04, 0.04, 0.02). Evenly spaced p-values all +// receive the largest raw p as their q-value, the identity 5/j·j/100 +// = 0.05 on (0.01 .. 0.05). +func TestBenjaminiHochbergAdjusted(t *testing.T) { + got, err := BenjaminiHochberg(multipleTestP) + if err != nil { + t.Fatalf("BenjaminiHochberg: %v", err) + } + want := []float64{0.02, 0.04, 0.04, 0.02} + for i := range want { + if math.Abs(got[i]-want[i]) > 1e-15 { + t.Fatalf("BenjaminiHochberg[%d] = %.17g, want %.17g", i, got[i], want[i]) + } + } + uniform, err := BenjaminiHochberg([]float64{0.01, 0.02, 0.03, 0.04, 0.05}) + if err != nil { + t.Fatalf("BenjaminiHochberg uniform: %v", err) + } + for i, v := range uniform { + if math.Abs(v-0.05) > 1e-15 { + t.Fatalf("uniform q[%d] = %.17g, want 0.05", i, v) + } + } +} + +// TestBenjaminiHochbergLiterature pins the fourteen p-values tabulated +// in Benjamini and Hochberg (1995), Controlling the false discovery +// rate, the classic worked example, with the q-values derived by hand +// from the sorted vector: q_i = min over j ≥ i of 14·p_j/j. From the +// top the raw products fall 0.7590, 14·0.6528/13 = 0.7030, +// 14·0.5719/12 = 0.6672, 14·0.4262/11 = 0.5424, 14·0.3240/10 = 0.4536 +// and then 0.0714 on rank 9, so the running minimum holds 0.4536 +// across ranks 10 and 11, 0.6672 on rank 12 and 0.7030 on rank 13; +// below, 0.0602 on rank 8 and the running minimum carries 0.0596 from +// rank 7 back over rank 6 (whose raw product is 0.0648). +func TestBenjaminiHochbergLiterature(t *testing.T) { + p := []float64{ + 0.0001, 0.0004, 0.0019, 0.0095, 0.0201, 0.0278, 0.0298, + 0.0344, 0.0459, 0.3240, 0.4262, 0.5719, 0.6528, 0.7590, + } + got, err := BenjaminiHochberg(p) + if err != nil { + t.Fatalf("BenjaminiHochberg: %v", err) + } + want := []float64{ + 0.0014, 0.0028, 14 * 0.0019 / 3, 0.03325, 0.05628, 0.0596, 0.0596, + 0.0602, 0.0714, 0.4536, 14 * 0.4262 / 11, 14 * 0.5719 / 12, 14 * 0.6528 / 13, 0.7590, + } + for i := range want { + if math.Abs(got[i]-want[i]) > 1e-14 { + t.Fatalf("q[%d] = %.17g, want %.17g", i, got[i], want[i]) + } + } +} + +// TestMultipleTestProperties holds the contract every correction +// states: adjusted values never fall below their raw p-values, never +// leave [0, 1], and preserve the order of the raw p-values. +func TestMultipleTestProperties(t *testing.T) { + raw := []float64{0.5, 0.002, 0.2, 0.04, 0.009, 0.7, 0.0001, 0.13} + corrections := []struct { + name string + apply func([]float64) ([]float64, error) + }{ + {"Bonferroni", Bonferroni}, + {"Holm", Holm}, + {"BenjaminiHochberg", BenjaminiHochberg}, + } + for _, c := range corrections { + got, err := c.apply(raw) + if err != nil { + t.Fatalf("%s: %v", c.name, err) + } + for i := range raw { + if got[i] < raw[i] { + t.Fatalf("%s[%d] = %v sits below the raw %v", c.name, i, got[i], raw[i]) + } + if got[i] < 0 || got[i] > 1 { + t.Fatalf("%s[%d] = %v leaves [0, 1]", c.name, i, got[i]) + } + for j := range raw { + if raw[i] <= raw[j] && got[i] > got[j] { + t.Fatalf("%s inverts the order at (%d, %d): %v raw %v vs %v raw %v", + c.name, i, j, got[i], raw[i], got[j], raw[j]) + } + } + } + } + // Zeros stay zeros under every correction. + zeroed, err := Holm([]float64{0, 0.2}) + if err != nil || zeroed[0] != 0 { + t.Fatalf("Holm on a zero p-value = %v, %v, want 0 untouched", zeroed, err) + } + // Equal p-values receive equal adjustments through the running + // extreme. + tied, _ := Holm([]float64{0.05, 0.05}) + if tied[0] != tied[1] { + t.Fatalf("tied Holm adjustments differ: %v", tied) + } +} + +// TestMultipleTestErrors pins the input contract on all three +// corrections. +func TestMultipleTestErrors(t *testing.T) { + for _, apply := range []func([]float64) ([]float64, error){Bonferroni, Holm, BenjaminiHochberg} { + if _, err := apply(nil); err == nil { + t.Fatal("an empty vector: want an error") + } + if _, err := apply([]float64{0.1, math.NaN()}); err == nil { + t.Fatal("a NaN p-value: want an error") + } + if _, err := apply([]float64{0.1, math.Inf(1)}); err == nil { + t.Fatal("an infinite p-value: want an error") + } + if _, err := apply([]float64{1.5}); err == nil { + t.Fatal("p above 1: want an error") + } + if _, err := apply([]float64{-0.1}); err == nil { + t.Fatal("a negative p-value: want an error") + } + } +} diff --git a/stats/mvn.go b/stats/mvn.go new file mode 100644 index 0000000..7075626 --- /dev/null +++ b/stats/mvn.go @@ -0,0 +1,210 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The multivariate normal: densities through a Cholesky solve and +// draws through the same triangular factor, the two operations any +// Bayesian or Monte Carlo workflow reaches for. + +// mvnSymmetryEps is the relative tolerance the covariance's mirror +// check allows. A covariance assembled as A·Aᵀ can differ from its +// mirror by an ulp of rounding; a matrix that is genuinely asymmetric +// differs by far more than this, and reading only its lower triangle +// would describe a different distribution from the one handed in. +const mvnSymmetryEps = 1e-12 + +// mvnCholesky factors a symmetric positive-definite matrix into the +// lower triangular L with A = L·Lᵀ; a non-positive pivot names the row +// in the error, and a mirror pair that disagrees beyond a relative +// mvnSymmetryEps names the entry, because only the lower triangle is +// read. +func mvnCholesky(name string, cov *core.Array, d int) ([][]float64, error) { + // The covariance is read once into a flat local copy: the symmetry + // sweep and the O(d³) factorisation then index a plain slice. The + // copy carries the values FloatAt returns, so every difference, + // product and pivot below keeps its exact bits. + covVals := make([]float64, d*d) + if fs := rawFloats(cov); fs != nil { + copy(covVals, fs) + } else { + for i := range covVals { + covVals[i] = cov.FloatAt(i) + } + } + return mvnCholeskyFlat(name, covVals, d) +} + +// mvnCholeskyFlat factors a matrix already held flat row-major, the +// shared body of the array entry point and of a caller that assembles +// the matrix in place. The slice is read, never written. +func mvnCholeskyFlat(name string, covVals []float64, d int) ([][]float64, error) { + l := make([][]float64, d) + for i := range d { + l[i] = make([]float64, d) + } + if err := mvnCholeskyFlatInto(name, covVals, l, d); err != nil { + return nil, err + } + return l, nil +} + +// mvnCholeskyFlatInto is mvnCholeskyFlat on a caller-owned d×d +// destination, for a sweep that factors one covariance after another: +// the rows are cleared and refilled, and every entry the factorisation +// reads is one it has already written in the same pass. The arithmetic +// is mvnCholeskyFlat's own, unchanged. +func mvnCholeskyFlatInto(name string, covVals []float64, l [][]float64, d int) error { + for i := range d { + for j := range i { + lo, hi := covVals[i*d+j], covVals[j*d+i] + if math.Abs(lo-hi) > mvnSymmetryEps*math.Max(math.Abs(lo), math.Abs(hi)) { + return base.Errf("%s: the covariance is not symmetric at (%d, %d): %g against %g", + name, i+1, j+1, lo, hi) + } + } + } + for i := range d { + row := l[i] + clear(row) + for j := range i + 1 { + total := covVals[i*d+j] + for k := range j { + total -= l[i][k] * l[j][k] + } + if i == j { + if total <= 0 { + return base.Errf("%s: the covariance is not positive definite at row %d", name, i+1) + } + row[j] = math.Sqrt(total) + } else { + row[j] = total / l[j][j] + } + } + } + return nil +} + +// MultivariateNormalLogDensity evaluates the density of the +// d-dimensional normal N(mean, cov) at one point x, all three rank-1 +// and cov symmetric positive definite, through the Cholesky factor: +// the log determinant is twice the sum of the factor's diagonal and +// the quadratic form the squared norm of the forward solve. Every +// entry of all three inputs must be finite, and the covariance must +// mirror itself: a non-finite value or an asymmetric pair is refused +// by name rather than answered with a NaN or with the density of a +// different distribution. +func MultivariateNormalLogDensity(mean, cov, x *core.Array) (float64, error) { + const name = "MultivariateNormalLogDensity" + if mean.NDim() != 1 || x.NDim() != 1 || cov.NDim() != 2 { + return 0, base.Errf("%s: mean and x must be rank 1 and cov rank 2", name) + } + if mean.Dtype() == core.Complex || cov.Dtype() == core.Complex || x.Dtype() == core.Complex { + return 0, base.Errf("%s: complex inputs are not supported", name) + } + d := mean.Len() + if x.Len() != d || cov.Shape()[0] != d || cov.Shape()[1] != d { + return 0, base.Errf("%s: mean holds %d entries, x %d, cov %s", name, d, x.Len(), base.ShapeText(cov.Shape())) + } + if err := checkFinite(name, "the mean", mean); err != nil { + return 0, err + } + if err := checkFinite(name, "the covariance", cov); err != nil { + return 0, err + } + if err := checkFinite(name, "the point", x); err != nil { + return 0, err + } + l, err := mvnCholesky(name, cov, d) + if err != nil { + return 0, err + } + diff := make([]float64, d) + solve := make([]float64, d) + logDet := 0.0 + for i := range d { + diff[i] = x.FloatAt(i) - mean.FloatAt(i) + logDet += math.Log(l[i][i]) + total := diff[i] + for j := range i { + total -= l[i][j] * solve[j] + } + solve[i] = total / l[i][i] + } + quad := 0.0 + for i := range d { + quad += solve[i] * solve[i] + } + return -0.5*float64(d)*math.Log(2*math.Pi) - logDet - 0.5*quad, nil +} + +// MultivariateNormalDraws returns n draws from N(mean, cov) as an +// (n × d) array, one draw per row: standard normals through the +// generator, coloured by the Cholesky factor of the covariance. The +// draws are deterministic for a given generator state. The mean and +// the covariance must be finite and the covariance symmetric, the +// same contract MultivariateNormalLogDensity enforces. +func MultivariateNormalDraws(g *core.Generator, n int, mean, cov *core.Array) (*core.Array, error) { + const name = "MultivariateNormalDraws" + if g == nil { + return nil, base.Errf("%s: the generator is nil", name) + } + if n < 1 { + return nil, base.Errf("%s: the count must be at least 1, got %d", name, n) + } + if mean.NDim() != 1 || cov.NDim() != 2 { + return nil, base.Errf("%s: mean must be rank 1 and cov rank 2", name) + } + if mean.Dtype() == core.Complex || cov.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + d := mean.Len() + if cov.Shape()[0] != d || cov.Shape()[1] != d { + return nil, base.Errf("%s: mean holds %d entries but cov is %s", name, d, base.ShapeText(cov.Shape())) + } + if err := checkFinite(name, "the mean", mean); err != nil { + return nil, err + } + if err := checkFinite(name, "the covariance", cov); err != nil { + return nil, err + } + l, err := mvnCholesky(name, cov, d) + if err != nil { + return nil, err + } + out := core.New(core.Float, n, d) + vals := out.RawFloats() + z := make([]float64, d) + // The mean is read once: every draw reuses the same values the + // accessor walk returned. + meanVals := make([]float64, d) + if fs := rawFloats(mean); fs != nil { + copy(meanVals, fs) + } else { + for i := range meanVals { + meanVals[i] = mean.FloatAt(i) + } + } + for r := range n { + // One standard normal vector per draw, coloured by L: the + // shared z is what makes the off-diagonal covariance appear. + for j := range d { + z[j] = g.NormalUnit() + } + for i := range d { + total := meanVals[i] + for j := range i + 1 { + total += l[i][j] * z[j] + } + vals[r*d+i] = total + } + } + return out, nil +} diff --git a/stats/negbinom_far_pin_test.go b/stats/negbinom_far_pin_test.go new file mode 100644 index 0000000..66efd4b --- /dev/null +++ b/stats/negbinom_far_pin_test.go @@ -0,0 +1,60 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Pin for the negative binomial's far regime: when the summation's seed +// p^r underflows to an exact zero the recurrence can no longer start, +// and the CDF answers through the regularised beta identity instead. +// The law with r = 100 and p = 1e-4 has its mean at 999 900, so the CDF +// at the mean is near one half; the seed 1e-400 underflows outright, and +// the summation that ignored that answered an exact zero at every k. + +package stats + +import ( + "math" + "testing" +) + +func TestNegativeBinomialFarTail(t *testing.T) { + atMean, err := NegativeBinomialCDF(999900, 100, 1e-4) + if err != nil { + t.Fatalf("NegativeBinomialCDF(999900, 100, 1e-4): %v", err) + } + // The mean is r(1−p)/p = 999 900 and the spread is near 100 000, so + // the true value at the mean sits within a few 1e-3 of one half + // (exactly 0.5133 by the exact referent). The identity's accuracy in + // this regime is bounded by the lgamma noise of its front factor, so + // the pin holds a loose band and a strict ordering. + if math.Abs(atMean-0.5133007914) > 1e-6 { + t.Fatalf("NegativeBinomialCDF at the mean = %.17g, want 0.5133007914 to within 1e-6", atMean) + } + below, err := NegativeBinomialCDF(950000, 100, 1e-4) + if err != nil { + t.Fatalf("NegativeBinomialCDF(950000, 100, 1e-4): %v", err) + } + above, err := NegativeBinomialCDF(1050000, 100, 1e-4) + if err != nil { + t.Fatalf("NegativeBinomialCDF(1050000, 100, 1e-4): %v", err) + } + if !(below < atMean && atMean < above) { + t.Fatalf("the CDF is not increasing through the mean: %.17g, %.17g, %.17g", below, atMean, above) + } + far, err := NegativeBinomialCDF(2000000, 100, 1e-4) + if err != nil { + t.Fatalf("NegativeBinomialCDF(2000000, 100, 1e-4): %v", err) + } + if far < 0.9999 { + t.Fatalf("NegativeBinomialCDF two hundred standard deviations out = %.17g, want ~1", far) + } + // The quantile route runs on the same CDF: the median must sit near + // the mean the law's own moments give, a fraction of the spread + // below it on the skewed side. The exact referent puts the median at + // 996 569, and the pin holds it within a thousandth of the spread. + median, err := NegativeBinomialQuantile(0.5, 1e-4, 100) + if err != nil { + t.Fatalf("NegativeBinomialQuantile: %v", err) + } + if math.Abs(median-996569) > 100 { + t.Fatalf("NegativeBinomialQuantile(0.5) = %v, want within 100 of the exact median 996569", median) + } +} diff --git a/stats/noncentral.go b/stats/noncentral.go new file mode 100644 index 0000000..7b599d7 --- /dev/null +++ b/stats/noncentral.go @@ -0,0 +1,441 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" +) + +import ( + "math" +) + +// Noncentral distributions: the χ², t and F laws with a noncentrality +// parameter, the laws the power of every test in this package runs on. +// The χ² is the Poisson mixture of its central family, the exact +// identity a noncentral χ²(ν, λ) draw is χ²(ν + 2J) with J a +// Poisson(λ/2) count; the F mixes the same numerator against the +// central denominator, whose pieces carry the scale (ν₁+2i)/ν₁. The +// noncentral t runs Lenth's algorithm +// (Applied Statistics 38, 1989, pages 185 to 189): the CDF as a sum of +// even terms, the Poisson-weighted I_x(j+½, ν/2) of the folded |t|, +// and odd terms, the half-normal-weighted I_x(j+1, ν/2) that carry the +// sign of δ, over x = t²/(t²+ν). Every series here sums positive +// well-scaled terms until the remaining Poisson mass bounds the +// truncation below the floor, and reports an error rather than a +// silently truncated value if the budget runs out. The test file holds +// all three against a direct quadrature of E[Φ(t√(V/ν) − δ)] and +// against the closed forms of the degenerate corners. + +// noncentralTermFloor bounds the weight a mixture term may still carry +// when the sum stops: the remaining Poisson mass is below it, so the +// omitted tail cannot reach the 15th digit of the answer. +const noncentralTermFloor = 1e-18 + +// maxNoncentralTerms bounds the mixture loops. The weights peak at the +// index ⌊λ/2⌋ and the walk needs the peak plus a few standard +// deviations of Poisson spread to cross it, so the budget carries +// noncentralities up to roughly 2·10⁵ in the χ² and F and δ up to +// about 440 in the t; beyond that the refusal is explicit, and so is +// every weight the format cannot hold: they are computed term by term +// in log space, never climbed from a seed that could underflow to +// zero and take the whole sum with it. +const maxNoncentralTerms = 100000 + +// noncentralPoissonWeight is the Poisson(half) weight of the index i, +// computed term by term in log space. The multiplicative climb from +// the e^{−half} seed the series definitions start from underflows to +// an exact zero once half passes about 745, and a zero seed never +// recovers: every later weight would stay zero while the loop believed +// it had converged. Evaluating each weight from its own logarithm +// keeps the terms near the peak exact at any half the budget can walk +// past, and the genuinely negligible ones answer zero, which is what +// they are. +func noncentralPoissonWeight(half float64, i int) float64 { + if half == 0 { + if i == 0 { + return 1 + } + return 0 + } + return math.Exp(-half + float64(i)*math.Log(half) - logGamma(float64(i)+1)) +} + +// noncentralBudgetRefused reports the explicit refusal when the weight +// peak of a Poisson(half) mixture sits past the term budget, where the +// walk would stop early with a wrong answer instead. +func noncentralBudgetRefused(name string, half float64) error { + return base.Errf("%s: the weight peak at %d needs a walk past the %d-term budget; the mixture answers only up to that noncentrality", + name, int(math.Floor(half)), maxNoncentralTerms) +} + +// noncentralPeakInsideBudget reports whether the Poisson weight peak +// at ⌊half⌋ plus its spread sits inside the term budget. +func noncentralPeakInsideBudget(half float64) bool { + peak := math.Floor(half) + return float64(maxNoncentralTerms) >= peak+8*math.Sqrt(peak)+2 +} + +// noncentralOddWeight is the j-th half-normal weight of Lenth's odd +// series, δ·λ^j·p_0/(√(2π)·(2j+1)!!) with λ = δ² and p_0 the j = 0 +// Poisson weight, carrying the sign of δ. The double factorial comes +// out of its log-space form (2j+1)!! = 2^{j+1}Γ(j+3/2)/√π, so the +// weight is computed from its own logarithm like the even part's and +// underflows only once it is genuinely negligible. +func noncentralOddWeight(shift, half float64, j int) float64 { + if shift == 0 { + return 0 + } + lambda := shift * shift + if lambda == 0 { + return 0 + } + ln := math.Log(math.Abs(shift)) + float64(j)*math.Log(lambda) - half - + (float64(j)+1.5)*math.Ln2 - logGamma(float64(j)+1.5) + return math.Copysign(math.Exp(ln), shift) +} + +// NoncentralChiSquareCDF returns P(X ≤ x) for X ~ χ²(ν, λ), the +// Poisson(λ/2) mixture of central χ²(ν + 2i) CDFs, each through the +// existing GammaLower. The noncentrality λ must be non-negative and +// finite; λ = 0 answers through ChiSquareCDF exactly. +func NoncentralChiSquareCDF(x float64, df int, lambda float64) (float64, error) { + const name = "NoncentralChiSquareCDF" + if df < 1 { + return 0, base.Errf("%s: df must be ≥ 1, got %d", name, df) + } + if math.IsNaN(lambda) || lambda < 0 || math.IsInf(lambda, 0) { + return 0, base.Errf("%s: lambda must be finite and non-negative, got %g", name, lambda) + } + if math.IsNaN(x) { + return 0, base.Errf("%s: x must be a number, got %g", name, x) + } + if x <= 0 { + return 0, nil + } + if lambda == 0 { + return ChiSquareCDF(x, df) + } + return noncentralPoissonMixture(name, df, lambda, + func(i int) (float64, error) { + g, err := GammaLower(float64(df)/2+float64(i), x/2) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + return g, nil + }) +} + +// NoncentralChiSquareDensity returns the χ²(ν, λ) density at x, the +// same Poisson mixture with central χ² densities, each carrying a +// closed exponential-power form. The support convention gives 0 below +// x = 0; at x = 0 with df = 1 the density is the +Inf the +// noncentralities preserve, with df = 2 it is the finite limit +// e^{−λ/2}/2 of the j = 0 mixture term, and past df = 2 it is 0. +func NoncentralChiSquareDensity(x float64, df int, lambda float64) (float64, error) { + const name = "NoncentralChiSquareDensity" + if df < 1 { + return 0, base.Errf("%s: df must be ≥ 1, got %d", name, df) + } + if math.IsNaN(lambda) || lambda < 0 || math.IsInf(lambda, 0) { + return 0, base.Errf("%s: lambda must be finite and non-negative, got %g", name, lambda) + } + if math.IsNaN(x) { + return 0, base.Errf("%s: x must be a number, got %g", name, x) + } + if x < 0 || (x == 0 && df > 2) { + return 0, nil + } + if x == 0 { + // df = 1: the x^{−½} singularity the mixture integrates; df = 2: + // the j = 0 term is finite at the origin and its limit + // e^{−λ/2}/2 is the answer. + if df == 2 { + return 0.5 * math.Exp(-lambda/2), nil + } + return math.Inf(1), nil + } + lambdaHalf := lambda / 2 + if !noncentralPeakInsideBudget(lambdaHalf) { + return 0, noncentralBudgetRefused(name, lambdaHalf) + } + total := 0.0 + for i := range maxNoncentralTerms { + weight := noncentralPoissonWeight(lambdaHalf, i) + a := float64(df)/2 + float64(i) + density := weight * math.Exp(-x/2+(a-1)*math.Log(x)-a*math.Ln2-logGamma(a)) + total += density + if float64(i) > lambdaHalf+1 && weight < noncentralTermFloor { + return total, nil + } + } + return 0, base.Errf("%s: the mixture did not converge within %d terms for lambda = %g", + name, maxNoncentralTerms, lambda) +} + +// NoncentralChiSquareQuantile returns the q-quantile of χ²(ν, λ) by +// the same bracketed bisection the central tables use, seeded near the +// mean ν + λ. +func NoncentralChiSquareQuantile(q float64, df int, lambda float64) (float64, error) { + if df < 1 { + return 0, base.Errf("NoncentralChiSquareQuantile: df must be ≥ 1, got %d", df) + } + if math.IsNaN(lambda) || lambda < 0 || math.IsInf(lambda, 0) { + return 0, base.Errf("NoncentralChiSquareQuantile: lambda must be finite and non-negative, got %g", lambda) + } + return continuousQuantile("NoncentralChiSquareQuantile", q, float64(df)+lambda/2, + func(x float64) (float64, error) { + return NoncentralChiSquareCDF(x, df, lambda) + }, nil) +} + +// NoncentralFCDF returns P(X ≤ x) for X ~ F(ν₁, ν₂, λ), the numerator +// χ²(ν₁, λ) carried against the central denominator: the Poisson(λ/2) +// mixture of the scaled central pieces (ν₁+2i)/ν₁·F(ν₁+2i, ν₂), whose +// beta form sums I_{ν₁x/(ν₁x+ν₂)}((ν₁+2i)/2, ν₂/2) over the weights, +// through the existing BetaIncomplete. The λ = 0 corner is the central +// F exactly. +func NoncentralFCDF(x float64, df1, df2 int, lambda float64) (float64, error) { + const name = "NoncentralFCDF" + if df1 < 1 || df2 < 1 { + return 0, base.Errf("%s: df1 and df2 must be ≥ 1, got %d and %d", name, df1, df2) + } + if math.IsNaN(lambda) || lambda < 0 || math.IsInf(lambda, 0) { + return 0, base.Errf("%s: lambda must be finite and non-negative, got %g", name, lambda) + } + if math.IsNaN(x) { + return 0, base.Errf("%s: x must be a number, got %g", name, x) + } + if math.IsInf(x, 1) { + return 1, nil + } + if x <= 0 { + return 0, nil + } + // The beta argument saturates at 1 for an x so large the product + // ν₁x overflows, which is the CDF's own limit there. + numerator := float64(df1) * x + arg := 1.0 + if !math.IsInf(numerator, 1) { + arg = numerator / (numerator + float64(df2)) + } + return noncentralPoissonMixture(name, df1, lambda, + func(i int) (float64, error) { + p, err := BetaIncomplete(arg, (float64(df1)+2*float64(i))/2, float64(df2)/2) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + return p, nil + }) +} + +// NoncentralFQuantile returns the q-quantile of F(ν₁, ν₂, λ) by +// bracketed bisection, seeded at 1 in the neighbourhood of the F +// median. +func NoncentralFQuantile(q float64, df1, df2 int, lambda float64) (float64, error) { + if df1 < 1 || df2 < 1 { + return 0, base.Errf("NoncentralFQuantile: df1 and df2 must be ≥ 1, got %d and %d", df1, df2) + } + if math.IsNaN(lambda) || lambda < 0 || math.IsInf(lambda, 0) { + return 0, base.Errf("NoncentralFQuantile: lambda must be finite and non-negative, got %g", lambda) + } + return continuousQuantile("NoncentralFQuantile", q, 1, func(x float64) (float64, error) { + return NoncentralFCDF(x, df1, df2, lambda) + }, nil) +} + +// noncentralPoissonMixture sums w_i·term(i) over the Poisson(λ/2) +// weights w_i, the shared engine of the noncentral χ² and F. All terms +// are positive, so the running sum carries no cancellation; the walk +// stops past the weight peak once the weights have sunk below the +// floor. The weights are computed term by term in log space, so no +// noncentrality the budget can reach underflows the walk. +func noncentralPoissonMixture(name string, df int, lambda float64, + term func(i int) (float64, error)) (float64, error) { + half := lambda / 2 + if !noncentralPeakInsideBudget(half) { + return 0, noncentralBudgetRefused(name, half) + } + total := 0.0 + for i := range maxNoncentralTerms { + weight := noncentralPoissonWeight(half, i) + t, err := term(i) + if err != nil { + return 0, err + } + total += t * weight + if float64(i) > half+1 && weight < noncentralTermFloor { + return total, nil + } + } + return 0, base.Errf("%s: the mixture did not converge within %d terms for lambda = %g", + name, maxNoncentralTerms, lambda) +} + +// NoncentralTCDF returns P(T ≤ t) for T ~ t(ν, δ), through Lenth's +// even and odd series: the even part folds the law through |t|, the +// Poisson(δ²/2)-weighted beta ratios of the |t| event, and the odd +// part carries the sign of δ through the √(2/π)δ(δ²)^j/(2j+1)!! +// half-normal weights. The truncation floor is the remaining Poisson +// mass the error bound 2s(xodd − godd) tracks, the bound the published +// algorithm proves. δ = 0 answers through StudentTCDF exactly, and +// t = 0 through the closed corner Φ(−δ). +func NoncentralTCDF(t float64, df int, delta float64) (float64, error) { + const name = "NoncentralTCDF" + if df < 1 { + return 0, base.Errf("%s: df must be ≥ 1, got %d", name, df) + } + if math.IsNaN(t) || math.IsInf(t, 0) { + return 0, base.Errf("%s: t must be finite, got %g", name, t) + } + if math.IsNaN(delta) || math.IsInf(delta, 0) { + return 0, base.Errf("%s: delta must be finite, got %g", name, delta) + } + if delta == 0 { + return StudentTCDF(t, df) + } + // The series is derived on t ≥ 0; the reflection F(t; δ) = + // 1 − F(−t; −δ), an exact identity of the law, covers the rest. + flipped := false + magnitude, shift := t, delta + if t < 0 { + flipped = true + magnitude = -t + shift = -delta + } + x := magnitude * magnitude / (magnitude*magnitude + float64(df)) + if x == 0 { + // t = 0: the value collapses to Φ(−δ) exactly. + return NormalCDF(-delta), nil + } + lambda := shift * shift + half := lambda / 2 + if !noncentralPeakInsideBudget(half) { + return 0, noncentralBudgetRefused(name, half) + } + // p_j are the Poisson(half) weights of the even part, q_j the + // half-normal weights of the odd part, q_j = δ·λ^j·p_0/(√(2π)·(2j+1)!!). + // Both are computed term by term in log space: the multiplicative + // climb from the e^{−λ/2} seed underflows to an exact zero once + // |δ| passes about 39, and a zero seed never recovers, which used + // to answer a silent 0 for the whole law. + p := 0.5 * noncentralPoissonWeight(half, 0) + q := noncentralOddWeight(shift, half, 0) + remaining := 0.5 - p + a := 0.5 + b := float64(df) / 2 + rxb := math.Pow(1-x, b) + // ln B(a, b) at a = ½. + lnBeta := 0.5*math.Log(math.Pi) + logGamma(b) - logGamma(a+b) + xodd, err := BetaIncomplete(x, a, b) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + // godd and geven are the beta-integral pieces the recurrences peel + // off xodd and xeven, the subtraction forms of I_x(a+1, b) and + // I_x(a, b+1): one beta evaluation seeds the whole walk. + godd := 2 * rxb * math.Exp(a*math.Log(x)-lnBeta) + xeven := 1 - rxb + geven := b * x * rxb + total := p*xodd + q*xeven + for en := 1.0; en <= maxNoncentralTerms; en++ { + a++ + xodd -= godd + xeven -= geven + godd *= x * (a + b - 1) / a + geven *= x * (a + b - 0.5) / (a + 0.5) + p = 0.5 * noncentralPoissonWeight(half, int(en)) + q = noncentralOddWeight(shift, half, int(en)) + remaining -= p + total += p*xodd + q*xeven + if bound := 2 * remaining * (xodd - godd); bound <= noncentralTermFloor { + total += NormalCDF(-shift) + if flipped { + total = 1 - total + } + return min(1, max(0, total)), nil + } + } + return 0, base.Errf("%s: the series did not converge within %d terms for delta = %g", + name, maxNoncentralTerms, delta) +} + +// NoncentralTQuantile returns the q-quantile of t(ν, δ) by bracketed +// bisection on the signed axis: the law leans towards δ, so the +// bracket grows from the seed in both directions. +func NoncentralTQuantile(q float64, df int, delta float64) (float64, error) { + if df < 1 { + return 0, base.Errf("NoncentralTQuantile: df must be ≥ 1, got %d", df) + } + if math.IsNaN(delta) || math.IsInf(delta, 0) { + return 0, base.Errf("NoncentralTQuantile: delta must be finite, got %g", delta) + } + return signedQuantile("NoncentralTQuantile", q, math.Abs(delta)+1, + func(t float64) (float64, error) { + return NoncentralTCDF(t, df, delta) + }) +} + +// signedQuantile inverts a continuous CDF over the whole real axis, +// the signed twin of continuousQuantile: the bracket starts at ±seed +// and doubles outwards until the CDF straddles q, then halves to +// rounding level under the same unconditional convergence. +func signedQuantile(name string, q float64, seed float64, + cdf func(float64) (float64, error)) (float64, error) { + // NaN-rejecting on purpose, as in continuousQuantile. + if !(q >= 0 && q <= 1) { + return 0, base.Errf("%s: q must lie in [0, 1], got %g", name, q) + } + if q == 0 || q == 1 { + return 0, base.Errf("%s: q = %g has no finite quantile", name, q) + } + lo, hi := -seed, seed + fLo, err := cdf(lo) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + fHi, err := cdf(hi) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + for fLo > q { + lo *= 2 + if math.IsInf(lo, 0) { + return 0, base.Errf("%s: failed to bracket q = %g from below", name, q) + } + if fLo, err = cdf(lo); err != nil { + return 0, base.Errf("%s: %w", name, err) + } + } + for fHi < q { + hi *= 2 + if math.IsInf(hi, 0) { + return 0, base.Errf("%s: failed to bracket q = %g from above", name, q) + } + if fHi, err = cdf(hi); err != nil { + return 0, base.Errf("%s: %w", name, err) + } + } + converged := false + for range 4096 { + mid := (lo + hi) / 2 + if mid == lo || mid == hi { + converged = true + break + } + f, err := cdf(mid) + if err != nil { + return 0, base.Errf("%s: %w", name, err) + } + if f < q { + lo = mid + } else { + hi = mid + } + } + if !converged { + return 0, base.Errf("%s: the bisection for q = %g did not converge", name, q) + } + return (lo + hi) / 2, nil +} diff --git a/stats/noncentral_test.go b/stats/noncentral_test.go new file mode 100644 index 0000000..211081d --- /dev/null +++ b/stats/noncentral_test.go @@ -0,0 +1,497 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +// TestNoncentralChiSquareClosed pins the noncentral χ² on the exact +// closed form its df = 1 corner carries: χ²(1, λ) is the square of a +// N(√λ, 1) draw, so the CDF is Φ(√x−√λ) − Φ(−√x−√λ), and on the λ = 0 +// reduction to the central law. +func TestNoncentralChiSquareClosed(t *testing.T) { + for _, lambda := range []float64{0.5, 1, 4, 9, 25} { + for _, x := range []float64{0.5, 1, 2, 4, 9, 16} { + got, err := NoncentralChiSquareCDF(x, 1, lambda) + if err != nil { + t.Fatalf("NoncentralChiSquareCDF(%g, 1, %g): %v", x, lambda, err) + } + root := math.Sqrt(lambda) + want := NormalCDF(math.Sqrt(x)-root) - NormalCDF(-math.Sqrt(x)-root) + if math.Abs(got-want) > 1e-13 { + t.Fatalf("NoncentralChiSquareCDF(%g, 1, %g) = %.16g, want %.16g", x, lambda, got, want) + } + } + } + for _, df := range []int{1, 2, 5, 10} { + for _, x := range []float64{0.5, 2, 7} { + got, err := NoncentralChiSquareCDF(x, df, 0) + if err != nil { + t.Fatalf("NoncentralChiSquareCDF(%g, %d, 0): %v", x, df, err) + } + want, err := ChiSquareCDF(x, df) + if err != nil || math.Abs(got-want) > 1e-14 { + t.Fatalf("λ = 0 reduction at df %d: %v vs %v (%v)", df, got, want, err) + } + } + } + if v, _ := NoncentralChiSquareCDF(-1, 3, 2); v != 0 { + t.Fatalf("CDF below the support = %v, want 0", v) + } + if v, _ := NoncentralChiSquareDensity(-1, 3, 2); v != 0 { + t.Fatalf("density below the support = %v, want 0", v) + } +} + +// TestNoncentralChiSquareDensityIntegral integrates the density against +// the CDF. The walk runs after the substitution x = s², which leaves +// 2s·f(s²) smooth at the origin for every df (the density itself +// behaves like x^{df/2−1} there, too flat a start for Simpson's error +// estimate on the lower degrees). Simpson must then reproduce the CDF +// to better than 1e-10 relative, the double route the closed forms +// cannot cover. +func TestNoncentralChiSquareDensityIntegral(t *testing.T) { + type grid struct { + df int + lambda float64 + x float64 + } + for _, g := range []grid{ + {2, 1, 6}, {3, 1, 6}, {3, 4, 10}, {5, 3, 12}, {10, 25, 60}, + } { + n := 200000 + s := math.Sqrt(g.x) + sh := s / float64(n) + f := func(sv float64) float64 { + if sv == 0 { + return 0 + } + xv := sv * sv + d, err := NoncentralChiSquareDensity(xv, g.df, g.lambda) + if err != nil { + t.Fatalf("NoncentralChiSquareDensity: %v", err) + } + return 2 * sv * d + } + sum := f(0) + f(s) + for i := 1; i < n; i++ { + w := 4.0 + if i%2 == 0 { + w = 2 + } + sum += w * f(float64(i)*sh) + } + integral := sum * sh / 3 + cdf, err := NoncentralChiSquareCDF(g.x, g.df, g.lambda) + if err != nil { + t.Fatalf("NoncentralChiSquareCDF: %v", err) + } + if rel := math.Abs(integral-cdf) / cdf; rel > 1e-10 { + t.Fatalf("df = %d, λ = %g, x = %g: density integral %.15g vs CDF %.15g (rel %g)", + g.df, g.lambda, g.x, integral, cdf, rel) + } + } +} + +// noncentralTOracle evaluates E[Φ(t√(V/ν) − δ)] for V ~ χ²(ν) by +// Simpson after the substitution V = u², which leaves the integrand +// smooth for every ν; the u = 0 limit is finite only for ν = 1. +func noncentralTOracle(t float64, df int, delta float64) float64 { + uhi := 40.0 + n := 200000 + h := uhi / float64(n) + logGammaB := func(a float64) float64 { l, _ := math.Lgamma(a); return l } + f := func(u float64) float64 { + if u == 0 { + if df == 1 { + return math.Sqrt(2/math.Pi) * NormalCDF(-delta) + } + return 0 + } + v := u * u + fv := 2 * u * math.Exp((float64(df)/2-1)*math.Log(v)-v/2-(float64(df)/2)*math.Ln2-logGammaB(float64(df)/2)) + return NormalCDF(t*math.Sqrt(v/float64(df))-delta) * fv + } + sum := f(0) + f(uhi) + for i := 1; i < n; i++ { + w := 4.0 + if i%2 == 0 { + w = 2 + } + sum += w * f(float64(i)*h) + } + return sum * h / 3 +} + +// TestNoncentralTAgainstQuadrature holds the Lenth series against a +// direct quadrature of E[Φ(t√(V/ν) − δ)], an independent route that +// shares no code with the series, at a grid spanning both signs of t +// and δ and degrees of freedom from 1 to 10. The large-δ cases are the +// underflow round: a noncentrality whose Poisson weight seed e^{−δ²/2} +// is below the double floor used to silence the whole series, and each +// of them once answered a silent 0. Tolerance 1e-9, an order above the +// quadrature's own accuracy. +func TestNoncentralTAgainstQuadrature(t *testing.T) { + for _, df := range []int{1, 2, 5, 10} { + for _, delta := range []float64{-3, -0.5, 0.5, 2} { + for _, tv := range []float64{-2, -0.5, 0.4, 1, 3} { + got, err := NoncentralTCDF(tv, df, delta) + if err != nil { + t.Fatalf("NoncentralTCDF(%g, %d, %g): %v", tv, df, delta, err) + } + want := noncentralTOracle(tv, df, delta) + if math.Abs(got-want) > 1e-9 { + t.Fatalf("NoncentralTCDF(%g, %d, %g) = %.15g, want quadrature %.15g", + tv, df, delta, got, want) + } + } + } + } + for _, c := range []struct { + tv, delta float64 + df int + }{ + {45, 45, 5}, + {50, 40, 3}, + {-45, -45, 5}, + {100, 45, 5}, + } { + got, err := NoncentralTCDF(c.tv, c.df, c.delta) + if err != nil { + t.Fatalf("NoncentralTCDF(%g, %d, %g): %v", c.tv, c.df, c.delta, err) + } + want := noncentralTOracle(c.tv, c.df, c.delta) + if math.Abs(got-want) > 1e-9 { + t.Fatalf("NoncentralTCDF(%g, %d, %g) = %.15g, want quadrature %.15g", + c.tv, c.df, c.delta, got, want) + } + } +} + +// TestNoncentralUnderflowSurvival pins the deep-noncentrality window +// where the mixture weights' raw seed underflows: the χ² CDF at its +// own mean answers a half, the far tail answers a genuinely computed +// negligible value rather than a silent zero, the F CDF saturates at +// 1 past the overflow of ν₁x, and the χ² density at the origin keeps +// the finite df = 2 limit. Past the term budget the refusal is +// explicit. +func TestNoncentralUnderflowSurvival(t *testing.T) { + atMean, err := NoncentralChiSquareCDF(2005, 5, 2000) + if err != nil { + t.Fatal(err) + } + if atMean < 0.48 || atMean > 0.52 { + t.Fatalf("NoncentralChiSquareCDF at the mean 2005 = %g, want near a half", atMean) + } + tail, err := NoncentralChiSquareCDF(1005, 5, 2000) + if err != nil { + t.Fatal(err) + } + if !(tail >= 0 && tail < 1e-30) { + t.Fatalf("NoncentralChiSquareCDF(1005, 5, 2000) = %g, want a negligible non-negative tail", tail) + } + for _, x := range []float64{math.MaxFloat64, math.Inf(1)} { + f, err := NoncentralFCDF(x, 4, 10, 3) + if err != nil { + t.Fatal(err) + } + if math.Abs(f-1) > 1e-15 { + t.Fatalf("NoncentralFCDF(%g, 4, 10, 3) = %g, want 1 to rounding", x, f) + } + } + d, err := NoncentralChiSquareDensity(0, 2, 3) + if err != nil { + t.Fatal(err) + } + if want := 0.5 * math.Exp(-1.5); d != want { + t.Fatalf("NoncentralChiSquareDensity(0, 2, 3) = %g, want the limit %g", d, want) + } + if _, err := NoncentralChiSquareCDF(10, 5, 4e5); err == nil || !strings.Contains(err.Error(), "budget") { + t.Fatalf("lambda 4e5: error = %v, want the budget refusal", err) + } + if _, err := NoncentralTCDF(10, 5, 500); err == nil || !strings.Contains(err.Error(), "budget") { + t.Fatalf("delta 500: error = %v, want the budget refusal", err) + } +} + +// TestNoncentralTIdentityReductions pins the exact corners: δ = 0 is +// the central Student t, t = 0 is Φ(−δ), and the two reflection +// identities of the law hold to rounding. +func TestNoncentralTIdentityReductions(t *testing.T) { + for _, df := range []int{1, 3, 8} { + for _, tv := range []float64{-4, -1, 0.3, 2} { + got, err := NoncentralTCDF(tv, df, 0) + if err != nil { + t.Fatalf("NoncentralTCDF(%g, %d, 0): %v", tv, df, err) + } + want, err := StudentTCDF(tv, df) + if err != nil || got != want { + t.Fatalf("δ = 0 reduction at (%g, %d): %v vs %v (%v)", tv, df, got, want, err) + } + } + } + for _, delta := range []float64{-3, -0.5, 1, 4} { + got, _ := NoncentralTCDF(0, 5, delta) + if want := NormalCDF(-delta); math.Abs(got-want) > 1e-15 { + t.Fatalf("NoncentralTCDF(0, 5, %g) = %.16g, want Φ(−δ) = %.16g", delta, got, want) + } + } + for _, delta := range []float64{-2, 1.5} { + for _, tv := range []float64{-1, 0.7, 2} { + pos, _ := NoncentralTCDF(tv, 4, delta) + reflected, _ := NoncentralTCDF(-tv, 4, -delta) + if math.Abs(pos+reflected-1) > 1e-14 { + t.Fatalf("reflection broken at t = %g, δ = %g: %.17g", tv, delta, pos+reflected) + } + mirrored, _ := NoncentralTCDF(-tv, 4, delta) + flipped, _ := NoncentralTCDF(tv, 4, -delta) + if math.Abs(flipped-(1-mirrored)) > 1e-14 { + t.Fatalf("sign symmetry broken at t = %g, δ = %g: %.17g vs %.17g", + tv, delta, flipped, 1-mirrored) + } + } + } +} + +// TestNoncentralFIdentityReductions pins the noncentral F on its λ = 0 +// central reduction and on the df₁ = 1 identity with the noncentral t: +// F(1, ν, λ) is the squared t(ν, √λ), so P(F ≤ y) is the t CDF across +// ±√y. The CDF must also fall as the noncentrality grows. +func TestNoncentralFIdentityReductions(t *testing.T) { + for _, df2 := range []int{1, 4, 12} { + for _, x := range []float64{0.3, 1, 2.5} { + got, err := NoncentralFCDF(x, 1, df2, 0) + if err != nil { + t.Fatalf("NoncentralFCDF(%g, 1, %d, 0): %v", x, df2, err) + } + want, err := BetaIncomplete(x/(x+float64(df2)), 0.5, float64(df2)/2) + if err != nil || math.Abs(got-want) > 1e-14 { + t.Fatalf("central reduction at x = %g, df₂ = %d: %v vs %v (%v)", x, df2, got, want, err) + } + } + } + for _, lambda := range []float64{1, 4} { + for _, df2 := range []int{2, 6} { + for _, y := range []float64{0.5, 2, 6} { + got, err := NoncentralFCDF(y, 1, df2, lambda) + if err != nil { + t.Fatalf("NoncentralFCDF(%g, 1, %d, %g): %v", y, df2, lambda, err) + } + root := math.Sqrt(lambda) + hi, _ := NoncentralTCDF(math.Sqrt(y), df2, root) + lo, _ := NoncentralTCDF(-math.Sqrt(y), df2, root) + if math.Abs(got-(hi-lo)) > 1e-13 { + t.Fatalf("t² identity at y = %g, ν = %d, λ = %g: %.15g vs %.15g", + y, df2, lambda, got, hi-lo) + } + } + } + } + central, _ := NoncentralFCDF(2, 3, 8, 0) + shifted, _ := NoncentralFCDF(2, 3, 8, 5) + if shifted >= central { + t.Fatalf("a larger λ lowered the CDF from %g to %g", central, shifted) + } +} + +// TestNoncentralQuantileRoundTrips inverts each noncentral CDF and +// checks the CDF at the quantile returns q. +func TestNoncentralQuantileRoundTrips(t *testing.T) { + for _, q := range []float64{0.01, 0.25, 0.5, 0.9, 0.99} { + x, err := NoncentralChiSquareQuantile(q, 5, 3) + if err != nil { + t.Fatalf("NoncentralChiSquareQuantile(%g): %v", q, err) + } + back, _ := NoncentralChiSquareCDF(x, 5, 3) + if math.Abs(back-q) > 1e-10 { + t.Fatalf("χ² round trip q = %g: CDF(quantile) = %v", q, back) + } + tq, err := NoncentralTQuantile(q, 5, 2) + if err != nil { + t.Fatalf("NoncentralTQuantile(%g): %v", q, err) + } + tback, _ := NoncentralTCDF(tq, 5, 2) + if math.Abs(tback-q) > 1e-10 { + t.Fatalf("t round trip q = %g: CDF(quantile) = %v", q, tback) + } + fq, err := NoncentralFQuantile(q, 4, 10, 2) + if err != nil { + t.Fatalf("NoncentralFQuantile(%g): %v", q, err) + } + fback, _ := NoncentralFCDF(fq, 4, 10, 2) + if math.Abs(fback-q) > 1e-10 { + t.Fatalf("F round trip q = %g: CDF(quantile) = %v", q, fback) + } + } + // The t quantile leans towards δ. + neg, _ := NoncentralTQuantile(0.5, 5, -2) + if neg >= 0 { + t.Fatalf("median of t(5, −2) = %g, want negative", neg) + } +} + +// TestNoncentralFMonteCarlo rebuilds the noncentral F from the +// package's own samplers: a χ²(df₁+2J) numerator with J drawn from the +// Poisson, over an independent central χ²(df₂) denominator, checked +// against the analytic CDF the mixture code computes. The sampler and +// the mixture share no code. Statistical tolerance 0.01, far above the +// 2σ of 50 000 draws. +func TestNoncentralFMonteCarlo(t *testing.T) { + g := core.NewGenerator(11) + const n = 50000 + x := 2.0 + got, err := NoncentralFCDF(x, 4, 10, 3) + if err != nil { + t.Fatalf("NoncentralFCDF: %v", err) + } + j, err := PoissonDraws(g, n, 1.5) + if err != nil { + t.Fatalf("PoissonDraws: %v", err) + } + count := 0.0 + for i := range n { + df := min( + // The Poisson(1.5) tail never reaches 30; the fold is a + // contract guard, not a working branch. + 4+2*int(j.FloatAt(i)), 64) + num, err := ChiSquareDraws(g, 1, df) + if err != nil { + t.Fatalf("ChiSquareDraws: %v", err) + } + den, err := ChiSquareDraws(g, 1, 10) + if err != nil { + t.Fatalf("ChiSquareDraws: %v", err) + } + f := num.FloatAt(0) * 10 / (4 * den.FloatAt(0)) + if f <= x { + count++ + } + } + if math.Abs(count/n-got) > 0.01 { + t.Fatalf("sampler CDF = %.4f, analytic %.4f", count/n, got) + } +} + +// TestNoncentralChiSquareMonteCarlo rebuilds the noncentral χ² from +// PoissonDraws and ChiSquareDraws, the sampler route the mixture CDF +// has no code in common with, at a tolerance the 200 000 draws can +// carry. +func TestNoncentralChiSquareMonteCarlo(t *testing.T) { + g := core.NewGenerator(13) + const n = 200000 + x := 6.0 + got, err := NoncentralChiSquareCDF(x, 3, 4) + if err != nil { + t.Fatalf("NoncentralChiSquareCDF: %v", err) + } + j, err := PoissonDraws(g, n, 2) + if err != nil { + t.Fatalf("PoissonDraws: %v", err) + } + // One shared χ²(3) stream rescaled per draw would not follow + // χ²(3+2J), so the check walks the mixture identity the other way: + // P(χ²(3+2J) ≤ x) averaged over the drawn J equals the CDF. + total := 0.0 + for i := range n { + df := min(3+2*int(j.FloatAt(i)), 400) + p, err := ChiSquareCDF(x, df) + if err != nil { + t.Fatalf("ChiSquareCDF: %v", err) + } + total += p + } + if math.Abs(total/n-got) > 0.01 { + t.Fatalf("sampler-route CDF = %v, want ≈ %v", total/n, got) + } +} + +// TestNoncentralErrors pins the parameter contracts of the noncentral +// family. +func TestNoncentralErrors(t *testing.T) { + if _, err := NoncentralChiSquareCDF(1, 0, 1); err == nil { + t.Fatal("df = 0: want an error") + } + if _, err := NoncentralChiSquareCDF(1, 3, -1); err == nil { + t.Fatal("negative λ: want an error") + } + if _, err := NoncentralChiSquareCDF(1, 3, math.Inf(1)); err == nil { + t.Fatal("λ = +Inf: want an error") + } + if _, err := NoncentralChiSquareCDF(math.NaN(), 3, 1); err == nil { + t.Fatal("NaN x: want an error") + } + if _, err := NoncentralChiSquareDensity(1, 0, 1); err == nil { + t.Fatal("density df = 0: want an error") + } + if _, err := NoncentralChiSquareQuantile(0.5, 0, 1); err == nil { + t.Fatal("quantile df = 0: want an error") + } + if _, err := NoncentralChiSquareQuantile(1.5, 3, 1); err == nil { + t.Fatal("q outside [0, 1]: want an error") + } + if _, err := NoncentralChiSquareQuantile(0.5, 0, 1); err == nil { + t.Fatal("quantile df = 0: want an error") + } + if _, err := NoncentralChiSquareQuantile(0.5, 3, math.NaN()); err == nil { + t.Fatal("quantile NaN λ: want an error") + } + if v, err := NoncentralChiSquareCDF(math.Inf(1), 3, 2); err != nil || v != 1 { + t.Fatalf("CDF at +Inf = %v, %v, want 1", v, err) + } + if v, err := NoncentralChiSquareDensity(0, 1, 2); err != nil || !math.IsInf(v, 1) { + t.Fatalf("density at 0 with df = 1 = %v, %v, want +Inf", v, err) + } + if _, err := NoncentralChiSquareDensity(math.NaN(), 3, 2); err == nil { + t.Fatal("density NaN x: want an error") + } + if _, err := NoncentralFCDF(1, 0, 4, 1); err == nil { + t.Fatal("df1 = 0: want an error") + } + if _, err := NoncentralFCDF(1, 4, 0, 1); err == nil { + t.Fatal("df2 = 0: want an error") + } + if _, err := NoncentralFCDF(1, 4, 4, math.NaN()); err == nil { + t.Fatal("NaN λ: want an error") + } + if _, err := NoncentralFCDF(math.NaN(), 4, 4, 1); err == nil { + t.Fatal("NaN x: want an error") + } + if _, err := NoncentralFCDF(1, 4, 4, math.Inf(1)); err == nil { + t.Fatal("λ = +Inf: want an error") + } + if _, err := NoncentralFQuantile(0.5, 0, 4, 1); err == nil { + t.Fatal("quantile df1 = 0: want an error") + } + if _, err := NoncentralFQuantile(0.5, 4, 4, -1); err == nil { + t.Fatal("quantile negative λ: want an error") + } + if _, err := NoncentralFQuantile(0, 4, 4, 1); err == nil { + t.Fatal("q = 0: want an error") + } + if _, err := NoncentralTCDF(math.Inf(1), 3, 1); err == nil { + t.Fatal("t = +Inf: want an error") + } + if _, err := NoncentralTCDF(1, 3, math.NaN()); err == nil { + t.Fatal("NaN δ: want an error") + } + if _, err := NoncentralTQuantile(0.5, 0, 1); err == nil { + t.Fatal("quantile df = 0: want an error") + } + if _, err := NoncentralTQuantile(0.5, 3, math.Inf(-1)); err == nil { + t.Fatal("δ = −Inf: want an error") + } + if _, err := NoncentralTQuantile(1.5, 3, 1); err == nil { + t.Fatal("q above 1: want an error") + } + if _, err := NoncentralTQuantile(0, 3, 1); err == nil { + t.Fatal("q = 0: want an error") + } + if _, err := NoncentralTQuantile(1, 3, 1); err == nil { + t.Fatal("q = 1: want an error") + } +} diff --git a/stats/pca.go b/stats/pca.go new file mode 100644 index 0000000..37372aa --- /dev/null +++ b/stats/pca.go @@ -0,0 +1,413 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "cmp" + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Principal component analysis: the rotation of a cloud of +// observations onto the axes along which it actually spreads, with the +// variances of those axes, the coordinates of every observation on +// them, and the whitening transform that flattens the cloud to unit +// covariance. The covariance comes from the package's own +// CovarianceMatrix; its eigendecomposition is computed here, because +// the package's linear algebra has so far only needed LU solves, and a +// symmetric eigensolver is carried as a cyclic Jacobi iteration: +// unconditionally convergent on symmetric matrices, quadratically so +// near the answer, and exactly the right instrument for the small +// dense covariance matrices a PCA runs on. + +// PCAResult carries the decomposition of an observation array. +type PCAResult struct { + // Mean holds the column means of the observations the fit ran on. + Mean []float64 + // Loadings is the (p, p) rotation: entry (j, k) is the loading of + // variable j on component k, columns ordered by falling explained + // variance, orthonormal as columns. Every column's largest-magnitude + // loading is positive, the first index winning a tie, so the axes + // have a fixed orientation and two runs on the same data agree + // sign for sign. + Loadings *core.Array + // Scores is the (n, p) array of the observations' coordinates on + // the components: the centred observations times the loadings. + Scores *core.Array + // ExplainedVariance holds each component's eigenvalue of the + // covariance, in the loadings' order; ExplainedVarianceRatio + // divides it by the total variance, so the ratios sum to one. + ExplainedVariance []float64 + ExplainedVarianceRatio []float64 + // Whitening and Unwhitening are the (p, p) transforms between the + // centred observations and the unit-covariance representation: + // whitening maps a centred row x to x·Whitening, unwhitening maps + // it back. Both are nil when the covariance is rank deficient, + // where no whitening transform exists. + Whitening *core.Array + Unwhitening *core.Array +} + +// PCA decomposes the observations in a, an (n, p) array whose rows are +// observations and columns variables, onto their principal +// components. At least two observations are needed, the input must be +// real and finite, and data with no variance at all has no +// decomposition to report. A rank-deficient covariance decomposes +// normally, with zero variances on the empty components; only the +// whitening transforms, which would divide by those variances, are +// withheld. +func PCA(a *core.Array) (*PCAResult, error) { + const name = "PCA" + covArr, err := CovarianceMatrix(a) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + n, p := a.Shape()[0], a.Shape()[1] + // The covariance is read once through its payload where it has one; + // the values are the ones FloatAt returns. + covVals := rawFloats(covArr) + cov := make([][]float64, p) + total := 0.0 + for i := range p { + cov[i] = make([]float64, p) + if covVals != nil { + copy(cov[i], covVals[i*p:i*p+p]) + } else { + for j := range p { + cov[i][j] = covArr.FloatAt(i*p + j) + } + } + total += cov[i][i] + } + if total == 0 { + return nil, base.Errf("%s: the observations have no variance", name) + } + values, vectors := jacobiEigen(cov) + // A covariance is positive semidefinite by construction, so a + // negative eigenvalue beyond a rounding-scale fraction of the + // largest means the decomposition cannot be trusted; anything + // within it is rounding and clamps to an exactly empty component. + worst := 0.0 + for _, v := range values { + worst = max(worst, math.Abs(v)) + } + for k, v := range values { + if v < 0 { + if v < -1e-9*worst { + return nil, base.Errf("%s: the covariance decomposed to the negative eigenvalue %g, not a covariance", name, v) + } + values[k] = 0 + } + } + // Column means, then the scores of every observation. + means := make([]float64, p) + rows := rawFloats(a) + for j := range p { + s := 0.0 + if rows != nil { + for i := range n { + s += rows[i*p+j] + } + } else { + for i := range n { + s += a.FloatAt(i*p + j) + } + } + means[j] = s / float64(n) + } + scores := make([]float64, 0, n*p) + for i := range n { + for k := range p { + s := 0.0 + vk := vectors[k] + for j := range p { + var xj float64 + if rows != nil { + xj = rows[i*p+j] + } else { + xj = a.FloatAt(i*p + j) + } + s += (xj - means[j]) * vk[j] + } + scores = append(scores, s) + } + } + // The loadings land in the documented (variable, component) + // layout: entry (j, k) is component k's loading on variable j. + loadings := make([]float64, 0, p*p) + for j := range p { + for k := range p { + loadings = append(loadings, vectors[k][j]) + } + } + out := &PCAResult{ + Mean: means, + Loadings: floatsToArray(loadings, []int{p, p}), + Scores: floatsToArray(scores, []int{n, p}), + ExplainedVariance: values, + ExplainedVarianceRatio: make([]float64, p), + } + for k := range p { + out.ExplainedVarianceRatio[k] = values[k] / total + } + if values[p-1] > 0 { + whitening := make([]float64, 0, p*p) + unwhitening := make([]float64, 0, p*p) + for i := range p { + for k := range p { + whitening = append(whitening, vectors[k][i]/math.Sqrt(values[k])) + } + } + for k := range p { + for j := range p { + unwhitening = append(unwhitening, math.Sqrt(values[k])*vectors[k][j]) + } + } + out.Whitening = floatsToArray(whitening, []int{p, p}) + out.Unwhitening = floatsToArray(unwhitening, []int{p, p}) + } + return out, nil +} + +// Whiten maps the observations in x, an (n, p) array on the fit's own +// variables, to their unit-covariance representation: the centred +// observations expressed on the components and scaled by each +// component's standard deviation. The result has the identity as its +// covariance, up to the fit's rank: a rank-deficient fit refuses to +// whiten rather than divide by an empty component's variance. +func (r *PCAResult) Whiten(x *core.Array) (*core.Array, error) { + const name = "Whiten" + if r == nil { + return nil, base.Errf("%s: no fit to whiten with", name) + } + if r.Whitening == nil { + return nil, base.Errf("%s: the fit's covariance is rank deficient and no whitening transform exists", name) + } + if x.NDim() != 2 { + return nil, base.Errf("%s: the observations must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if x.Dtype() == core.Complex { + return nil, base.Errf("%s: complex observations are not supported", name) + } + if err := checkFinite(name, "the observations", x); err != nil { + return nil, err + } + n, p := x.Shape()[0], len(r.Mean) + if x.Shape()[1] != p { + return nil, base.Errf("%s: the observations have %d columns, the fit %d", name, x.Shape()[1], p) + } + rows := rawFloats(x) + out := core.New(core.Float, n, p) + vals := out.RawFloats() + // The transform is read once into a plain slice: the values are the + // ones FloatAt returns. + transform := make([]float64, p*p) + if fs := rawFloats(r.Whitening); fs != nil { + copy(transform, fs) + } else { + for i := range transform { + transform[i] = r.Whitening.FloatAt(i) + } + } + for i := range n { + for k := range p { + s := 0.0 + for j := range p { + var xj float64 + if rows != nil { + xj = rows[i*p+j] + } else { + xj = x.FloatAt(i*p + j) + } + s += (xj - r.Mean[j]) * transform[j*p+k] + } + vals[i*p+k] = s + } + } + return out, nil +} + +// Unwhiten inverts Whiten: it maps a unit-covariance representation +// back to the observations' own coordinates, the observations +// themselves recovered exactly to rounding. +func (r *PCAResult) Unwhiten(z *core.Array) (*core.Array, error) { + const name = "Unwhiten" + if r == nil { + return nil, base.Errf("%s: no fit to unwhiten with", name) + } + if r.Unwhitening == nil { + return nil, base.Errf("%s: the fit's covariance is rank deficient and no unwhitening transform exists", name) + } + if z.NDim() != 2 { + return nil, base.Errf("%s: the whitened observations must be rank 2, got shape %s", name, base.ShapeText(z.Shape())) + } + if z.Dtype() == core.Complex { + return nil, base.Errf("%s: complex observations are not supported", name) + } + if err := checkFinite(name, "the whitened observations", z); err != nil { + return nil, err + } + n, p := z.Shape()[0], len(r.Mean) + if z.Shape()[1] != p { + return nil, base.Errf("%s: the whitened observations have %d columns, the fit %d", name, z.Shape()[1], p) + } + rows := rawFloats(z) + out := core.New(core.Float, n, p) + vals := out.RawFloats() + // The transform is read once into a plain slice: the values are the + // ones FloatAt returns. + transform := make([]float64, p*p) + if fs := rawFloats(r.Unwhitening); fs != nil { + copy(transform, fs) + } else { + for i := range transform { + transform[i] = r.Unwhitening.FloatAt(i) + } + } + for i := range n { + for j := range p { + s := 0.0 + for k := range p { + var zk float64 + if rows != nil { + zk = rows[i*p+k] + } else { + zk = z.FloatAt(i*p + k) + } + s += zk * transform[k*p+j] + } + vals[i*p+j] = r.Mean[j] + s + } + } + return out, nil +} + +// jacobiEigen eigendecomposes the symmetric matrix a by the cyclic +// Jacobi iteration: every sweep applies the plane rotation that zeroes +// each off-diagonal entry in turn, accumulating the rotations into the +// eigenvector matrix. The iteration is unconditionally convergent for +// symmetric matrices and quadratically convergent near the answer, so +// a handful of sweeps carries any small dense matrix to machine +// precision. The eigenvalues come back in falling order with the +// matching orthonormal eigenvectors as columns, each column oriented +// so its largest-magnitude entry is positive, the first index winning +// a tie. +func jacobiEigen(a [][]float64) (values []float64, vectors [][]float64) { + p := len(a) + m := make([][]float64, p) + for i := range p { + m[i] = append([]float64(nil), a[i]...) + } + v := make([][]float64, p) + for i := range p { + v[i] = make([]float64, p) + v[i][i] = 1 + } + scale := 0.0 + for i := range p { + for j := i; j < p; j++ { + scale += m[i][j] * m[i][j] + } + } + for range 100 { + off := 0.0 + for i := range p { + for j := i + 1; j < p; j++ { + off += m[i][j] * m[i][j] + } + } + if math.Sqrt(off) <= 1e-14*math.Sqrt(scale) { + break + } + for q := 1; q < p; q++ { + for i := 0; i < q; i++ { + aiq := m[i][q] + if aiq == 0 { + continue + } + theta := (m[q][q] - m[i][i]) / (2 * aiq) + sign := 1.0 + if theta < 0 { + sign = -1 + } + t := sign / (math.Abs(theta) + math.Sqrt(theta*theta+1)) + c := 1 / math.Sqrt(t*t+1) + s := t * c + tangent := s / (1 + c) + aii := m[i][i] + aqq := m[q][q] + m[i][i] = aii - t*aiq + m[q][q] = aqq + t*aiq + m[i][q] = 0 + m[q][i] = 0 + for k := range p { + if k == i || k == q { + continue + } + aki := m[k][i] + akq := m[k][q] + m[k][i] = aki - s*(akq+tangent*aki) + m[i][k] = m[k][i] + m[k][q] = akq + s*(aki-tangent*akq) + m[q][k] = m[k][q] + } + for k := range p { + vki := v[k][i] + vkq := v[k][q] + v[k][i] = vki - s*(vkq+tangent*vki) + v[k][q] = vkq + s*(vki-tangent*vkq) + } + } + } + } + values = make([]float64, p) + for i := range p { + values[i] = m[i][i] + } + // Falling order, ties left in their original order. + order := make([]int, p) + for i := range p { + order[i] = i + } + slices.SortStableFunc(order, func(a, b int) int { + return cmp.Compare(values[b], values[a]) + }) + sortedValues := make([]float64, p) + sortedVectors := make([][]float64, p) + for k := range p { + sortedValues[k] = values[order[k]] + sortedVectors[k] = make([]float64, p) + for i := range p { + sortedVectors[k][i] = v[i][order[k]] + } + } + fixEigenSigns(sortedVectors) + return sortedValues, sortedVectors +} + +// fixEigenSigns orients every eigenvector, held as a row of v, so its +// largest-magnitude entry is positive, the first index winning a tie. +// An eigendecomposition is defined only up to each eigenvector's sign, +// and a decomposition that reports different signs on the same input +// twice would be no decomposition at all. +func fixEigenSigns(v [][]float64) { + for k := range v { + worst := 0.0 + index := 0 + for i := range v[k] { + if a := math.Abs(v[k][i]); a > worst { + worst = a + index = i + } + } + if v[k][index] < 0 { + for i := range v[k] { + v[k][i] = -v[k][i] + } + } + } +} diff --git a/stats/pca_test.go b/stats/pca_test.go new file mode 100644 index 0000000..97abf9a --- /dev/null +++ b/stats/pca_test.go @@ -0,0 +1,441 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// pcaFixture builds a seeded (n, 4) observation array with correlated +// columns of clearly different scales, the ordinary material a PCA +// runs on. +func pcaFixture(t *testing.T, n int, seed int64) *core.Array { + t.Helper() + g := core.NewGenerator(31) + x0 := make([]float64, n) + vals := make([]float64, 0, 4*n) + for i := range n { + x0[i] = g.NormalUnit() + } + for i := range n { + x1 := 0.8*x0[i] + 0.6*g.NormalUnit() + x2 := -0.5*x0[i] + g.NormalUnit() + x3 := 0.3 * g.NormalUnit() + vals = append(vals, x0[i], x1, x2, x3) + } + return mustFromFloats(t, vals, n, 4) +} + +// TestPCALongAxis puts a two-cluster anisotropic cloud under the +// decomposition: two Gaussian blobs strung along a thirty-degree axis +// must come back with the first component along that axis, the +// measured angle against the truth, and nearly all the variance on it. +func TestPCALongAxis(t *testing.T) { + const n = 80 + g := core.NewGenerator(29) + const theta = math.Pi / 6 + cos, sin := math.Cos(theta), math.Sin(theta) + vals := make([]float64, 0, 4*n) + for shift := 0.0; shift <= 10; shift += 10 { + for range n { + u := 5 * g.NormalUnit() + v := 0.5 * g.NormalUnit() + vals = append(vals, + u*cos-v*sin+shift*cos, + u*sin+v*cos+shift*sin) + } + } + design := mustFromFloats(t, vals, 2*n, 2) + res, err := PCA(design) + if err != nil { + t.Fatalf("PCA: %v", err) + } + // The angle of a component's axis, read off its two loadings and + // folded into (−π/2, π/2], where the fixed sign convention leaves + // it. Loadings entry (j, k) sits at j·p+k. + angleOf := func(k int) float64 { + a := math.Atan2(res.Loadings.FloatAt(1*2+k), res.Loadings.FloatAt(0*2+k)) + if a > math.Pi/2 { + a -= math.Pi + } + if a <= -math.Pi/2 { + a += math.Pi + } + return a + } + // Axis directions are defined only up to a half turn, so the + // short axis's angle is folded like the first before comparing. + first := angleOf(0) + second := angleOf(1) + secondWant := theta + math.Pi/2 + if secondWant > math.Pi/2 { + secondWant -= math.Pi + } + t.Logf("first axis at %.4f rad against %.4f, second at %.4f against %.4f, explaining %.4f of the variance", + first, theta, second, secondWant, res.ExplainedVarianceRatio[0]) + if math.Abs(first-theta) > 0.05 { + t.Fatalf("the first component sits at %.4f rad, want the long axis %.4f", first, theta) + } + if math.Abs(second-secondWant) > 0.05 { + t.Fatalf("the second component sits at %.4f rad, want the short axis %.4f", second, secondWant) + } + if res.ExplainedVarianceRatio[0] < 0.98 { + t.Fatalf("the long axis explains only %.4f of the variance", res.ExplainedVarianceRatio[0]) + } +} + +// TestPCAVariancesAndLoadings pins the spectral accounting: the +// eigenvalues sum to the covariance's trace, the ratios sum to one, +// they arrive in falling order, the loadings are orthonormal as +// columns, and the loadings and eigenvalues rebuild the covariance +// they came from. +func TestPCAVariancesAndLoadings(t *testing.T) { + const n, p = 60, 4 + a := pcaFixture(t, n, 31) + res, err := PCA(a) + if err != nil { + t.Fatalf("PCA: %v", err) + } + // The trace, computed here straight from the data. + means := make([]float64, p) + for j := range p { + s := 0.0 + for i := range n { + s += a.FloatAt(i*p + j) + } + means[j] = s / float64(n) + } + trace := 0.0 + for j := range p { + s := 0.0 + for i := range n { + d := a.FloatAt(i*p+j) - means[j] + s += d * d + } + trace += s / float64(n-1) + } + total := 0.0 + for k := range p { + total += res.ExplainedVariance[k] + } + if math.Abs(total-trace) > 1e-10*math.Max(1, trace) { + t.Fatalf("the eigenvalues sum to %.12g, want the trace %.12g", total, trace) + } + ratioSum := 0.0 + for k := range p { + ratioSum += res.ExplainedVarianceRatio[k] + if k > 0 && res.ExplainedVariance[k] > res.ExplainedVariance[k-1] { + t.Fatalf("the variances are not descending at %d", k) + } + } + if math.Abs(ratioSum-1) > 1e-12 { + t.Fatalf("the ratios sum to %.16g, want 1", ratioSum) + } + // Orthonormal columns: LᵀL is the identity. + for j := range p { + for k := j; k < p; k++ { + dot := 0.0 + for i := range p { + dot += res.Loadings.FloatAt(i*p+j) * res.Loadings.FloatAt(i*p+k) + } + want := 0.0 + if j == k { + want = 1 + } + if math.Abs(dot-want) > 1e-10 { + t.Fatalf("loadings %d and %d have inner product %.4g, want %.4g", j, k, dot, want) + } + } + } + // The covariance rebuilds: L·D·Lᵀ against the entries the package + // computed. + cov, err := CovarianceMatrix(a) + if err != nil { + t.Fatalf("CovarianceMatrix: %v", err) + } + for i := range p { + for j := range p { + s := 0.0 + for k := range p { + s += res.ExplainedVariance[k] * res.Loadings.FloatAt(i*p+k) * res.Loadings.FloatAt(j*p+k) + } + if math.Abs(s-cov.FloatAt(i*p+j)) > 1e-9*math.Max(1, math.Abs(cov.FloatAt(i*p+j))) { + t.Fatalf("the rebuilt covariance entry (%d, %d) is %.12g, want %.12g", i, j, s, cov.FloatAt(i*p+j)) + } + } + } +} + +// TestPCAScoresWhitenRoundTrip pins the transforms: the scores are +// the centred observations times the loadings, the whitened data is +// the scores scaled by the components' standard deviations, the +// covariance of the whitened data is the identity, and unwhitening +// returns the centred observations. +func TestPCAScoresWhitenRoundTrip(t *testing.T) { + const n, p = 60, 4 + a := pcaFixture(t, n, 31) + res, err := PCA(a) + if err != nil { + t.Fatalf("PCA: %v", err) + } + // Scores against their definition. + for i := range n { + for k := range p { + s := 0.0 + for j := range p { + s += (a.FloatAt(i*p+j) - res.Mean[j]) * res.Loadings.FloatAt(j*p+k) + } + if math.Abs(s-res.Scores.FloatAt(i*p+k)) > 1e-9 { + t.Fatalf("score (%d, %d) is %.12g, want the centred row times the loading %.12g", + i, k, res.Scores.FloatAt(i*p+k), s) + } + } + } + // Whitened against the scores scaled by the component scales, and + // with identity covariance. + z, err := res.Whiten(a) + if err != nil { + t.Fatalf("Whiten: %v", err) + } + for i := range n { + for k := range p { + want := res.Scores.FloatAt(i*p+k) / math.Sqrt(res.ExplainedVariance[k]) + if math.Abs(z.FloatAt(i*p+k)-want) > 1e-8 { + t.Fatalf("whitened (%d, %d) is %.12g, want the scaled score %.12g", + i, k, z.FloatAt(i*p+k), want) + } + } + } + zcov, err := CovarianceMatrix(z) + if err != nil { + t.Fatalf("CovarianceMatrix of the whitened data: %v", err) + } + for i := range p { + for j := range p { + want := 0.0 + if i == j { + want = 1 + } + if math.Abs(zcov.FloatAt(i*p+j)-want) > 1e-9 { + t.Fatalf("the whitened covariance (%d, %d) is %.6g, want %.6g", i, j, zcov.FloatAt(i*p+j), want) + } + } + } + // Unwhitening returns the observations themselves: the round trip + // closes exactly. + back, err := res.Unwhiten(z) + if err != nil { + t.Fatalf("Unwhiten: %v", err) + } + for i := range n { + for j := range p { + if math.Abs(back.FloatAt(i*p+j)-a.FloatAt(i*p+j)) > 1e-9 { + t.Fatalf("the round trip returned %.12g at (%d, %d), want the observation %.12g", + back.FloatAt(i*p+j), i, j, a.FloatAt(i*p+j)) + } + } + } +} + +// TestPCATransposedSpectrum pins the transform consistency across the +// transposed problem: for a square matrix centred along both axes the +// covariance of the rows and the covariance of the columns are the +// Gram pair XXᵀ and XᵀX, which share their spectrum exactly, and the +// decomposition must see the same eigenvalues from either side. +func TestPCATransposedSpectrum(t *testing.T) { + const n = 6 + g := core.NewGenerator(37) + vals := make([]float64, 0, n*n) + for range n * n { + vals = append(vals, g.NormalUnit()) + } + // Centre along the columns and then along the rows, so both + // readings of the matrix describe the same centred scatter. + for i := range n { + mean := 0.0 + for j := range n { + mean += vals[i*n+j] + } + mean /= float64(n) + for j := range n { + vals[i*n+j] -= mean + } + } + for j := range n { + mean := 0.0 + for i := range n { + mean += vals[i*n+j] + } + mean /= float64(n) + for i := range n { + vals[i*n+j] -= mean + } + } + a := mustFromFloats(t, vals, n, n) + transposed := make([]float64, 0, n*n) + for i := range n { + for j := range n { + transposed = append(transposed, vals[j*n+i]) + } + } + at := mustFromFloats(t, transposed, n, n) + res, err := PCA(a) + if err != nil { + t.Fatalf("PCA: %v", err) + } + resT, err := PCA(at) + if err != nil { + t.Fatalf("PCA of the transposed data: %v", err) + } + for k := range n { + if math.Abs(res.ExplainedVariance[k]-resT.ExplainedVariance[k]) > 1e-8 { + t.Fatalf("eigenvalue %d: %.10g from the rows, %.10g from the columns", + k, res.ExplainedVariance[k], resT.ExplainedVariance[k]) + } + } +} + +// TestPCASignConvention pins the orientation rule directly on the +// helper: every eigenvector row turns so its largest-magnitude entry +// is positive, the first index winning a tie, and a row already +// oriented stays untouched. +func TestPCASignConvention(t *testing.T) { + // Row 0 ties at 0.6 across indices 0 and 1, index 0 negative: the + // first index wins, so the row flips. Row 1's largest entry is + // −0.9: it flips. Row 2's largest entry is 0.7: it stays. + v := [][]float64{ + {-0.6, 0.6, 0.1}, + {0.2, -0.9, 0.4}, + {0.1, 0.7, -0.2}, + } + fixEigenSigns(v) + want := [][]float64{ + {0.6, -0.6, -0.1}, + {-0.2, 0.9, -0.4}, + {0.1, 0.7, -0.2}, + } + for i := range 3 { + if !slices.Equal(v[i], want[i]) { + t.Fatalf("orientation wrong at row %d: %v, want %v", i, v[i], want[i]) + } + } + // And on a real fit: every component's heaviest loading positive. + a := pcaFixture(t, 40, 31) + res, err := PCA(a) + if err != nil { + t.Fatalf("PCA: %v", err) + } + for k := range 4 { + worst, index := 0.0, 0 + for i := range 4 { + if magnitude := math.Abs(res.Loadings.FloatAt(i*4 + k)); magnitude > worst { + worst = magnitude + index = i + } + } + if res.Loadings.FloatAt(index*4+k) < 0 { + t.Fatalf("component %d is oriented against the convention", k) + } + } +} + +// TestPCAValidation refuses the inputs without a decomposition and +// withholds the whitening transforms where they do not exist. +func TestPCAValidation(t *testing.T) { + good := pcaFixture(t, 20, 31) + if _, err := PCA(mustFromFloats(t, []float64{1, 2, 3}, 3)); err == nil || !strings.Contains(err.Error(), "2-D array") { + t.Fatalf("a rank 1 array: got %v, want the rank refusal", err) + } + if _, err := PCA(mustFromFloats(t, []float64{1, 2}, 1, 2)); err == nil || !strings.Contains(err.Error(), "at least two observations") { + t.Fatalf("a single observation: got %v, want the observation floor refusal", err) + } + if _, err := PCA(core.New(core.Complex, 4, 2)); err == nil || !strings.Contains(err.Error(), "complex observations") { + t.Fatalf("complex observations: got %v, want the complex refusal", err) + } + if _, err := PCA(mustFromFloats(t, []float64{1, 2, math.NaN(), 4, 5, 6, 7, 8}, 4, 2)); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("non-finite observations: got %v, want the non-finite refusal", err) + } + constant := mustFromFloats(t, []float64{1, 2, 1, 2, 1, 2, 1, 2}, 4, 2) + if _, err := PCA(constant); err == nil || !strings.Contains(err.Error(), "no variance") { + t.Fatalf("data with no variance: got %v, want the variance refusal", err) + } + // A duplicated column: the decomposition stands, the whitening + // transforms do not exist and are withheld. + singular := mustFromFloats(t, []float64{ + 1, 1, 2, + 2, 2, 1, + 3, 3, 0, + 4, 4, 1, + 5, 5, 2, + 6, 6, 3, + }, 6, 3) + res, err := PCA(singular) + if err != nil { + t.Fatalf("PCA on a singular covariance: %v", err) + } + if res.Whitening != nil || res.Unwhitening != nil { + t.Fatalf("a rank-deficient fit published whitening transforms") + } + if _, err := res.Whiten(singular); err == nil || !strings.Contains(err.Error(), "rank deficient") { + t.Fatalf("Whiten on a rank-deficient fit: got %v, want the rank-deficiency refusal", err) + } + if _, err := res.Unwhiten(singular); err == nil || !strings.Contains(err.Error(), "rank deficient") { + t.Fatalf("Unwhiten on a rank-deficient fit: got %v, want the rank-deficiency refusal", err) + } + // Shape and content checks on the transforms. + fit, err := PCA(good) + if err != nil { + t.Fatalf("PCA: %v", err) + } + if _, err := fit.Whiten(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2)); err == nil || !strings.Contains(err.Error(), "columns, the fit") { + t.Fatalf("Whiten with a column mismatch: got %v, want the column refusal", err) + } + if _, err := fit.Whiten(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, math.NaN(), 9, 10, 11, 12, 13, 14, 15, 16}, 4, 4)); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("Whiten with non-finite observations: got %v, want the non-finite refusal", err) + } + if _, err := fit.Unwhiten(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2)); err == nil || !strings.Contains(err.Error(), "columns, the fit") { + t.Fatalf("Unwhiten with a column mismatch: got %v, want the column refusal", err) + } + var noFit *PCAResult + if _, err := noFit.Whiten(good); err == nil || !strings.Contains(err.Error(), "no fit to whiten") { + t.Fatalf("a nil fit whitened: got %v, want the nil-fit refusal", err) + } + if _, err := noFit.Unwhiten(good); err == nil || !strings.Contains(err.Error(), "no fit to unwhiten") { + t.Fatalf("a nil fit unwhitened: got %v, want the nil-fit refusal", err) + } + complexData := core.New(core.Complex, 4, 4) + if _, err := fit.Whiten(complexData); err == nil || !strings.Contains(err.Error(), "complex observations") { + t.Fatalf("Whiten with complex observations: got %v, want the complex refusal", err) + } + if _, err := fit.Unwhiten(complexData); err == nil || !strings.Contains(err.Error(), "complex observations") { + t.Fatalf("Unwhiten with complex observations: got %v, want the complex refusal", err) + } + // Integer observations reach the transforms through the widening + // accessor instead of a raw float payload, and whiten back out + // identically. Five rows of four generic columns keep the + // covariance full rank. + integers := mustFromInts(t, []int64{ + 1, 2, 3, 4, + 2, 4, 6, 3, + 3, 6, 2, 9, + 4, 3, 8, 1, + 5, 7, 1, 2, + }, 5, 4) + intFit, err := PCA(integers) + if err != nil { + t.Fatalf("PCA on integer observations: %v", err) + } + z, err := intFit.Whiten(integers) + if err != nil { + t.Fatalf("Whiten on integer observations: %v", err) + } + if _, err := intFit.Unwhiten(z); err != nil { + t.Fatalf("Unwhiten on integer observations: %v", err) + } +} diff --git a/stats/pins_test.go b/stats/pins_test.go new file mode 100644 index 0000000..b3c3520 --- /dev/null +++ b/stats/pins_test.go @@ -0,0 +1,1006 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Behaviour pins for the package's entry points: every test here +// holds a corrected defect to the answer it must keep, among them the +// NaN quantile hang, the rolling dtype panics, the Quantile NaN +// panic, the Histogram2D complex panic and unbounded bin counts, the +// lower-half NormalQuantile bracket, the saturated GLM likelihood, +// the weighted regression intercept statistics, the missing +// finiteness and symmetry refusals, the BinomialDraws count guard, +// the Kolmogorov series' truncation, the NaN-parameter and +// complex-input sweeps, the KolmogorovSmirnovTest NaN loop and the +// degenerate-response F statistic. Each one fails with the old +// behaviour restored. + +package stats + +import ( + "math" + "math/big" + "testing" + "time" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// pinVector builds a rank-1 float array. +func pinVector(t *testing.T, vals ...float64) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, len(vals)) + if err != nil { + t.Fatalf("FromFloats(%v): %v", vals, err) + } + return a +} + +// pinWatchdog runs f in its own goroutine and fails the test if it +// has not returned within limit. It is what keeps the hang regression a +// failure rather than a stuck test binary. +func pinWatchdog(t *testing.T, limit time.Duration, what string, f func()) { + t.Helper() + done := make(chan struct{}) + go func() { + defer close(done) + f() + }() + select { + case <-done: + case <-time.After(limit): + t.Fatalf("%s has not returned after %v: the loop does not terminate", what, limit) + } +} + +// TestDiscreteQuantileNaNQuantileDoesNotHang pins the P0: q = NaN passed +// the written-out range guard (`q < 0 || q > 1` is false for NaN), every +// bracket comparison below was then false too, and the doubling loop +// never terminated. Both discrete quantiles must refuse NaN instead. +func TestDiscreteQuantileNaNQuantileDoesNotHang(t *testing.T) { + for _, tc := range []struct { + name string + call func() (float64, error) + }{ + {"PoissonQuantile", func() (float64, error) { return PoissonQuantile(math.NaN(), 3) }}, + {"BinomialQuantile", func() (float64, error) { return BinomialQuantile(math.NaN(), 0.5, 5) }}, + } { + t.Run(tc.name, func(t *testing.T) { + var err error + pinWatchdog(t, 10*time.Second, tc.name+"(NaN)", func() { + _, err = tc.call() + }) + if err == nil { + t.Fatalf("%s(NaN) returned a value, want a refusal", tc.name) + } + }) + } +} + +// bigSqrt2 is sqrt(2) at the working precision. +func bigSqrt2() *big.Float { + const prec = 256 + return new(big.Float).SetPrec(prec).Sqrt(new(big.Float).SetPrec(prec).SetInt64(2)) +} + +// pinIntArray builds a rank-1 int array. +func pinIntArray(t *testing.T, vals ...int64) *core.Array { + t.Helper() + a, err := core.FromInts(vals, len(vals)) + if err != nil { + t.Fatalf("FromInts(%v): %v", vals, err) + } + return a +} + +// pinFloat32Array builds a rank-1 float32 array. +func pinFloat32Array(t *testing.T, vals ...float32) *core.Array { + t.Helper() + a, err := core.FromFloat32s(vals, len(vals)) + if err != nil { + t.Fatalf("FromFloat32s(%v): %v", vals, err) + } + return a +} + +// pinSameFloats pins two float slices to the bit for the dtype tests. +func pinSameFloats(t *testing.T, label string, got, want []float64) { + t.Helper() + if len(got) != len(want) { + t.Fatalf("%s has length %d, want %d", label, len(got), len(want)) + } + for i := range got { + if got[i] != want[i] { + t.Fatalf("%s[%d] = %.17g, want %.17g", label, i, got[i], want[i]) + } + } +} + +// TestRollingFamilyAcceptsIntAndFloat32 pins the P0: rollingCheck +// returned RawFloats()[:n], which is a nil slice with capacity 0 for an +// int or a float32 payload, so all four reductions panicked. The int, +// float32 and float routes must agree value for value. +func TestRollingFamilyAcceptsIntAndFloat32(t *testing.T) { + vals := []float64{1, 4, 2, 8, 5, 7} + floats := pinVector(t, vals...) + ints := pinIntArray(t, 1, 4, 2, 8, 5, 7) + f32 := pinFloat32Array(t, 1, 4, 2, 8, 5, 7) + cases := []struct { + name string + call func(a *core.Array) (*core.Array, error) + }{ + {"RollingMean", func(a *core.Array) (*core.Array, error) { return RollingMean(a, 3) }}, + {"RollingSum", func(a *core.Array) (*core.Array, error) { return RollingSum(a, 2) }}, + {"RollingMax", func(a *core.Array) (*core.Array, error) { return RollingMax(a, 3) }}, + {"RollingMin", func(a *core.Array) (*core.Array, error) { return RollingMin(a, 2) }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + want, err := tc.call(floats) + if err != nil { + t.Fatalf("float route: %v", err) + } + wantVals := make([]float64, want.Len()) + for i := range wantVals { + wantVals[i] = want.FloatAt(i) + } + for _, dt := range []struct { + name string + a *core.Array + }{{"int", ints}, {"float32", f32}} { + got, err := tc.call(dt.a) + if err != nil { + t.Fatalf("%s route: %v", dt.name, err) + } + gotVals := make([]float64, got.Len()) + for i := range gotVals { + gotVals[i] = got.FloatAt(i) + } + pinSameFloats(t, dt.name, gotVals, wantVals) + } + }) + } +} + +// TestRollingExtremesAllNaNWindow pins the P2: the window fold seeded +// its running best with ±Inf and only replaced it on a strict +// comparison, so an all-NaN window answered -Inf/+Inf where core.Min +// and core.Max answer NaN. A NaN that shares a window with a number +// must still lose. +func TestRollingExtremesAllNaNWindow(t *testing.T) { + allNaN := pinVector(t, math.NaN(), math.NaN()) + max, err := RollingMax(allNaN, 2) + if err != nil { + t.Fatalf("RollingMax: %v", err) + } + if !math.IsNaN(max.FloatAt(0)) { + t.Fatalf("RollingMax of an all-NaN window = %v, want NaN as core.Max gives", max.FloatAt(0)) + } + min, err := RollingMin(allNaN, 2) + if err != nil { + t.Fatalf("RollingMin: %v", err) + } + if !math.IsNaN(min.FloatAt(0)) { + t.Fatalf("RollingMin of an all-NaN window = %v, want NaN as core.Min gives", min.FloatAt(0)) + } + // A NaN never wins against a number. + mixed := pinVector(t, math.NaN(), 5, 2, math.NaN()) + wantMax := []float64{5, 5, 2} + wantMin := []float64{5, 2, 2} + gotMax, err := RollingMax(mixed, 2) + if err != nil { + t.Fatalf("RollingMax: %v", err) + } + gotMin, err := RollingMin(mixed, 2) + if err != nil { + t.Fatalf("RollingMin: %v", err) + } + for i := range wantMax { + if gotMax.FloatAt(i) != wantMax[i] { + t.Fatalf("RollingMax[%d] = %v, want %v (NaN never wins)", i, gotMax.FloatAt(i), wantMax[i]) + } + if gotMin.FloatAt(i) != wantMin[i] { + t.Fatalf("RollingMin[%d] = %v, want %v (NaN never wins)", i, gotMin.FloatAt(i), wantMin[i]) + } + } +} + +// TestQuantileRejectsNaNQuantile pins the P0: `q < 0 || q > 1` is +// false for NaN, so the quantile reached int(math.Floor(NaN)) and +// indexed the sorted values with it, panicking instead of refusing. +func TestQuantileRejectsNaNQuantile(t *testing.T) { + a := pinVector(t, 1, 2, 3, 4) + for _, qs := range [][]float64{{math.NaN()}, {0.5, math.NaN()}, {math.Inf(1)}, {math.Inf(-1)}} { + out, err := Quantile(a, qs) + if err == nil { + t.Fatalf("Quantile(%v) = %v, want a refusal", qs, out) + } + } + // The documented domain still works. + out, err := Quantile(a, []float64{0.25, 0.5, 0.75}) + if err != nil { + t.Fatalf("Quantile: %v", err) + } + pinSameFloats(t, "Quantile", []float64{out.FloatAt(0), out.FloatAt(1), out.FloatAt(2)}, + []float64{1.75, 2.5, 3.25}) +} + +// TestHistogram2DRejectsComplexInput pins the P0: every other entry +// point of the package refuses complex input by name, while Histogram2D +// walked its float accessor into a nil int payload and panicked. +func TestHistogram2DRejectsComplexInput(t *testing.T) { + real := pinVector(t, 1, 2, 3) + cx, err := core.FromComplexes([]complex128{1, 2, 3}, 3) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + if _, _, _, err := Histogram2D(cx, real, 2, 2); err == nil { + t.Fatal("complex x accepted") + } + if _, _, _, err := Histogram2D(real, cx, 2, 2); err == nil { + t.Fatal("complex y accepted") + } + // The real path is untouched. + counts, _, _, err := Histogram2D(real, real, 2, 2) + if err != nil { + t.Fatalf("Histogram2D: %v", err) + } + if counts.Shape()[0] != 2 || counts.Shape()[1] != 2 { + t.Fatalf("shape %v, want [2 2]", counts.Shape()) + } +} + +// TestHistogramBinCountsAreBounded pins the P1: neither Histogram nor +// Histogram2D bounded its bin count, so a hostile request either asked +// the allocator for an unbounded buffer or, with a product that wraps, +// produced an empty count payload the first write then indexed. The +// guard must fire before anything is allocated. +func TestHistogramBinCountsAreBounded(t *testing.T) { + a := pinVector(t, 0, 1) + if _, _, err := Histogram(a, maxHistBins+1); err == nil { + t.Fatal("a bin count past the limit was accepted") + } + if _, _, err := Histogram(a, 1<<21); err == nil { + t.Fatal("a 2^21-bin request was accepted") + } + counts, edges, err := Histogram(a, 4) + if err != nil { + t.Fatalf("a modest histogram: %v", err) + } + if counts.Len() != 4 || edges.Len() != 5 { + t.Fatalf("lengths %d and %d, want 4 and 5", counts.Len(), edges.Len()) + } + if _, _, _, err := Histogram2D(a, a, maxHistBins+1, 1); err == nil { + t.Fatal("a per-axis bin count past the limit was accepted") + } + // 1025 * 1024 = 1049600 is past the total, each axis alone is not. + if _, _, _, err := Histogram2D(a, a, 1025, 1024); err == nil { + t.Fatal("a bin product past the limit was accepted") + } + // A wide but legal request goes through: the guard refuses what is + // past the limit, not what is merely large. + if _, _, _, err := Histogram2D(a, a, 512, 512); err != nil { + t.Fatalf("a 512 x 512 histogram: %v", err) + } +} + +// TestNormalQuantileLowerTail pins the P1: continuousQuantile brackets +// upwards from the seed only, and Phi(0) = 0.5 is the infimum it can +// reach, so every q < 0.5 failed to bracket. The mirrored tail must +// reproduce an independent high-precision inverse of Phi. +func TestNormalQuantileLowerTail(t *testing.T) { + qs := []float64{1e-6, 1e-5, 1e-4, 1e-3, 0.01, 0.025, 0.05, 0.1, 0.25, 0.4, 0.4999} + worst := 0.0 + for _, q := range qs { + z, err := NormalQuantile(q) + if err != nil { + t.Fatalf("NormalQuantile(%g): %v", q, err) + } + if z >= 0 { + t.Fatalf("NormalQuantile(%g) = %v, want a negative quantile", q, z) + } + ref := bigNormalQuantile(q) + dev := math.Abs(z - ref) + if dev > worst { + worst = dev + } + if dev > 1e-9 { + t.Fatalf("NormalQuantile(%g) = %.17g, the 256-bit bisection reference is %.17g (off by %.3g)", + q, z, ref, dev) + } + // The symmetry is exact in the implementation and must hold to + // the last bit: the mirrored call returns the negated value. + mirror, err := NormalQuantile(1 - q) + if err != nil { + t.Fatalf("NormalQuantile(%g): %v", 1-q, err) + } + if z+mirror != 0 { + t.Fatalf("NormalQuantile(%g) + NormalQuantile(%g) = %.17g, want exactly 0", q, 1-q, z+mirror) + } + // And the library's own CDF inverts the result. + if dev := math.Abs(NormalCDF(z) - q); dev > 1e-12 { + t.Fatalf("NormalCDF(NormalQuantile(%g)) is off by %.3g", q, dev) + } + } + t.Logf("worst deviation from the 256-bit reference: %.3g", worst) + // The upper half was already correct and must stay so. + upper, err := NormalQuantile(0.975) + if err != nil { + t.Fatalf("NormalQuantile(0.975): %v", err) + } + if math.Abs(upper-1.9599639845400532) > 1e-12 { + t.Fatalf("NormalQuantile(0.975) = %.17g, want 1.9599639845400532", upper) + } +} + +// TestNormalQuantileNaNRejected pins the P2: NaN passed the written-out +// range guard and every comparison of the bisection returned false, so +// the seed came back as the answer (v = 1 with a nil error). +func TestNormalQuantileNaNRejected(t *testing.T) { + for _, q := range []float64{math.NaN(), math.Inf(1), math.Inf(-1), -0.5, 1.5, 0, 1} { + v, err := NormalQuantile(q) + if err == nil { + t.Fatalf("NormalQuantile(%v) = %v, want a refusal", q, v) + } + } + if v, err := StudentTQuantile(math.NaN(), 5); err == nil { + t.Fatalf("StudentTQuantile(NaN, 5) = %v, want a refusal", v) + } +} + +// TestLogisticRegressionSaturatedLikelihoodFinite pins the P1: the +// convergence exit recomputed the sigmoid without the clamp the loop +// relies on, so one far-out covariate saturated Fitted to exactly 0/1 +// and turned LogLikelihood into 0·log(0) = NaN while Converged was +// true. The reported values must be the ones the loop maximised. +func TestLogisticRegressionSaturatedLikelihoodFinite(t *testing.T) { + design, err := core.FromFloats([]float64{ + 1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, 1, 1e3, + }, 7, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + y := pinVector(t, 0, 1, 0, 1, 1, 0, 1) + res, err := LogisticRegression(design, y) + if err != nil { + t.Fatalf("LogisticRegression: %v", err) + } + if !res.Converged { + t.Fatal("the fit did not converge") + } + manual := 0.0 + for i, p := range res.Fitted { + if p <= 0 || p >= 1 { + t.Fatalf("Fitted[%d] = %v, want a probability strictly inside (0, 1)", i, p) + } + yv := y.FloatAt(i) + manual += yv*math.Log(p) + (1-yv)*math.Log(1-p) + } + if math.IsNaN(res.LogLikelihood) || math.IsInf(res.LogLikelihood, 0) { + t.Fatalf("LogLikelihood = %v, want the finite value the loop maximised", res.LogLikelihood) + } + if math.Abs(res.LogLikelihood-manual) > 1e-12 { + t.Fatalf("LogLikelihood = %.17g, recomputing from Fitted gives %.17g", + res.LogLikelihood, manual) + } + // The inference must stay finite as well. + for j := range res.Coefficients { + mustFinite := []float64{res.Coefficients[j], res.StandardErrors[j], + res.ZStatistics[j], res.PValues[j]} + for _, v := range mustFinite { + if math.IsNaN(v) || math.IsInf(v, 0) { + t.Fatalf("coefficient %d has a non-finite statistic: %v", j, mustFinite) + } + } + } +} + +// TestWeightedLinearRegressionInterceptStatistics pins the P1: the +// intercept column of the unweighted design stops being constant once +// the sqrt weights are applied, so the delegated fit reported the +// no-intercept conventions (uncentred TSS, DModel = p, an F test +// against the zero model) for a regression the caller described with an +// intercept. The statistics must be the weighted-theory ones, against +// the weighted mean. +func TestWeightedLinearRegressionInterceptStatistics(t *testing.T) { + design, err := core.FromFloats([]float64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5}, 6, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + yVals := []float64{10, 10.1, 9.9, 10, 10.2, 9.8} + y := pinVector(t, yVals...) + plain, err := LinearRegression(design, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + // Uniform weights are the ordinary fit, statistic for statistic. + uniform, err := WeightedLinearRegression(design, y, pinVector(t, 1, 1, 1, 1, 1, 1)) + if err != nil { + t.Fatalf("WeightedLinearRegression: %v", err) + } + if uniform.DModel != plain.DModel || uniform.RSquared != plain.RSquared || + uniform.FStatistic != plain.FStatistic || uniform.FPValue != plain.FPValue { + t.Fatalf("uniform weights report DModel=%d R2=%.17g F=%.17g p=%.17g, the ordinary fit DModel=%d R2=%.17g F=%.17g p=%.17g", + uniform.DModel, uniform.RSquared, uniform.FStatistic, uniform.FPValue, + plain.DModel, plain.RSquared, plain.FStatistic, plain.FPValue) + } + for _, w6 := range []float64{1.0000001, 2} { + w := pinVector(t, 1, 1, 1, 1, 1, w6) + res, err := WeightedLinearRegression(design, y, w) + if err != nil { + t.Fatalf("WeightedLinearRegression(w6=%g): %v", w6, err) + } + // The reference from first principles, on the residuals the call + // itself reports: sumW, the weighted mean, the weighted centred + // total and the weighted residual sum. + sumW, sumWY, rss := 0.0, 0.0, 0.0 + weights := []float64{1, 1, 1, 1, 1, w6} + for i, wi := range weights { + sumW += wi + sumWY += wi * yVals[i] + rss += wi * res.Residuals[i] * res.Residuals[i] + } + meanW := sumWY / sumW + tss := 0.0 + for i, wi := range weights { + d := yVals[i] - meanW + tss += wi * d * d + } + wantR2 := 1 - rss/tss + wantF := (tss - rss) / 1 / (rss / 4) + if res.DModel != 1 { + t.Fatalf("w6=%g: DModel = %d, want 1 (the design carries the intercept)", w6, res.DModel) + } + if math.Abs(res.RSquared-wantR2) > 1e-12 { + t.Fatalf("w6=%g: R2 = %.17g, the weighted-centred reference is %.17g", w6, res.RSquared, wantR2) + } + if math.Abs(res.FStatistic-wantF) > 1e-9*wantF { + t.Fatalf("w6=%g: F = %.17g, the weighted-centred reference is %.17g", w6, res.FStatistic, wantF) + } + // For a single slope F = t^2, so the model p-value is the + // slope's own two-sided tail. + if math.Abs(res.FPValue-res.PValues[1]) > 1e-9 { + t.Fatalf("w6=%g: model p = %.17g, the slope's two-sided tail is %.17g", + w6, res.FPValue, res.PValues[1]) + } + } + // The report's exact acceptance values on this data with w6 = 2, the + // weighted-centred values (R2 0.172938829787234, F 0.836401640003216, + // p 0.4121703772248073 from a numerical integration of F(1,4)): + // the defect reported 0.9998404595, 12534.0045019697 and + // 2.54531602909285e-08 instead. + res, err := WeightedLinearRegression(design, y, pinVector(t, 1, 1, 1, 1, 1, 2)) + if err != nil { + t.Fatalf("WeightedLinearRegression: %v", err) + } + if math.Abs(res.RSquared-0.172938829787234) > 1e-9 { + t.Fatalf("R2 = %.17g, the weighted-centred value is 0.172938829787234", res.RSquared) + } + if math.Abs(res.FStatistic-0.836401640003216) > 1e-9 { + t.Fatalf("F = %.17g, the weighted-centred value is 0.836401640003216", res.FStatistic) + } + if math.Abs(res.FPValue-0.4121703772248073) > 1e-9 { + t.Fatalf("model p = %.17g, the weighted-centred value is 0.4121703772248073", res.FPValue) + } + if res.DModel != 1 || res.DResidual != 4 { + t.Fatalf("degrees of freedom (%d, %d), want (1, 4)", res.DModel, res.DResidual) + } + // A single weight a hair away from uniform must not move the + // statistics: the flip the defect produced was 0.05 to 0.9998. + hair, err := WeightedLinearRegression(design, y, pinVector(t, 1, 1, 1, 1, 1, 1.0000001)) + if err != nil { + t.Fatalf("WeightedLinearRegression: %v", err) + } + if math.Abs(hair.RSquared-uniform.RSquared) > 1e-6 { + t.Fatalf("a 1e-7 weight change moved R2 from %.17g to %.17g", uniform.RSquared, hair.RSquared) + } + if math.Abs(hair.FPValue-uniform.FPValue) > 1e-6 { + t.Fatalf("a 1e-7 weight change moved the model p from %.17g to %.17g", + uniform.FPValue, hair.FPValue) + } +} + +// TestBinomialDrawsRejectsNonPositiveCount pins the P2: the only draw +// function in the file without the n < 1 guard returned a nil array +// with a nil error, which a caller checking only the error dereferences. +func TestBinomialDrawsRejectsNonPositiveCount(t *testing.T) { + g := core.NewGenerator(1) + for _, n := range []int{-1, 0} { + out, err := BinomialDraws(g, n, 4, 0.5) + if err == nil { + t.Fatalf("BinomialDraws(n=%d) = %v, want a refusal", n, out) + } + if out != nil { + t.Fatalf("BinomialDraws(n=%d) returned a non-nil array with the refusal", n) + } + } + out, err := BinomialDraws(g, 3, 4, 0.5) + if err != nil { + t.Fatalf("BinomialDraws: %v", err) + } + if out.Len() != 3 { + t.Fatalf("length %d, want 3", out.Len()) + } +} + +// TestKolmogorovSmirnovRejectsNonFiniteSamples pins an unbounded loop +// found while verifying the package's finiteness policy: the merge walk +// in KolmogorovSmirnovTest advances only past values that compare true +// against the current support point, and a NaN never does, so the walk +// re-reads the same element for ever. A +Inf sample walks through but +// distorts the distance. The sibling tests (MannWhitneyU, ANOVAOneWay) +// refuse a non-finite sample by name and so must this one. +func TestKolmogorovSmirnovRejectsNonFiniteSamples(t *testing.T) { + clean := pinVector(t, 1, 2, 3, 4, 5) + nan := pinVector(t, 1, 2, math.NaN(), 4, 5) + inf := pinVector(t, 1, 2, math.Inf(1), 4, 5) + for _, tc := range []struct { + name string + a, b *core.Array + }{ + {"NaN first sample", nan, clean}, + {"NaN second sample", clean, nan}, + {"+Inf first sample", inf, clean}, + {"+Inf second sample", clean, inf}, + } { + t.Run(tc.name, func(t *testing.T) { + var err error + pinWatchdog(t, 5*time.Second, "KolmogorovSmirnovTest("+tc.name+")", func() { + _, _, err = KolmogorovSmirnovTest(tc.a, tc.b) + }) + if err == nil { + t.Fatal("a non-finite sample was accepted") + } + }) + } + // The finite path is untouched: identical samples are distance zero + // with p = 1, and the report's {1,2,3} against {1.5,2.5,3.5} case + // gives d = 1/3 exactly. + d, p, err := KolmogorovSmirnovTest(clean, clean) + if err != nil { + t.Fatalf("KolmogorovSmirnovTest: %v", err) + } + if d != 0 || p != 1 { + t.Fatalf("identical samples give d = %g, p = %g, want 0 and 1", d, p) + } + d, _, err = KolmogorovSmirnovTest(pinVector(t, 1, 2, 3), pinVector(t, 1.5, 2.5, 3.5)) + if err != nil { + t.Fatalf("KolmogorovSmirnovTest: %v", err) + } + // The largest gap is the first support point's 1/3, computed from + // integer counts: allow the couple of ulps that division costs. + if math.Abs(d-1.0/3.0) > 1e-15 { + t.Fatalf("d = %.17g, want 1/3", d) + } +} + +// TestKernelDensityRejectsNonFiniteInput pins the first half of the +// finiteness gap: KernelDensity never scanned the sample or the +// evaluation points, so a NaN came back as a NaN estimate with a nil +// error and a +Inf silently dropped that sample from the average (the +// estimate was then the one of the remaining n-1 points). The clean +// path is pinned against the hand-written KDE definition. +func TestKernelDensityRejectsNonFiniteInput(t *testing.T) { + sample := pinVector(t, 0, 1, 2, 3) + points := pinVector(t, 0.5, 2.5) + for _, tc := range []struct { + name string + s, p *core.Array + }{ + {"NaN sample", pinVector(t, 0, 1, math.NaN(), 3), points}, + {"+Inf sample", pinVector(t, 0, 1, math.Inf(1), 3), points}, + {"NaN point", sample, pinVector(t, 0.5, math.NaN())}, + {"-Inf point", sample, pinVector(t, 0, math.Inf(-1))}, + } { + t.Run(tc.name, func(t *testing.T) { + if out, err := KernelDensity(tc.s, 0.5, tc.p); err == nil { + t.Fatalf("a non-finite input was accepted, out = %v", out) + } + }) + } + // The definition, written out: 1/(n h sqrt(2 pi)) * sum_k + // exp(-((x - sample[k])/h)^2 / 2) at x = 0.5 and x = 2.5, which are + // mirror images of each other in this sample and so agree exactly. + for _, x := range []float64{0.5, 2.5} { + want := 0.0 + for _, s := range []float64{0, 1, 2, 3} { + z := (x - s) / 0.5 + want += math.Exp(-0.5 * z * z) + } + want /= 4 * 0.5 * math.Sqrt(2*math.Pi) + out, err := KernelDensity(sample, 0.5, pinVector(t, x)) + if err != nil { + t.Fatalf("KernelDensity(%g): %v", x, err) + } + if got := out.FloatAt(0); got != want { + t.Fatalf("KernelDensity at %g = %.17g, the hand-written KDE is %.17g", x, got, want) + } + } +} + +// TestMultivariateNormalRejectsNonFiniteInput pins the second half of +// the finiteness gap: the covariance was scanned for non-finite +// entries, the mean and the point were not, so a NaN mean or point came +// back as a NaN log density and drew NaN vectors, all with a nil error. +func TestMultivariateNormalRejectsNonFiniteInput(t *testing.T) { + mean := pinVector(t, 0, 0) + cov, err := core.FromFloats([]float64{1, 0.5, 0.5, 1}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + nanCov, err := core.FromFloats([]float64{1, math.NaN(), math.NaN(), 1}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + x := pinVector(t, 1, 1) + nanMean := pinVector(t, math.NaN(), 0) + infMean := pinVector(t, math.Inf(1), 0) + nanPoint := pinVector(t, 1, math.NaN()) + g := core.NewGenerator(7) + for _, tc := range []struct { + name string + call func() error + }{ + {"density NaN mean", func() error { _, err := MultivariateNormalLogDensity(nanMean, cov, x); return err }}, + {"density +Inf mean", func() error { _, err := MultivariateNormalLogDensity(infMean, cov, x); return err }}, + {"density NaN covariance", func() error { _, err := MultivariateNormalLogDensity(mean, nanCov, x); return err }}, + {"density NaN point", func() error { _, err := MultivariateNormalLogDensity(mean, cov, nanPoint); return err }}, + {"draws NaN mean", func() error { _, err := MultivariateNormalDraws(g, 2, nanMean, cov); return err }}, + {"draws NaN covariance", func() error { _, err := MultivariateNormalDraws(g, 2, mean, nanCov); return err }}, + } { + t.Run(tc.name, func(t *testing.T) { + if err := tc.call(); err == nil { + t.Fatal("a non-finite input was accepted") + } + }) + } + // The clean path is pinned by hand: with mean 0, cov [[1,.5],[.5,1]] + // and x = (1,1), the log density is -ln(2 pi) - 0.5 ln(0.75) - 2/3. + got, err := MultivariateNormalLogDensity(mean, cov, x) + if err != nil { + t.Fatalf("MultivariateNormalLogDensity: %v", err) + } + want := -math.Log(2*math.Pi) - 0.5*math.Log(0.75) - 2.0/3.0 + if math.Abs(got-want) > 1e-12 { + t.Fatalf("log density = %.17g, hand value %.17g", got, want) + } + draws, err := MultivariateNormalDraws(g, 3, mean, cov) + if err != nil { + t.Fatalf("MultivariateNormalDraws: %v", err) + } + if draws.Len() != 6 { + t.Fatalf("draws length %d, want 6", draws.Len()) + } +} + +// TestMultivariateNormalRejectsAsymmetricCovariance pins the P2: the +// Cholesky factor read only the lower triangle, so an asymmetric matrix +// was silently symmetrised and the reported density described a +// different distribution from the one handed in. A mirror that differs +// beyond a relative tolerance is refused by name; a mirror that differs +// by rounding only is still accepted. +func TestMultivariateNormalRejectsAsymmetricCovariance(t *testing.T) { + mean := pinVector(t, 0, 0) + x := pinVector(t, 1, 1) + asym, err := core.FromFloats([]float64{1, 100, 0.5, 1}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + for _, tc := range []struct { + name string + call func() error + }{ + {"density", func() error { _, err := MultivariateNormalLogDensity(mean, asym, x); return err }}, + {"draws", func() error { _, err := MultivariateNormalDraws(core.NewGenerator(3), 2, mean, asym); return err }}, + } { + t.Run(tc.name, func(t *testing.T) { + if err := tc.call(); err == nil { + t.Fatal("the asymmetric covariance was accepted") + } + }) + } + // Rounding-level asymmetry is not asymmetry: 0.5 + 1e-15 against 0.5 + // is a mirror pair an A*A^T assembly can produce. + sym, err := core.FromFloats([]float64{1, 0.5 + 1e-15, 0.5, 1}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err := MultivariateNormalLogDensity(mean, sym, x) + if err != nil { + t.Fatalf("a rounding-level mirror difference was refused: %v", err) + } + want := -math.Log(2*math.Pi) - 0.5*math.Log(0.75) - 2.0/3.0 + if math.Abs(got-want) > 1e-12 { + t.Fatalf("log density = %.17g, the symmetric hand value is %.17g", got, want) + } +} + +// TestKolmogorovSeriesTruncationRule pins the inert defect: the early +// exit asked whether the current term was 1e-14 times the next, which is +// exp(2(2k+1)l^2) > 1 for every l > 0, so it could never fire and the +// comment described a termination check that was not there. The +// truncation is now the first term below the rounding level; the count +// must be exactly that term, and the summed series must agree with a +// directly summed one to the last bit. +func TestKolmogorovSeriesTruncationRule(t *testing.T) { + term := func(k int, lambda float64) float64 { + return 2 * math.Exp(-2*float64(k)*float64(k)*lambda*lambda) + } + for _, lambda := range []float64{0.2, 0.3, 0.5, 0.75, 1, 1.5, 2, 3, 5} { + n := kolmogorovTermCount(lambda) + if n < 1 || n > maxKolmogorovTerms { + t.Fatalf("lambda = %g: %d terms is outside [1, %d]", lambda, n, maxKolmogorovTerms) + } + if term(n, lambda) >= kolmogorovEps { + t.Fatalf("lambda = %g: term %d is %g, still above the %g rounding level: the series stops too early", + lambda, n, term(n, lambda), kolmogorovEps) + } + if n > 1 && term(n-1, lambda) < kolmogorovEps { + t.Fatalf("lambda = %g: term %d is already below the rounding level: the series adds terms that cannot move the sum", + lambda, n-1) + } + direct := 0.0 + for k := 1; k <= 400; k++ { + if k%2 == 0 { + direct -= term(k, lambda) + } else { + direct += term(k, lambda) + } + } + want := min(1, max(0, direct)) + if got := kolmogorovTail(lambda); got != want { + t.Fatalf("lambda = %g: tail = %.17g, a directly summed 400-term series gives %.17g", lambda, got, want) + } + } + // The documented short circuit below lambda = 0.2 is untouched. + if got := kolmogorovTail(0.1); got != 1 { + t.Fatalf("kolmogorovTail(0.1) = %g, want 1", got) + } +} + +// bootstrapMean is the bootstrap statistic the sweeps below resample with. +func bootstrapMean(a *core.Array) (float64, error) { + return core.Mean(a) +} + +// TestNaNParametersRejected pins the class sweep over the package's +// float parameters. Each guard used to be written out (`x < 0 || x > 1` +// and friends), which a NaN walks through because every comparison +// against a NaN is false: the parameter reached the algorithm and the +// call answered a value, a NaN, or nothing at all. The NaN-rejecting +// form `!(x >= 0 && x <= 1)` must turn every one of them into the +// refusal the caller deserves. The watchdog turns a regression to the +// unbounded discrete-quantile loop into a failure rather than a stuck +// test binary. +func TestNaNParametersRejected(t *testing.T) { + g := core.NewGenerator(1) + clean := pinVector(t, 1, 2, 3, 4, 5) + nanSample := pinVector(t, 1, 2, math.NaN(), 4, 5) + cases := []struct { + name string + call func() error + }{ + {"GammaLower(NaN shape)", func() error { _, err := GammaLower(math.NaN(), 1); return err }}, + {"GammaLower(NaN x)", func() error { _, err := GammaLower(1, math.NaN()); return err }}, + {"GammaUpper(NaN shape)", func() error { _, err := GammaUpper(math.NaN(), 1); return err }}, + {"GammaUpper(NaN x)", func() error { _, err := GammaUpper(1, math.NaN()); return err }}, + {"BetaIncomplete(NaN x)", func() error { _, err := BetaIncomplete(math.NaN(), 2, 2); return err }}, + {"BetaIncomplete(NaN a)", func() error { _, err := BetaIncomplete(0.5, math.NaN(), 2); return err }}, + {"BetaIncomplete(NaN b)", func() error { _, err := BetaIncomplete(0.5, 2, math.NaN()); return err }}, + {"ExponentialCDF(NaN x)", func() error { _, err := ExponentialCDF(math.NaN(), 2); return err }}, + {"ExponentialCDF(NaN rate)", func() error { _, err := ExponentialCDF(1, math.NaN()); return err }}, + {"GammaCDF(NaN x)", func() error { _, err := GammaCDF(math.NaN(), 2, 3); return err }}, + {"GammaCDF(NaN shape)", func() error { _, err := GammaCDF(1, math.NaN(), 3); return err }}, + {"GammaCDF(NaN rate)", func() error { _, err := GammaCDF(1, 2, math.NaN()); return err }}, + {"ChiSquareCDF(NaN x)", func() error { _, err := ChiSquareCDF(math.NaN(), 2); return err }}, + {"StudentTCDF(NaN t)", func() error { _, err := StudentTCDF(math.NaN(), 5); return err }}, + {"PoissonCDF(NaN lambda)", func() error { _, err := PoissonCDF(3, math.NaN()); return err }}, + {"BinomialCDF(NaN p)", func() error { _, err := BinomialCDF(3, 10, math.NaN()); return err }}, + {"ExponentialQuantile(NaN q)", func() error { _, err := ExponentialQuantile(math.NaN(), 2); return err }}, + {"ExponentialQuantile(NaN rate)", func() error { _, err := ExponentialQuantile(0.5, math.NaN()); return err }}, + {"GammaQuantile(NaN q)", func() error { _, err := GammaQuantile(math.NaN(), 2, 3); return err }}, + {"GammaQuantile(NaN shape)", func() error { _, err := GammaQuantile(0.5, math.NaN(), 3); return err }}, + {"GammaQuantile(NaN rate)", func() error { _, err := GammaQuantile(0.5, 2, math.NaN()); return err }}, + {"ChiSquareQuantile(NaN q)", func() error { _, err := ChiSquareQuantile(math.NaN(), 2); return err }}, + {"StudentTQuantile(NaN q)", func() error { _, err := StudentTQuantile(math.NaN(), 5); return err }}, + {"PoissonQuantile(NaN q)", func() error { _, err := PoissonQuantile(math.NaN(), 3); return err }}, + {"PoissonQuantile(NaN lambda)", func() error { _, err := PoissonQuantile(0.5, math.NaN()); return err }}, + {"BinomialQuantile(NaN q)", func() error { _, err := BinomialQuantile(math.NaN(), 0.5, 5); return err }}, + {"BinomialQuantile(NaN p)", func() error { _, err := BinomialQuantile(0.5, math.NaN(), 5); return err }}, + {"ExponentialDraws(NaN rate)", func() error { _, err := ExponentialDraws(g, 3, math.NaN()); return err }}, + {"GammaDraws(NaN shape)", func() error { _, err := GammaDraws(g, 3, math.NaN(), 2); return err }}, + {"GammaDraws(NaN rate)", func() error { _, err := GammaDraws(g, 3, 2, math.NaN()); return err }}, + {"PoissonDraws(NaN lambda)", func() error { _, err := PoissonDraws(g, 3, math.NaN()); return err }}, + {"BinomialDraws(NaN p)", func() error { _, err := BinomialDraws(g, 3, 4, math.NaN()); return err }}, + {"ChiSquareGoodnessOfFit(NaN observed)", func() error { + _, _, _, err := ChiSquareGoodnessOfFit(nanSample, clean) + return err + }}, + {"ChiSquareGoodnessOfFit(NaN expected)", func() error { + _, _, _, err := ChiSquareGoodnessOfFit(clean, nanSample) + return err + }}, + {"BootstrapCI(NaN level)", func() error { _, _, err := BootstrapCI(clean, bootstrapMean, math.NaN(), 20, 1); return err }}, + {"TrimmedMean(NaN fraction)", func() error { _, err := TrimmedMean(clean, math.NaN()); return err }}, + {"KernelDensity(NaN bandwidth)", func() error { _, err := KernelDensity(clean, math.NaN(), clean); return err }}, + {"WelchTTest(NaN sample)", func() error { _, _, _, err := WelchTTest(nanSample, clean); return err }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + var err error + pinWatchdog(t, 5*time.Second, tc.name, func() { err = tc.call() }) + if err == nil { + t.Fatalf("%s answered with no error for a NaN parameter", tc.name) + } + }) + } + // No guard overreaches: the legal domain still answers, and the two + // values below are hand-checkable. The incomplete beta is the + // continued fraction's own rounding level, a few ulps at 0.5. + if p, err := BetaIncomplete(0.5, 2, 2); err != nil || math.Abs(p-0.5) > 1e-15 { + t.Fatalf("BetaIncomplete(0.5, 2, 2) = %v, %v, want 0.5 by symmetry", p, err) + } + if p, err := GammaLower(1, 1); err != nil || math.Abs(p-(1-math.Exp(-1))) > 1e-15 { + t.Fatalf("GammaLower(1, 1) = %v, %v, want 1 - e^-1", p, err) + } + if q, err := ExponentialQuantile(0.5, 3); err != nil || math.Abs(q-math.Ln2/3) > 1e-15 { + t.Fatalf("ExponentialQuantile(0.5, 3) = %v, %v, want ln 2 / 3", q, err) + } + if p, err := PoissonQuantile(0.5, 3); err != nil || p != 3 { + t.Fatalf("PoissonQuantile(0.5, 3) = %v, %v, want 3", p, err) + } +} + +// TestComplexInputsRejected pins the dtype sweep: every entry point in +// the package refuses a complex array by name, and the ones exercised +// here (the three MVN inputs, BootstrapCI's data, KernelDensity's +// sample and points) walked the float accessor into a nil payload and +// panicked instead. +func TestComplexInputsRejected(t *testing.T) { + cx2, err := core.FromComplexes([]complex128{1 + 1i, 2 - 1i}, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + cx5, err := core.FromComplexes([]complex128{1 + 1i, 2, 3, 4, 5}, 5) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + cxMat, err := core.FromComplexes([]complex128{1, 0, 0, 1}, 2, 2) + if err != nil { + t.Fatalf("FromComplexes: %v", err) + } + mean := pinVector(t, 0, 0) + x := pinVector(t, 1, 1) + cov, err := core.FromFloats([]float64{1, 0.5, 0.5, 1}, 2, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + g := core.NewGenerator(3) + for _, tc := range []struct { + name string + call func() error + }{ + {"MVN density complex mean", func() error { _, err := MultivariateNormalLogDensity(cx2, cov, x); return err }}, + {"MVN density complex covariance", func() error { _, err := MultivariateNormalLogDensity(mean, cxMat, x); return err }}, + {"MVN density complex point", func() error { _, err := MultivariateNormalLogDensity(mean, cov, cx2); return err }}, + {"MVN draws complex mean", func() error { _, err := MultivariateNormalDraws(g, 2, cx2, cov); return err }}, + {"MVN draws complex covariance", func() error { _, err := MultivariateNormalDraws(g, 2, mean, cxMat); return err }}, + {"BootstrapCI complex data", func() error { _, _, err := BootstrapCI(cx5, bootstrapMean, 0.95, 20, 1); return err }}, + {"KernelDensity complex sample", func() error { _, err := KernelDensity(cx5, 0.5, cx5); return err }}, + {"KernelDensity complex points", func() error { _, err := KernelDensity(pinVector(t, 1, 2, 3, 4, 5), 0.5, cx5); return err }}, + {"Histogram2D complex sample", func() error { _, _, _, err := Histogram2D(cx5, cx5, 2, 2); return err }}, + } { + t.Run(tc.name, func(t *testing.T) { + if err := tc.call(); err == nil { + t.Fatal("a complex input was accepted") + } + }) + } + // The real path is untouched: a constant sample has a degenerate + // bootstrap interval at its own value. + lo, hi, err := BootstrapCI(pinVector(t, 5, 5, 5, 5), bootstrapMean, 0.95, 50, 1) + if err != nil { + t.Fatalf("BootstrapCI: %v", err) + } + if lo != 5 || hi != 5 { + t.Fatalf("the constant sample gives [%g, %g], want [5, 5]", lo, hi) + } +} + +// TestLinearRegressionConstantResponseStatistics pins the NaN guard a +// degenerate response needs: with no variation in y the fit reproduces +// it exactly, so the explained and residual sums are both zero and the +// F statistic is 0/0. A NaN there would spread through every consumer +// of the result, where the finite zero is the value the tail reads as +// p = 1: there is no evidence of a model. +func TestLinearRegressionConstantResponseStatistics(t *testing.T) { + design, err := core.FromFloats([]float64{1, 0, 1, 1, 1, 2, 1, 3}, 4, 2) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + y := pinVector(t, 5, 5, 5, 5) + for _, tc := range []struct { + name string + fit func() (*LinearRegressionResult, error) + }{ + {"ordinary", func() (*LinearRegressionResult, error) { return LinearRegression(design, y) }}, + {"weighted", func() (*LinearRegressionResult, error) { + return WeightedLinearRegression(design, y, pinVector(t, 1, 1, 1, 2)) + }}, + } { + t.Run(tc.name, func(t *testing.T) { + res, err := tc.fit() + if err != nil { + t.Fatalf("the constant-response fit: %v", err) + } + if math.IsNaN(res.FStatistic) || res.FStatistic != 0 { + t.Fatalf("FStatistic = %v, want the finite 0", res.FStatistic) + } + if res.FPValue != 1 { + t.Fatalf("FPValue = %v, want 1", res.FPValue) + } + if res.DModel != 1 || res.DResidual != 2 { + t.Fatalf("degrees of freedom (%d, %d), want (1, 2)", res.DModel, res.DResidual) + } + }) + } +} + +// bigPhi returns Phi(x) in 256-bit arithmetic through the Taylor +// series of the error function at x/sqrt(2). Nothing here calls the +// library, so it is an independent reference for the quantile tests. +func bigPhi(x *big.Float) *big.Float { + const prec = 256 + one := new(big.Float).SetPrec(prec).SetInt64(1) + y := new(big.Float).SetPrec(prec).Quo(new(big.Float).SetPrec(prec).Set(x), bigSqrt2()) + neg := y.Sign() < 0 + ay := new(big.Float).SetPrec(prec).Abs(y) + // erf(ay) = (2/sqrt(pi)) * sum_{n>=0} (-1)^n ay^(2n+1) / (n! (2n+1)). + term := new(big.Float).SetPrec(prec).Set(ay) // ay^(2n+1)/n! at n = 0 + sum := new(big.Float).SetPrec(prec).Set(ay) + ay2 := new(big.Float).SetPrec(prec).Mul(ay, ay) + for n := 1; ; n++ { + term.Mul(term, ay2) + term.Quo(term, new(big.Float).SetPrec(prec).SetInt64(int64(n))) + term.Neg(term) + contrib := new(big.Float).SetPrec(prec).Quo(term, + new(big.Float).SetPrec(prec).SetInt64(int64(2*n+1))) + sum.Add(sum, contrib) + if contrib.Sign() == 0 || contrib.MantExp(nil) < -400 { + break + } + } + // 2/sqrt(pi). + twoOverSqrtPi := new(big.Float).SetPrec(prec).Quo( + new(big.Float).SetPrec(prec).SetInt64(2), + new(big.Float).SetPrec(prec).Sqrt(new(big.Float).SetPrec(prec).SetFloat64(math.Pi))) + erf := new(big.Float).SetPrec(prec).Mul(sum, twoOverSqrtPi) + if neg { + erf.Neg(erf) + } + // Phi(x) = (1 + erf(x/sqrt(2)))/2. + return new(big.Float).SetPrec(prec).Quo( + new(big.Float).SetPrec(prec).Add(one, erf), + new(big.Float).SetPrec(prec).SetInt64(2)) +} + +// bigNormalQuantile inverts bigPhi by bisection on [-12, 12], the +// independent reference the NormalQuantile test compares against. +func bigNormalQuantile(q float64) float64 { + const prec = 256 + target := new(big.Float).SetPrec(prec).SetFloat64(q) + lo := new(big.Float).SetPrec(prec).SetFloat64(-12) + hi := new(big.Float).SetPrec(prec).SetFloat64(12) + mid := new(big.Float).SetPrec(prec) + for range 220 { + mid.Add(lo, hi) + mid.Quo(mid, new(big.Float).SetPrec(prec).SetInt64(2)) + if bigPhi(mid).Cmp(target) < 0 { + lo.Set(mid) + } else { + hi.Set(mid) + } + } + mid.Add(lo, hi) + mid.Quo(mid, new(big.Float).SetPrec(prec).SetInt64(2)) + out, _ := mid.Float64() + return out +} diff --git a/stats/poissonregression_test.go b/stats/poissonregression_test.go new file mode 100644 index 0000000..8add67c --- /dev/null +++ b/stats/poissonregression_test.go @@ -0,0 +1,293 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// TestPoissonRegressionRecoversCoefficients fits a generated count +// response whose truth is known: with 1500 samples the Newton fit must +// land within a few standard errors of the generating coefficients, +// the Wald statistics must flag the slope, and the fitted means must +// increase in the direction of the true slope. +func TestPoissonRegressionRecoversCoefficients(t *testing.T) { + g := core.NewGenerator(7) + const n = 1500 + design := core.New(core.Float, n, 2) + y := core.New(core.Float, n) + for i := range n { + xv := -2 + 4*g.Unit() + design.RawFloats()[i*2] = 1 + design.RawFloats()[i*2+1] = xv + counts, err := PoissonDraws(g, 1, math.Exp(0.3+0.9*xv)) + if err != nil { + t.Fatalf("PoissonDraws: %v", err) + } + y.RawFloats()[i] = counts.FloatAt(0) + } + res, err := PoissonRegression(design, y) + if err != nil { + t.Fatalf("PoissonRegression: %v", err) + } + if !res.Converged { + t.Fatal("the fit reported no convergence") + } + if math.Abs(res.Coefficients[0]-0.3) > 4*res.StandardErrors[0] { + t.Fatalf("intercept = %.4g (%.4g SE), outside four SEs of 0.3", + res.Coefficients[0], res.StandardErrors[0]) + } + if math.Abs(res.Coefficients[1]-0.9) > 4*res.StandardErrors[1] { + t.Fatalf("slope = %.4g (%.4g SE), outside four SEs of 0.9", + res.Coefficients[1], res.StandardErrors[1]) + } + if res.ZStatistics[1] <= 3 { + t.Fatalf("slope z = %.4g, want a clearly non-zero effect", res.ZStatistics[1]) + } + if !(res.Fitted[n-1] > res.Fitted[0]) { + t.Fatalf("fitted means not increasing: %g then %g", res.Fitted[0], res.Fitted[n-1]) + } +} + +// TestPoissonRegressionInterceptOnlyClosedForm pins the one case with +// a closed-form maximum likelihood estimate: with an intercept-only +// design the estimate is log of the mean count and its standard error +// is 1/sqrt(n·mean). +func TestPoissonRegressionInterceptOnlyClosedForm(t *testing.T) { + const n = 40 + design := core.New(core.Float, n, 1) + y := core.New(core.Float, n) + total := 0.0 + for i := range n { + design.RawFloats()[i] = 1 + y.RawFloats()[i] = float64(i % 7) + total += float64(i % 7) + } + meanY := total / n + res, err := PoissonRegression(design, y) + if err != nil { + t.Fatalf("PoissonRegression: %v", err) + } + if math.Abs(res.Coefficients[0]-math.Log(meanY)) > 1e-9 { + t.Fatalf("intercept = %.15g, want log(%.15g) = %.15g", + res.Coefficients[0], meanY, math.Log(meanY)) + } + wantSE := 1 / math.Sqrt(n*meanY) + if math.Abs(res.StandardErrors[0]-wantSE) > 1e-9*wantSE { + t.Fatalf("SE = %.15g, want %.15g", res.StandardErrors[0], wantSE) + } + for i := range n { + if math.Abs(res.Fitted[i]-meanY) > 1e-9 { + t.Fatalf("fitted[%d] = %.15g, want %.15g", i, res.Fitted[i], meanY) + } + } + // The score must vanish at the optimum. + score := 0.0 + for i := range n { + score += res.Fitted[i] - y.FloatAt(i) + } + if math.Abs(score) > 1e-6 { + t.Fatalf("score at the optimum = %.3g, want 0", score) + } +} + +// TestPoissonRegressionStandardErrorsClosedForm checks the Wald +// inference against a closed-form 2x2 inverse of the Fisher +// information built from the fitted means: no solve, only the +// adjugate formula. +func TestPoissonRegressionStandardErrorsClosedForm(t *testing.T) { + // Near-deterministic counts: large means make the rounding to + // integers a relative perturbation under 1e-3, so the fit must + // recover the generating coefficients tightly. + const n = 25 + truth := [2]float64{0.2, 0.5} + design := core.New(core.Float, n, 2) + y := core.New(core.Float, n) + for i := range n { + xv := float64(i) + design.RawFloats()[i*2] = 1 + design.RawFloats()[i*2+1] = xv + y.RawFloats()[i] = math.Round(math.Exp(truth[0] + truth[1]*xv)) + } + res, err := PoissonRegression(design, y) + if err != nil { + t.Fatalf("PoissonRegression: %v", err) + } + for j := range 2 { + if math.Abs(res.Coefficients[j]-truth[j]) > 4e-3 { + t.Fatalf("coefficient %d = %.6g, want %.6g within rounding noise", + j, res.Coefficients[j], truth[j]) + } + } + // Fisher information from the fitted means, inverted by the + // adjugate. + a, b, c := 0.0, 0.0, 0.0 + for i := range n { + x1 := design.FloatAt(i * 2) + x2 := design.FloatAt(i*2 + 1) + mu := res.Fitted[i] + a += x1 * x1 * mu + b += x1 * x2 * mu + c += x2 * x2 * mu + } + det := a*c - b*b + wantSE0 := math.Sqrt(c / det) + wantSE1 := math.Sqrt(a / det) + if math.Abs(res.StandardErrors[0]-wantSE0) > 1e-9*wantSE0 { + t.Fatalf("SE0 = %.15g, want %.15g", res.StandardErrors[0], wantSE0) + } + if math.Abs(res.StandardErrors[1]-wantSE1) > 1e-9*wantSE1 { + t.Fatalf("SE1 = %.15g, want %.15g", res.StandardErrors[1], wantSE1) + } + z := res.ZStatistics[0] + // The p-value reference is computed independently of the tail + // formula the fit uses: Simpson quadrature of the Gaussian density, + // every summand positive, sharing nothing with NormalCDF or + // math.Erfc. The intercept's z lands far past the z ≈ 8.3 where the + // algebraic form 2·(1−Φ(z)) cancels to an exact zero, so this pin + // fails against the cancelled form, which answers 0 here. + wantP := gaussTailReference(math.Abs(z)) + if math.Abs(res.PValues[0]-wantP) > 1e-6*wantP { + t.Fatalf("p-value = %.15g, want %.15g", res.PValues[0], wantP) + } + // The slope's z is in the hundreds: its tail underflows the float64 + // range entirely, and an exact zero is the correctly rounded answer + // both the fit and the reference produce. + if res.PValues[1] != 0 { + t.Fatalf("p-value of the slope = %.15g, want the underflowed 0", res.PValues[1]) + } +} + +// TestPoissonRegressionRefusals checks the response and design +// validation and the honest failure of a singular Fisher information +// from a duplicated column. +func TestPoissonRegressionRefusals(t *testing.T) { + y4 := mustFromFloats(t, []float64{1, 0, 1, 0}, 4) + if _, err := PoissonRegression(core.New(core.Float, 4), y4); err == nil { + t.Fatal("a rank 1 design was accepted") + } + design := core.New(core.Float, 4, 2) + for i := range 4 { + design.RawFloats()[i*2] = 1 + design.RawFloats()[i*2+1] = float64(i) + } + rank2 := core.New(core.Float, 2, 2) + if _, err := PoissonRegression(design, rank2); err == nil { + t.Fatal("a rank 2 response was accepted") + } + if _, err := PoissonRegression(design, mustFromFloats(t, []float64{1, 0, 1}, 3)); err == nil { + t.Fatal("a length mismatch was accepted") + } + if _, err := PoissonRegression(design, mustFromFloats(t, []float64{1, 0, 1, -2}, 4)); err == nil { + t.Fatal("a negative count was accepted") + } + if _, err := PoissonRegression(design, mustFromFloats(t, []float64{1, 0, 1, 0.5}, 4)); err == nil { + t.Fatal("a fractional count was accepted") + } + if _, err := PoissonRegression(design, mustFromFloats(t, []float64{1, 0, math.NaN(), 0}, 4)); err == nil { + t.Fatal("a NaN count was accepted") + } + if _, err := PoissonRegression(design, mustFromFloats(t, []float64{1, 0, math.Inf(1), 0}, 4)); err == nil { + t.Fatal("an infinite count was accepted") + } + bad := mustFromFloats(t, []float64{1, math.NaN(), 0, 1, 2, 0, 1, 2}, 4, 2) + if _, err := PoissonRegression(bad, y4); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("PoissonRegression with a NaN design: %v", err) + } + // n <= p has no unique fit. + square := core.New(core.Float, 2, 2) + square.RawFloats()[0] = 1 + square.RawFloats()[3] = 1 + if _, err := PoissonRegression(square, mustFromFloats(t, []float64{1, 2}, 2)); err == nil { + t.Fatal("n <= p was accepted") + } + // A duplicated column makes the Fisher information singular at the + // start. + dup := core.New(core.Float, 4, 2) + for i := range 4 { + dup.RawFloats()[i*2] = float64(i) + dup.RawFloats()[i*2+1] = float64(i) + } + if _, err := PoissonRegression(dup, y4); err == nil || !strings.Contains(err.Error(), "singular") { + t.Fatalf("a duplicated column: %v", err) + } +} + +// TestPoissonRegressionIsDeterministic runs the same fit twice and +// requires identical coefficients bit for bit, as every Tensor entry +// point promises. +func TestPoissonRegressionIsDeterministic(t *testing.T) { + const n = 300 + build := func(seed int64) (*core.Array, *core.Array) { + g := core.NewGenerator(seed) + design := core.New(core.Float, n, 2) + y := core.New(core.Float, n) + for i := range n { + xv := -1 + 2*g.Unit() + design.RawFloats()[i*2] = 1 + design.RawFloats()[i*2+1] = xv + counts, err := PoissonDraws(g, 1, math.Exp(0.4+0.6*xv)) + if err != nil { + t.Fatalf("PoissonDraws: %v", err) + } + y.RawFloats()[i] = counts.FloatAt(0) + } + return design, y + } + d1, y1 := build(11) + d2, y2 := build(11) + r1, err := PoissonRegression(d1, y1) + if err != nil { + t.Fatalf("first fit: %v", err) + } + r2, err := PoissonRegression(d2, y2) + if err != nil { + t.Fatalf("second fit: %v", err) + } + for j := range r1.Coefficients { + if r1.Coefficients[j] != r2.Coefficients[j] { + t.Fatalf("coefficient %d differs: %.17g vs %.17g", + j, r1.Coefficients[j], r2.Coefficients[j]) + } + } + if r1.LogLikelihood != r2.LogLikelihood { + t.Fatalf("log likelihood differs: %.17g vs %.17g", + r1.LogLikelihood, r2.LogLikelihood) + } +} + +// TestPoissonRegressionFloat32Design runs the fit over a float32 +// design, the widened path the raw-slice fast path must agree with. +func TestPoissonRegressionFloat32Design(t *testing.T) { + const n = 120 + f64 := core.New(core.Float, n, 2) + f32 := core.New(core.Float32, n, 2) + y := core.New(core.Float, n) + for i := range n { + xv := -2 + 4*float64(i)/float64(n-1) + f64.RawFloats()[i*2] = 1 + f64.RawFloats()[i*2+1] = xv + f32.SetFloatAt(i*2, 1) + f32.SetFloatAt(i*2+1, xv) + y.RawFloats()[i] = math.Round(math.Exp(0.2 + 0.3*xv)) + } + ref, err := PoissonRegression(f64, y) + if err != nil { + t.Fatalf("float64 fit: %v", err) + } + got, err := PoissonRegression(f32, y) + if err != nil { + t.Fatalf("float32 fit: %v", err) + } + for j := range ref.Coefficients { + if math.Abs(got.Coefficients[j]-ref.Coefficients[j]) > 1e-6 { + t.Fatalf("coefficient %d: float32 %.12g vs float64 %.12g", + j, got.Coefficients[j], ref.Coefficients[j]) + } + } +} diff --git a/stats/pool_reuse_pin_test.go b/stats/pool_reuse_pin_test.go new file mode 100644 index 0000000..c7767f7 --- /dev/null +++ b/stats/pool_reuse_pin_test.go @@ -0,0 +1,56 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// The pooled pair-slope buffer is sliced to the pair count before the +// median reads it: a reused buffer carries a longer tail from an earlier +// fit, and the median must never see it. The sequence below makes the +// long fit prime the pool and the short fits reuse it; a garbage +// collection between the calls would drop the entry and weaken the pin, +// so the fits run back to back. + +func TestTheilSenPooledReuseKeepsExactSlope(t *testing.T) { + // Both fixtures sit exactly on y = 3x - 2, so every pair slope is + // exactly 3 and the answers are exact whatever the pair order. + makeLine := func(n int) (*core.Array, *core.Array) { + xs := make([]float64, n) + ys := make([]float64, n) + for i := range n { + xs[i] = float64(i) + 0.5 + ys[i] = 3*xs[i] - 2 + } + x, err := core.FromFloats(xs, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + y, err := core.FromFloats(ys, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return x, y + } + bigX, bigY := makeLine(300) + smallX, smallY := makeLine(40) + if _, _, err := TheilSenRegression(bigX, bigY); err != nil { + t.Fatalf("long fit: %v", err) + } + for round := range 3 { + intercept, slope, err := TheilSenRegression(smallX, smallY) + if err != nil { + t.Fatalf("short fit %d: %v", round, err) + } + if slope != 3 { + t.Fatalf("short fit %d: slope = %v, want exactly 3", round, slope) + } + if intercept != -2 { + t.Fatalf("short fit %d: intercept = %v, want exactly -2", round, intercept) + } + } +} diff --git a/stats/quantile_newton_test.go b/stats/quantile_newton_test.go new file mode 100644 index 0000000..ba2317f --- /dev/null +++ b/stats/quantile_newton_test.go @@ -0,0 +1,243 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Precision pins for the bracketed Newton quantile walk: on a grid of +// extreme and central probabilities the Newton answer must invert the +// float64 CDF at least as well as the bisection fallback it replaced, +// distribution by distribution, against either the independent +// 256-bit inverse of Phi or a 200-bit refinement of the CDF's own +// crossing. The comparison is made at the resolution the format +// actually delivers: a float64 CDF is a staircase whose steps are an +// ulp of the probability wide, so answers on the step the crossing +// sits on are equally exact by construction. + +package stats + +import ( + "math" + "math/big" + "testing" +) + +// newtonBigBetaLogNormaliser returns ln B(a, b) for the beta density. +func newtonBigBetaLogNormaliser(a, b float64) float64 { + lb, _ := math.Lgamma(a + b) + la, _ := math.Lgamma(a) + lb2, _ := math.Lgamma(b) + return lb - la - lb2 +} + +// newtonQuantileBigReference refines the crossing of the float64 CDF +// through q by bisection carried out in 200-bit arithmetic, three +// hundred rounds: the bracket collapses onto the exact transition +// point between the last sample below q and the first above it, the +// limit both float64 walks approximate. The returned float64 is that +// crossing rounded, the correctly-rounded inverse of the float64 CDF. +func newtonQuantileBigReference(t *testing.T, name string, lo, hi, q float64, + cdf func(float64) (float64, error)) float64 { + t.Helper() + const prec = 200 + half := new(big.Float).SetPrec(prec).SetFloat64(0.5) + l := new(big.Float).SetPrec(prec).SetFloat64(lo) + h := new(big.Float).SetPrec(prec).SetFloat64(hi) + m := new(big.Float).SetPrec(prec) + for range 300 { + m.Add(l, h) + m.Mul(m, half) + mf, _ := m.Float64() + f, err := cdf(mf) + if err != nil { + t.Fatalf("%s: reference CDF at %g: %v", name, mf, err) + } + if f < q { + l.Set(m) + } else { + h.Set(m) + } + } + m.Add(l, h) + m.Mul(m, half) + out, _ := m.Float64() + return out +} + +// newtonStepWidth is the x-width of one probability step of the +// float64 CDF at the crossing: an ulp of the inverted probability +// over the density. An answer within a step and a half of the exact +// crossing sits on the crossing's own step of the staircase and +// inverts the float64 CDF as exactly as the format allows. +func newtonStepWidth(p, density float64) float64 { + if !(density > 0) { + return math.Inf(1) + } + return 1.5 * newtonUlp(p) / density +} + +// newtonUlp is one step of the float64 probability grid at p. +func newtonUlp(p float64) float64 { + return math.Nextafter(p, math.Inf(1)) - p +} + +// TestQuantileNewtonMatchesBisectionPrecision holds the bracketed +// Newton walk against the bisection fallback: both invert the same +// float64 CDF, the reference is the exact +// crossing refined at 200 bits (for the normal law the independent +// 256-bit inverse of Phi, the pins_test reference), and the acceptance +// bar is distribution-wise: beyond the format's own step width the +// Newton walk's worst distance must be the bisection's or better, and +// its worst CDF residual within one step of the bisection's. +func TestQuantileNewtonMatchesBisectionPrecision(t *testing.T) { + qs := []float64{1e-12, 0.001, 0.25, 0.5, 0.75, 0.999, 1 - 1e-12} + // The mirrored laws ride the same reflection the public quantiles + // use, since the bracket cannot start below Phi(0) = 0.5; the gamma + // and beta brackets walk down to their support's floor instead, and + // the beta is clamped to that support, which the internal bracket + // needs. The beta rides the same machinery through BetaIncomplete. + laws := []struct { + name string + reflect bool + cdf func(float64) (float64, error) + pdf func(float64) float64 + seed float64 + }{ + {"Normal", true, func(x float64) (float64, error) { return NormalCDF(x), nil }, normalPdf, 1}, + {"Gamma(2.5, 1.2)", false, func(x float64) (float64, error) { return GammaCDF(x, 2.5, 1.2) }, + func(x float64) float64 { return gammaShapeRatePdf(x, 2.5, 1.2) }, 2.5 / 1.2}, + {"Beta(2, 3)", false, func(x float64) (float64, error) { + if x >= 1 { + return 1, nil + } + if x <= 0 { + return 0, nil + } + return BetaIncomplete(x, 2, 3) + }, func(x float64) float64 { + if x <= 0 || x >= 1 { + return 0 + } + return math.Exp(-newtonBigBetaLogNormaliser(2, 3) + math.Log(x) + 2*math.Log1p(-x)) + }, 0.4}, + {"StudentT(5)", true, func(x float64) (float64, error) { return StudentTCDF(x, 5) }, + func(x float64) float64 { return studentTPdf(x, 5) }, 1}, + } + for _, law := range laws { + t.Run(law.name, func(t *testing.T) { + worstDistBisect, worstResBisect := 0.0, 0.0 + worstDistNewton, worstResNewton := 0.0, 0.0 + for _, q := range qs { + // The mirrored laws bracket on the reflected probability + // below the seed's reach, exactly as the public quantile + // does, and both answers come back negated. + qEff, sign := q, 1.0 + if law.reflect && q < 0.5 { + qEff, sign = 1-q, -1 + } + bisect, err := continuousQuantile("bisect", qEff, law.seed, law.cdf, nil) + if err != nil { + t.Fatalf("bisection q = %g: %v", q, err) + } + newton, err := continuousQuantile("newton", qEff, law.seed, law.cdf, law.pdf) + if err != nil { + t.Fatalf("newton q = %g: %v", q, err) + } + bisect, newton = sign*bisect, sign*newton + // The bracket for the reference: the two answers plus a + // float on each side straddle the crossing. + lo := math.Nextafter(min(bisect, newton), math.Inf(-1)) + hi := math.Nextafter(max(bisect, newton), math.Inf(1)) + var ref float64 + if law.name == "Normal" { + ref = bigNormalQuantile(q) + } else { + ref = newtonQuantileBigReference(t, law.name, lo, hi, q, law.cdf) + } + // The bar lives in probability space, where inversion + // quality lives: the float64 CDF is a staircase whose + // steps and evaluation noise are ulps of the inverted + // probability, so the Newton answer passes when its CDF + // residual is the bisection's or better, one staircase's + // worth of quantisation aside. A flat stretch of the CDF + // (the t law rounds to a constant on a whole plateau + // around its median) answers both walks the same value + // and passes by construction. + resBisect, err := law.cdf(bisect) + if err != nil { + t.Fatalf("bisection residual q = %g: %v", q, err) + } + resNewton, err := law.cdf(newton) + if err != nil { + t.Fatalf("newton residual q = %g: %v", q, err) + } + noise := 16 * newtonUlp(qEff) + if math.Abs(resNewton-q) > math.Abs(resBisect-q)+noise { + t.Fatalf("%s q = %g: newton's CDF residual %.3g is past the bisection's %.3g", + law.name, q, resNewton-q, resBisect-q) + } + if r := math.Abs(resBisect-q) / newtonUlp(qEff); r > worstResBisect { + worstResBisect = r + } + if r := math.Abs(resNewton-q) / newtonUlp(qEff); r > worstResNewton { + worstResNewton = r + } + if d := math.Abs(bisect - ref); d > worstDistBisect { + worstDistBisect = d + } + if d := math.Abs(newton - ref); d > worstDistNewton { + worstDistNewton = d + } + } + t.Logf("bisection: worst distance %.3g, worst CDF residual %.3g steps", + worstDistBisect, worstResBisect) + t.Logf("newton: worst distance %.3g, worst CDF residual %.3g steps", + worstDistNewton, worstResNewton) + }) + } +} + +// TestNormalQuantileNewtonIndependentReference runs the normal law's +// grid against the 256-bit Taylor-series inverse of Phi directly: the +// Newton answer must stay within the pins_test bar at every +// grid point, widened only where the float64 CDF's ulp at the +// inverted probability sets a coarser limit than any inverse of it +// can beat, and the mirrored lower tail must stay exactly symmetric. +func TestNormalQuantileNewtonIndependentReference(t *testing.T) { + for _, q := range []float64{1e-12, 0.001, 0.25, 0.5, 0.75, 0.999, 1 - 1e-12} { + z, err := NormalQuantile(q) + if err != nil { + t.Fatalf("NormalQuantile(%g): %v", q, err) + } + ref := bigNormalQuantile(q) + allowed := 1e-9 + if d := normalPdf(z); d > 0 { + // The lower half inverts the reflected probability, so the + // format's resolution there is the ulp of 1 - q. + allowed = max(allowed, 2*newtonStepWidth(max(q, 1-q), d)) + } + if dev := math.Abs(z - ref); dev > allowed { + t.Fatalf("NormalQuantile(%g) = %.17g, the 256-bit reference is %.17g (off by %.3g, allowed %.3g)", + q, z, ref, dev, allowed) + } + // The symmetry stays exact on the mirrored route, which the + // public quantile takes strictly below one half. + if q < 0.5 { + mirror, err := NormalQuantile(1 - q) + if err != nil { + t.Fatalf("NormalQuantile(%g): %v", 1-q, err) + } + if z+mirror != 0 { + t.Fatalf("NormalQuantile(%g) + NormalQuantile(%g) = %.17g, want exactly 0", + q, 1-q, z+mirror) + } + } + } + // The extreme tail keeps its dedicated bisection route, untouched + // by the Newton walk, and answers a probability the reflection + // cannot represent, at the accuracy that route has always given. + z, err := NormalQuantile(1e-15) + if err != nil { + t.Fatalf("NormalQuantile(1e-15): %v", err) + } + if back := NormalCDF(z); math.Abs(back-1e-15) > 5e-17 { + t.Fatalf("NormalCDF(NormalQuantile(1e-15)) = %g, want 1e-15", back) + } +} diff --git a/stats/quantile_test.go b/stats/quantile_test.go new file mode 100644 index 0000000..2370ac6 --- /dev/null +++ b/stats/quantile_test.go @@ -0,0 +1,67 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" +) + +// TestQuantile pins linear interpolation between order statistics. +func TestQuantile(t *testing.T) { + a := mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10}, 10) + q, err := Quantile(a, []float64{0.25}) + if err != nil { + t.Fatal(err) + } + // Linear interpolation: position 0.25*(9)=2.25 lands between 3 and 4. + if v := q.FloatAt(0); math.Abs(v-3.25) > 1e-9 { + t.Errorf("q25: %v, want 3.25", v) + } + med, _ := Quantile(a, []float64{0.5}) + if v := med.FloatAt(0); math.Abs(v-5.5) > 1e-9 { + t.Errorf("median quantile: %v", v) + } +} + +// TestQuantileDegenerateInputs pins the guard that used to panic: +// Quantile on an empty sample. +func TestQuantileDegenerateInputs(t *testing.T) { + empty := mustFloats(t, []float64{}, 0) + if _, err := Quantile(empty, []float64{0.5}); err == nil { + t.Error("Quantile empty: expected error") + } +} + +// TestBinCounts counts values in equally wide bins. +func TestBinCounts(t *testing.T) { + vals := mustFloats(t, []float64{-3, 5, 12, 25}, 4) + counts, err := BinCounts(vals, 2) + if err != nil { + t.Fatal(err) + } + if counts.Len() != 2 { + t.Fatalf("counts len: %d", counts.Len()) + } + if c1 := counts.FloatAt(0); c1 < 1 { + t.Errorf("first bin empty: %v", counts.RawFloats()) + } +} + +// TestComplexStatsErrors pins the refusals for complex input: no +// ordering, so no median, standard deviation or histogram. +func TestComplexStatsErrors(t *testing.T) { + c := mustComplexes(t, []complex128{complex(1, 1)}, 1) + + if _, err := Median(c); err == nil || !strings.Contains(err.Error(), "no median") { + t.Fatalf("Median complex: %v", err) + } + if _, err := Std(c); err == nil || !strings.Contains(err.Error(), "no float standard deviation") { + t.Fatalf("Std complex: %v", err) + } + if _, _, err := Histogram(c, 2); err == nil || !strings.Contains(err.Error(), "no histogram") { + t.Fatalf("Histogram complex: %v", err) + } +} diff --git a/stats/quantilereg.go b/stats/quantilereg.go new file mode 100644 index 0000000..30a1935 --- /dev/null +++ b/stats/quantilereg.go @@ -0,0 +1,458 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Quantile regression: the linear fit that asks not for the mean of +// the response but for its tau-th conditional quantile, by minimising +// the check loss +// +// ρ_τ(r) = r·τ when r ≥ 0, r·(τ − 1) when r < 0, +// +// the asymmetric absolute loss that rewards a fitted line for putting +// the right fraction of the data beneath it. The implementation is the +// Frisch-Newton interior-point method on the dual, the form +// Portnoy and Koenker put the method in: the check loss problem is +// the linear program +// +// min Σᵢ τ·uᵢ + (1−τ)·vᵢ subject to u − v = y − Xβ, u, v ≥ 0, +// +// and its dual asks for +// +// max wᵀy subject to Xᵀw = (1−τ)·Xᵀ1, w ∈ [0, 1]ⁿ. +// +// At the optimum an observation with a positive residual carries w = 1, +// one with a negative residual w = 0, and the observations the fit +// reproduces exactly carry w strictly inside the box. A logarithmic +// barrier is laid on the box, every iteration solves the barrier's +// Newton system exactly in the p×p form XᵀD⁻¹X through the shared LU +// solve, and the barrier parameter falls geometrically. The start is +// exactly feasible, w = 1−τ for every observation, and the primal fit +// is read back off the interior set at every iteration, so the run +// records a monotone descent of the true check loss. + +// The documented schedule of the interior-point loop: the barrier +// parameter starts at the mean absolute residual of the ordinary +// least squares start, shrinks by this factor every iteration, and the +// run is settled when no coefficient of the dual moves by more than +// the tolerance, or when the parameter has fallen sixteen orders of +// magnitude, whichever comes first. +const ( + quantileMaxIterations = 100 + quantileMuShrink = 0.25 + quantileMuFloorRatio = 1e-16 + quantileStepFraction = 0.995 + quantileTolerance = 1e-14 +) + +// QuantileRegressionResult carries a quantile regression fit. +type QuantileRegressionResult struct { + // Coefficients are the quantile estimates β̂, one per design + // column, in the design's own order. The intercept, supplied by + // the caller as a constant column, is estimated like any other + // coefficient: the check loss pulls it to the response's tau-th + // quantile at x = 0. + Coefficients []float64 + // Fitted and Residuals align with the rows of the design. + Fitted []float64 + Residuals []float64 + // Tau is the quantile the fit minimises the check loss for. + Tau float64 + // CheckLoss is the minimised check loss Σᵢ ρ_τ(rᵢ) at the fit. + CheckLoss float64 + // Objective records the best check loss seen after every + // interior-point iteration, starting from the ordinary least + // squares start. It is the instrument that shows the optimisation + // descending: monotone non-increasing by construction, because an + // iterate only enters the record by improving on every iterate + // before it. + Objective []float64 + // Iterations counts the interior-point iterations taken; + // Converged reports whether the run settled by its own stopping + // rules. + Iterations int + Converged bool +} + +// QuantileRegression fits y = X·β for the tau-th conditional quantile +// by the Frisch-Newton interior-point method on the dual of the check +// loss program. The design carries n rows and p columns exactly as +// LinearRegression's, the intercept included by the caller as a +// constant column when wanted, and the same validations apply: n > p, +// a full-rank design, real finite input. tau must lie strictly inside +// (0, 1): the closed ends have no regression answer, tau = 0 and +// tau = 1 being the envelope of the data rather than a fit. +// +// The run starts from the ordinary least squares fit and follows the +// barrier path in the dual, where the equality constraint is satisfied +// exactly from the first iterate to the last. The primal fit is read +// back off the dual's interior set, the observations the optimal fit +// reproduces exactly, by least squares over that set; a degenerate +// optimum, whose interior set is underdetermined, is resolved to the +// smallest-norm member of the optimal face by a ridge of the order of +// the solve's own rounding. +func QuantileRegression(x, y *core.Array, tau float64) (*QuantileRegressionResult, error) { + const name = "QuantileRegression" + if x.NDim() != 2 { + return nil, base.Errf("%s: the design must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 { + return nil, base.Errf("%s: the response must be rank 1", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + n, p := x.Shape()[0], x.Shape()[1] + if y.Len() != n { + return nil, base.Errf("%s: the design has %d rows but the response %d", name, n, y.Len()) + } + if n <= p { + return nil, base.Errf("%s: need n > p, got %d observations and %d columns", name, n, p) + } + if p == 0 { + return nil, base.Errf("%s: the design must carry at least one column", name) + } + // NaN compares false against both bounds, so this refuses it too. + if !(tau > 0 && tau < 1) { + return nil, base.Errf("%s: tau must lie strictly inside (0, 1), got %g", name, tau) + } + if err := checkFinite(name, "the design", x); err != nil { + return nil, err + } + if err := checkFinite(name, "the response", y); err != nil { + return nil, err + } + fx := rawFloats(x) + fy := rawFloats(y) + if fy == nil { + // The response is a view or a narrower dtype: widen it once so + // the barrier loops below sweep a plain slice, the same bits + // the widening accessor returns. + fy = make([]float64, n) + for i := range n { + fy[i] = y.FloatAt(i) + } + } + + // The start is the ordinary least squares answer. It is not the + // optimum of any asymmetric loss, but it is inside the basin, its + // residual scale sets the barrier's starting parameter, and a + // design the shared LU cannot solve is refused here, rank + // deficiency and all. + beta, err := leastSquaresSolve(name, n, p, x, y, fx, fy, nil) + if err != nil { + return nil, err + } + residuals := make([]float64, n) + regressionResiduals(n, p, x, y, fx, fy, beta, residuals) + meanAbs := 0.0 + for _, r := range residuals { + meanAbs += math.Abs(r) + } + meanAbs /= float64(n) + lossAt := func(b []float64) float64 { + total := 0.0 + for i := range n { + r := fy[i] + for j := range p { + var xj float64 + if fx != nil { + xj = fx[i*p+j] + } else { + xj = x.FloatAt(i*p + j) + } + r -= b[j] * xj + } + if r >= 0 { + total += r * tau + } else { + // r and tau − 1 are both negative, so the product is + // the positive |r|·(1 − tau) the check loss asks for. + total += r * (tau - 1) + } + } + return total + } + loss := lossAt(beta) + out := &QuantileRegressionResult{ + Coefficients: beta, + Tau: tau, + CheckLoss: loss, + Objective: []float64{loss}, + } + if meanAbs == 0 { + // The ordinary least squares fit reproduces the response + // exactly: the check loss is zero and no asymmetric loss can + // beat zero. The run is settled before it begins. + out.Fitted = make([]float64, n) + out.Residuals = make([]float64, n) + copy(out.Residuals, residuals) + for i := range n { + out.Fitted[i] = fy[i] + } + out.Converged = true + return out, nil + } + + // The dual barrier state. w = 1−τ for every observation satisfies + // the equality Xᵀw = (1−τ)·Xᵀ1 identically and sits exactly in the + // middle of the box, the interior point the method asks for. + mu := meanAbs + muFloor := mu * quantileMuFloorRatio + w := make([]float64, n) + dw := make([]float64, n) + gradient := make([]float64, n) + curvature := make([]float64, n) + cols := make([]float64, n) + colSums := make([]float64, p) + // The backtracking buffer holds only complete candidate vectors: every + // iteration rewrites all n slots before the objective reads any, so one + // buffer serves the whole run. + trial := make([]float64, n) + for i := range n { + w[i] = 1 - tau + } + for i := range n { + for j := range p { + var xj float64 + if fx != nil { + xj = fx[i*p+j] + } else { + xj = x.FloatAt(i*p + j) + } + colSums[j] += xj + } + } + bestBeta := append([]float64(nil), beta...) + bestLoss := loss + converged := false + iterations := quantileMaxIterations + // The barrier system's storage rides the whole run: the upper + // triangle is refilled by accumulation from an explicit zero and + // its mirror copies the lower one, and the right-hand wrapper is + // fixed because the solve writes through rhs in place. + normal := make([][]float64, p) + for j := range p { + normal[j] = make([]float64, p) + } + rhs := make([]float64, p) + solveRHS := [][]float64{rhs} + // The barrier objective Φ = wᵀy + μΣ(log w + log(1−w)) rises + // along the Newton direction of this concave problem by + // ∇ΦᵀΔw = ΔwᵀDΔw ≥ 0 identically, so a backtracking half-step + // finds an ascent all the way to the barrier maximum. + objective := func(wv []float64, muv float64) float64 { + total := 0.0 + for i := range n { + total += wv[i]*fy[i] + muv*(math.Log(wv[i])+math.Log(1-wv[i])) + } + return total + } + for iter := 1; iter <= quantileMaxIterations; iter++ { + // The Newton system of the barrier problem, eliminated to the + // p×p form XᵀD⁻¹X Δλ = XᵀD⁻¹z: z the barrier gradient + // y + μ(1/w − 1/(1−w)), D⁻¹ the reciprocal of the barrier's + // curvature μ((1−w)² + w²)/(w²(1−w)²). The solve is exact, the + // equality direction XᵀΔw = 0 comes out of it, and the step + // keeps the equality exactly because it started exact. + for j := range p { + nj := normal[j] + for k := j; k < p; k++ { + nj[k] = 0 + } + } + clear(rhs) + for i := range n { + if fx != nil { + for j := range p { + cols[j] = fx[i*p+j] + } + } else { + for j := range p { + cols[j] = x.FloatAt(i*p + j) + } + } + oneMinus := 1 - w[i] + gradient[i] = fy[i] + mu*(1/w[i]-1/oneMinus) + curvature[i] = w[i] * w[i] * oneMinus * oneMinus / + (mu * (oneMinus*oneMinus + w[i]*w[i])) + scale := gradient[i] * curvature[i] + for j := range p { + rhs[j] += cols[j] * scale + for k := j; k < p; k++ { + normal[j][k] += cols[j] * cols[k] * curvature[i] + } + } + } + for j := range p { + for k := range j { + normal[j][k] = normal[k][j] + } + } + solved, err := base.SolveSystem(name, normal, solveRHS) + if err != nil { + return nil, base.Errf("%s: the barrier system is singular (%w)", name, err) + } + deltaLambda := solved[0] + mostMove := 0.0 + for i := range n { + pred := 0.0 + if fx != nil { + for j := range p { + pred += fx[i*p+j] * deltaLambda[j] + } + } else { + for j := range p { + pred += x.FloatAt(i*p+j) * deltaLambda[j] + } + } + dw[i] = curvature[i] * (gradient[i] - pred) + if m := math.Abs(dw[i]); m > mostMove { + mostMove = m + } + } + // Fraction to the boundary: the largest step that keeps every + // w strictly inside the box, taken at a documented fraction of + // it and never more than the full Newton step. + t := 1.0 + for i := range n { + if dw[i] < 0 { + t = min(t, quantileStepFraction*(-w[i]/dw[i])) + } + if dw[i] > 0 { + t = min(t, quantileStepFraction*((1-w[i])/dw[i])) + } + } + // The barrier objective rises along the Newton direction of + // this concave problem (the note above the closure), so a + // backtracking half-step finds an ascent all the way to the + // barrier maximum. + current := objective(w, mu) + for { + for i := range n { + trial[i] = w[i] + t*dw[i] + } + if objective(trial, mu) >= current || t < quantileTolerance { + break + } + t /= 2 + } + if t < quantileTolerance || mostMove < quantileTolerance { + // The barrier problem is stationary at this μ: further + // pressure buys nothing until μ falls, and μ is about to. + // The run settles when the barrier itself is exhausted. + converged = true + iterations = iter + break + } + copy(w, trial) + // The primal fit read off this iterate's interior set, and the + // monotone record it feeds. + candidate, err := quantileRecover(name, n, p, x, y, fx, fy, w) + if err != nil { + return nil, err + } + if l := lossAt(candidate); l < bestLoss { + bestLoss = l + copy(bestBeta, candidate) + } + out.Objective = append(out.Objective, bestLoss) + mu *= quantileMuShrink + if mu < muFloor { + converged = true + iterations = iter + break + } + } + if !converged { + return nil, base.Errf("%s: %d iterations did not converge", name, quantileMaxIterations) + } + out.Iterations = iterations + out.Converged = true + out.CheckLoss = bestLoss + copy(out.Coefficients, bestBeta) + out.Fitted = make([]float64, n) + out.Residuals = make([]float64, n) + regressionResiduals(n, p, x, y, fx, fy, out.Coefficients, out.Residuals) + for i := range n { + out.Fitted[i] = fy[i] - out.Residuals[i] + } + return out, nil +} + +// quantileRecover reads a primal fit off a dual barrier point. The +// observations whose w sits strictly inside the box are the ones the +// optimal fit reproduces exactly, their residual being zero, so the +// least squares over that interior set is the fit; the degenerate +// case, an interior set that underdetermines the coefficients, is +// resolved to the smallest-norm member of the optimal face by a ridge +// a trillionth of the normal equations' own scale, far below the +// solve's meaningful digits. An empty interior set falls back to the +// plain least squares fit, the same reading the first iterate's +// all-interior box gives. +func quantileRecover(name string, n, p int, x, y *core.Array, fx, fy []float64, w []float64) ([]float64, error) { + const band = 1e-5 + normal := make([][]float64, p) + for j := range p { + normal[j] = make([]float64, p) + } + rhs := make([]float64, p) + trace := 0.0 + count := 0 + for i := range n { + if !(w[i] > band && w[i] < 1-band) { + continue + } + count++ + var yv float64 + if fy != nil { + yv = fy[i] + } else { + yv = y.FloatAt(i) + } + for a := range p { + var xa float64 + if fx != nil { + xa = fx[i*p+a] + } else { + xa = x.FloatAt(i*p + a) + } + rhs[a] += xa * yv + for b := range p { + var xb float64 + if fx != nil { + xb = fx[i*p+b] + } else { + xb = x.FloatAt(i*p + b) + } + normal[a][b] += xa * xb + } + } + } + if count < p { + // An interior set too small to identify the coefficients: the + // whole box was effectively at its bounds, and the plain least + // squares fit is the honest reading. The caller's + // best-by-check-loss record keeps this from ever worsening the + // answer. + return leastSquaresSolve(name, n, p, x, y, fx, fy, nil) + } + for a := range p { + trace += normal[a][a] + } + for a := range p { + normal[a][a] += 1e-12 * trace / float64(p) + } + solved, err := base.SolveSystem(name, normal, [][]float64{rhs}) + if err != nil { + return nil, base.Errf("%s: the interior set is degenerate (%w)", name, err) + } + return solved[0], nil +} diff --git a/stats/quantilereg_test.go b/stats/quantilereg_test.go new file mode 100644 index 0000000..7991334 --- /dev/null +++ b/stats/quantilereg_test.go @@ -0,0 +1,296 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// quantileSubgradient evaluates the check loss's subgradient at a fit +// and the slack the zero-band residuals buy it. At the optimum of +// Σρ_τ(y − Xβ) there must be a choice of s ∈ [τ−1, τ] per observation +// with Xᵀs = 0; residuals above the band fix their s (τ above the +// fit, τ−1 below it), and residuals within 1e-6 of zero, where no +// floating point fit can place them exactly, are free to slide in +// [τ−1, τ]. The certificate passes when every column's forced +// gradient is within the slack its zero-band residuals can absorb. +func quantileSubgradient(t *testing.T, x, y *core.Array, beta []float64, tau float64) (grad, free []float64) { + t.Helper() + const band = 1e-6 + n, p := x.Shape()[0], x.Shape()[1] + grad = make([]float64, p) + free = make([]float64, p) + for i := range n { + r := y.FloatAt(i) + for j := range p { + r -= beta[j] * x.FloatAt(i*p+j) + } + var s float64 + switch { + case r > band: + s = tau + case r < -band: + s = tau - 1 + default: + for j := range p { + free[j] += math.Abs(x.FloatAt(i*p + j)) + } + continue + } + for j := range p { + grad[j] += x.FloatAt(i*p+j) * s + } + } + for j := range p { + free[j] *= math.Max(tau, 1-tau) + } + return grad, free +} + +// quantileHeteroFixture builds the seeded heteroscedastic model the +// coverage and monotonicity pins share: x uniform on [−2, 2], the +// response moved by 1 + 2x and widened by a scale that grows with +// |x|, so the tau-th conditional quantile is a genuinely different +// line from the mean's and no single slope fits every tau. +func quantileHeteroFixture(t *testing.T, n int, seed int64) (*core.Array, *core.Array) { + t.Helper() + g := core.NewGenerator(seed) + vals := make([]float64, 0, 2*n) + yv := make([]float64, 0, n) + for range n { + x := 4*g.Unit() - 2 + y := 1 + 2*x + (0.3+0.3*math.Abs(x))*g.NormalUnit() + vals = append(vals, 1, x) + yv = append(yv, y) + } + return mustFromFloats(t, vals, n, 2), mustFromFloats(t, yv, n) +} + +// TestQuantileMedianCertificate pins tau = 0.5 against the L1 median +// regression answer, twice over. An intercept-only design has the +// sample median for its answer, and the general design is verified +// against the subgradient optimality certificate, which is the exact +// first-order condition of the check loss rather than another +// algorithm's output. +func TestQuantileMedianCertificate(t *testing.T) { + // Intercept only, odd sample: the fit is the median, the single + // point every absolute deviation bends toward. + const n = 25 + g := core.NewGenerator(19) + ys := make([]float64, n) + for i := range n { + ys[i] = math.Round(40*g.Unit() - 20) + } + y := mustFromFloats(t, ys, n) + ones := make([]float64, n) + for i := range n { + ones[i] = 1 + } + design := mustFromFloats(t, ones, n, 1) + res, err := QuantileRegression(design, y, 0.5) + if err != nil { + t.Fatalf("QuantileRegression: %v", err) + } + if !res.Converged { + t.Fatalf("the intercept-only fit did not converge") + } + median, err := Median(y) + if err != nil { + t.Fatalf("Median: %v", err) + } + if math.Abs(res.Coefficients[0]-median) > 1e-7 { + t.Fatalf("the tau = 0.5 intercept-only fit is %.12g, want the median %.12g", res.Coefficients[0], median) + } + // A general design against the subgradient certificate. + xs := make([]float64, 0, 2*n) + yv := make([]float64, 0, n) + for range n { + x := 4*g.Unit() - 2 + yv = append(yv, 1+2*x+0.4*g.NormalUnit()) + xs = append(xs, 1, x) + } + design2 := mustFromFloats(t, xs, n, 2) + y2 := mustFromFloats(t, yv, n) + fit, err := QuantileRegression(design2, y2, 0.5) + if err != nil { + t.Fatalf("QuantileRegression: %v", err) + } + grad, free := quantileSubgradient(t, design2, y2, fit.Coefficients, 0.5) + for j := range 2 { + t.Logf("tau = 0.5 column %d: forced gradient %.3g against slack %.3g", j, grad[j], free[j]) + if math.Abs(grad[j]) > free[j]+1e-6 { + t.Fatalf("column %d fails the subgradient certificate: %.3g against slack %.3g", j, grad[j], free[j]) + } + } +} + +// TestQuantileObjectiveMonotone instruments the interior-point run on +// the heteroscedastic fixture: the recorded objective is the running +// best check loss, monotone non-increasing from the ordinary least +// squares start, and the record has to show real descent: the +// tau = 0.9 quantile line is not the mean line, and the check loss at +// the optimum must sit measurably below the start. +func TestQuantileObjectiveMonotone(t *testing.T) { + design, y := quantileHeteroFixture(t, 400, 23) + fit, err := QuantileRegression(design, y, 0.9) + if err != nil { + t.Fatalf("QuantileRegression: %v", err) + } + if !fit.Converged { + t.Fatalf("the interior-point run did not converge") + } + if len(fit.Objective) < 3 { + t.Fatalf("the run recorded only %d objectives, too few to show a descent", len(fit.Objective)) + } + for k := 1; k < len(fit.Objective); k++ { + if fit.Objective[k] > fit.Objective[k-1] { + t.Fatalf("the objective rose at record %d: %.12g after %.12g", + k, fit.Objective[k], fit.Objective[k-1]) + } + } + if math.Abs(fit.Objective[len(fit.Objective)-1]-fit.CheckLoss) > 1e-9 { + t.Fatalf("the record ends at %.12g but the fit reports %.12g", + fit.Objective[len(fit.Objective)-1], fit.CheckLoss) + } + t.Logf("check loss descended from %.6f to %.6f over %d iterations", + fit.Objective[0], fit.CheckLoss, fit.Iterations) + if fit.Objective[0]-fit.CheckLoss < 1 { + t.Fatalf("the tau = 0.9 fit barely left the mean fit: descent %.6g", fit.Objective[0]-fit.CheckLoss) + } +} + +// TestQuantileTracksConditionalQuantile measures the fit where it +// matters: on six hundred held-out draws of the heteroscedastic model, +// the tau = 0.9 line must cover ninety percent of the responses and +// the tau = 0.1 line ten percent, the defining property of a +// conditional quantile estimate. +func TestQuantileTracksConditionalQuantile(t *testing.T) { + design, y := quantileHeteroFixture(t, 400, 23) + fit90, err := QuantileRegression(design, y, 0.9) + if err != nil { + t.Fatalf("QuantileRegression tau 0.9: %v", err) + } + fit10, err := QuantileRegression(design, y, 0.1) + if err != nil { + t.Fatalf("QuantileRegression tau 0.1: %v", err) + } + if !fit90.Converged || !fit10.Converged { + t.Fatalf("a coverage fit did not converge") + } + // The fitted slopes track the true conditional quantile slope 2. + if math.Abs(fit90.Coefficients[1]-2) > 0.5 { + t.Fatalf("the tau = 0.9 slope is %.4f, far from the truth 2", fit90.Coefficients[1]) + } + g := core.NewGenerator(97) + const held = 600 + below90, below10 := 0, 0 + for range held { + x := 4*g.Unit() - 2 + yi := 1 + 2*x + (0.3+0.3*math.Abs(x))*g.NormalUnit() + q90 := fit90.Coefficients[0] + fit90.Coefficients[1]*x + q10 := fit10.Coefficients[0] + fit10.Coefficients[1]*x + if yi <= q90 { + below90++ + } + if yi <= q10 { + below10++ + } + } + coverage90 := float64(below90) / held + coverage10 := float64(below10) / held + t.Logf("held-out coverage: tau = 0.9 covers %.3f, tau = 0.1 covers %.3f", coverage90, coverage10) + if coverage90 < 0.84 || coverage90 > 0.96 { + t.Fatalf("the tau = 0.9 fit covers %.3f of held-out draws, want near 0.9", coverage90) + } + if coverage10 < 0.04 || coverage10 > 0.16 { + t.Fatalf("the tau = 0.1 fit covers %.3f of held-out draws, want near 0.1", coverage10) + } +} + +// TestQuantileExactFit walks the zero-loss corner: a constant response +// is reproduced exactly by the start, the check loss is zero, and the +// run is settled before the first interior-point iteration. +func TestQuantileExactFit(t *testing.T) { + design := mustFromFloats(t, []float64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4}, 5, 2) + y := mustFromFloats(t, []float64{4.2, 4.2, 4.2, 4.2, 4.2}, 5) + res, err := QuantileRegression(design, y, 0.3) + if err != nil { + t.Fatalf("QuantileRegression: %v", err) + } + if !res.Converged { + t.Fatalf("the exact fit did not report convergence") + } + if res.CheckLoss != 0 { + t.Fatalf("the exact fit reports check loss %g, want 0", res.CheckLoss) + } + if res.Iterations != 0 { + t.Fatalf("the exact fit spent %d iterations, want 0", res.Iterations) + } + for i := range 5 { + if math.Abs(res.Fitted[i]-4.2) > 1e-9 || math.Abs(res.Residuals[i]) > 1e-9 { + t.Fatalf("the exact fit moved row %d: fitted %.12g", i, res.Fitted[i]) + } + } +} + +// TestQuantileOnIntegerArrays exercises the widening accessor's +// fallback paths in the loss and the recovery: an integer design and +// response reach the fit through FloatAt rather than a raw float +// payload. +func TestQuantileOnIntegerArrays(t *testing.T) { + design := mustFromInts(t, []int64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, 1, 6}, 7, 2) + y := mustFromInts(t, []int64{2, 3, 4, 5, 6, 7, 8}, 7) + res, err := QuantileRegression(design, y, 0.75) + if err != nil { + t.Fatalf("QuantileRegression on integer input: %v", err) + } + if !res.Converged { + t.Fatalf("the integer-input fit did not converge") + } + grad, free := quantileSubgradient(t, design, y, res.Coefficients, 0.75) + for j := range 2 { + if math.Abs(grad[j]) > free[j]+1e-6 { + t.Fatalf("column %d fails the subgradient certificate: %.3g against slack %.3g", j, grad[j], free[j]) + } + } +} + +// TestQuantileValidation refuses the malformed inputs: tau outside the +// open interval, wrong shapes, non-finite samples and a rank-deficient +// design. +func TestQuantileValidation(t *testing.T) { + design := mustFromFloats(t, []float64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5}, 6, 2) + y := mustFromFloats(t, []float64{1, 3, 2, 5, 4, 7}, 6) + for _, tau := range []float64{0, 1, -0.5, 1.5, math.NaN()} { + if _, err := QuantileRegression(design, y, tau); err == nil || !strings.Contains(err.Error(), "strictly inside") { + t.Fatalf("tau = %g: got %v, want the tau refusal", tau, err) + } + } + if _, err := QuantileRegression(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 6), y, 0.5); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("a rank 1 design: got %v, want the rank refusal", err) + } + if _, err := QuantileRegression(design, mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2), 0.5); err == nil || !strings.Contains(err.Error(), "must be rank 1") { + t.Fatalf("a rank 2 response: got %v, want the rank refusal", err) + } + if _, err := QuantileRegression(design, mustFromFloats(t, []float64{1, 2, 3}, 3), 0.5); err == nil || !strings.Contains(err.Error(), "rows but the response") { + t.Fatalf("a length mismatch: got %v, want the length refusal", err) + } + if _, err := QuantileRegression(mustFromFloats(t, []float64{1, 0, 1, 1}, 2, 2), mustFromFloats(t, []float64{1, 2}, 2), 0.5); err == nil || !strings.Contains(err.Error(), "n > p") { + t.Fatalf("n = p: got %v, want the n > p refusal", err) + } + if _, err := QuantileRegression(design, mustFromFloats(t, []float64{1, 2, 3, math.NaN(), 5, 7}, 6), 0.5); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("a non-finite response: got %v, want the non-finite refusal", err) + } + if _, err := QuantileRegression(core.New(core.Complex, 6, 2), y, 0.5); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("complex input: got %v, want the complex refusal", err) + } + singular := mustFromFloats(t, []float64{1, 2, 3, 1, 2, 3, 1, 2, 3, 1, 2, 3}, 4, 3) + if _, err := QuantileRegression(singular, mustFromFloats(t, []float64{1, 2, 3, 4}, 4), 0.5); err == nil || !strings.Contains(err.Error(), "singular") { + t.Fatalf("a rank-deficient design: got %v, want the singular-system refusal", err) + } +} diff --git a/stats/rankcorr.go b/stats/rankcorr.go new file mode 100644 index 0000000..bd4e817 --- /dev/null +++ b/stats/rankcorr.go @@ -0,0 +1,278 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import ( + "cmp" + "math" + "slices" +) + +// Rank correlations over paired observations: Spearman's ρ and +// Kendall's τ-b. Both refuse length mismatches, fewer than two pairs, +// complex or non-finite input, the contract every inference entry +// point of the package shares: a NaN would otherwise flow silently +// through the ranks and the pair counts. + +// SpearmanRho returns Spearman's rank correlation of the paired +// samples x and y: Pearson's correlation evaluated on the mid-ranks, +// the averaged ranks ties share. Computing it as the Pearson of the +// mid-ranks is exactly the tie-corrected form, so tied observations +// lower ρ honestly instead of pretending a finer ordering than the +// data carry. A sample whose ranks have zero variance (all values +// equal) is refused: the correlation is undefined, not zero. +func SpearmanRho(x, y *core.Array) (float64, error) { + const name = "SpearmanRho" + rx, ry, err := pairedSamples(name, x, y) + if err != nil { + return 0, err + } + rankX := midRanks(rx) + rankY := midRanks(ry) + n := float64(len(rx)) + mx, my := 0.0, 0.0 + for i := range rx { + mx += rankX[i] + my += rankY[i] + } + mx /= n + my /= n + sxx, syy, sxy := 0.0, 0.0, 0.0 + for i := range rx { + dx := rankX[i] - mx + dy := rankY[i] - my + sxx += dx * dx + syy += dy * dy + sxy += dx * dy + } + if sxx == 0 || syy == 0 { + return 0, base.Errf("%s: a sample whose ranks all agree has no rank correlation", name) + } + return sxy / math.Sqrt(sxx*syy), nil +} + +// KendallTau returns Kendall's τ-b for the paired samples x and y: the +// concordant minus discordant pair count over the tie-corrected +// denominator √((n₀−n₁)(n₀−n₂)), with n₀ the pair total and n₁, n₂ the +// within-sample tied pair counts. The normalisation keeps τ inside +// [−1, 1] under ties and reaches ±1 on every perfectly monotone +// pairing, which the naive n₀ denominator cannot; pairs tied in both +// samples count for neither side. +func KendallTau(x, y *core.Array) (float64, error) { + const name = "KendallTau" + xs, ys, err := pairedSamples(name, x, y) + if err != nil { + return 0, err + } + // The pair counts are exact well below 2⁵³ pairs, so the walks run + // in float64 without an overflow thought. + n0 := float64(len(xs)) * float64(len(xs)-1) / 2 + tally := tallyPairs(xs, ys) + n1 := float64(tally.tiedX) + n2 := float64(tally.tiedY) + if n0 == n1 || n0 == n2 { + return 0, base.Errf("%s: a constant sample leaves the denominator zero", name) + } + concordant := float64(tally.concordant) + discordant := float64(tally.discordant) + return (concordant - discordant) / math.Sqrt((n0-n1)*(n0-n2)), nil +} + +// pairedSamples extracts two equal-length real samples, the shared +// gate of both rank correlations: complex input, a length mismatch, +// fewer than two pairs and non-finite values are all refused here, +// under the entry point's own name. +func pairedSamples(name string, x, y *core.Array) ([]float64, []float64, error) { + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, nil, base.Errf("%s: complex samples are not supported", name) + } + n := x.Len() + if y.Len() != n { + return nil, nil, base.Errf("%s: the samples have %d and %d observations", name, n, y.Len()) + } + if n < 2 { + return nil, nil, base.Errf("%s: at least two paired observations are needed, got %d", name, n) + } + if err := checkFinite(name, "the first sample", x); err != nil { + return nil, nil, err + } + if err := checkFinite(name, "the second sample", y); err != nil { + return nil, nil, err + } + xs := make([]float64, n) + if fs := rawFloats(x); fs != nil { + copy(xs, fs) + } else { + for i := range n { + xs[i] = x.FloatAt(i) + } + } + ys := make([]float64, n) + if fs := rawFloats(y); fs != nil { + copy(ys, fs) + } else { + for i := range n { + ys[i] = y.FloatAt(i) + } + } + return xs, ys, nil +} + +// midRanks returns the mid-ranks of vals: 1-based ranks with ties +// sharing the average of the ranks they span, the ordering both +// correlations count on. +func midRanks(vals []float64) []float64 { + order := make([]int, len(vals)) + for i := range order { + order[i] = i + } + slices.SortFunc(order, func(a, b int) int { return cmp.Compare(vals[a], vals[b]) }) + ranks := make([]float64, len(vals)) + for i := 0; i < len(order); { + j := i + for j < len(order) && vals[order[j]] == vals[order[i]] { + j++ + } + // The average of the 1-based ranks i+1 through j. + mid := float64(i+j+1) / 2 + for k := i; k < j; k++ { + ranks[order[k]] = mid + } + i = j + } + return ranks +} + +// rankPair is one paired observation, the sort key of the pair walk. +type rankPair struct{ x, y float64 } + +// pairTally holds the exact pair counts of a paired sample: the +// concordant and discordant pairs and the pairs tied within the first +// and within the second sample. Every field is the integer the O(n²) +// enumeration would produce, so the τ-b numerator and denominator built +// from them carry the same bits. +type pairTally struct { + concordant int64 + discordant int64 + tiedX int64 + tiedY int64 +} + +// tallyPairs counts the pairs of two paired samples in O(n log n). The +// pairs are first ordered by (x, y): a pair is then discordant exactly +// when the reordered y sequence descends, which the merge sort of +// sortTally counts as it sorts, and every remaining pair is either +// ascending in y or tied in y. Ascending pairs inside a block of equal +// x are tied in the predictor and count for neither side either, so the +// block's ordered pairs, its ties and its pairs tied in both are +// removed by one scan of the sorted pairs. A pair tied in either sample +// counts for neither side, which is the definition the enumeration +// applies by testing the product of the two differences against zero. +func tallyPairs(xs, ys []float64) pairTally { + n := len(xs) + pairs := make([]rankPair, n) + for i := range n { + pairs[i] = rankPair{x: xs[i], y: ys[i]} + } + slices.SortFunc(pairs, func(a, b rankPair) int { + if c := cmp.Compare(a.x, b.x); c != 0 { + return c + } + return cmp.Compare(a.y, b.y) + }) + ordered := make([]float64, n) + for i := range pairs { + ordered[i] = pairs[i].y + } + var tally pairTally + tally.discordant, tally.tiedY = sortTally(ordered) + total := int64(n) * int64(n-1) / 2 + // The blocks of equal x hold their y values ascending, so no pair + // inside one is an inversion: the block's pairs are tied in x, and + // the concordant count is what is left of the ascending pairs once + // the block's own ordered pairs are taken out. + tiedBoth := int64(0) + for lo := 0; lo < n; { + hi := lo + 1 + for hi < n && pairs[hi].x == pairs[lo].x { + hi++ + } + m := int64(hi - lo) + tally.tiedX += m * (m - 1) / 2 + for i := lo; i < hi; { + j := i + 1 + for j < hi && pairs[j].y == pairs[i].y { + j++ + } + c := int64(j - i) + tiedBoth += c * (c - 1) / 2 + i = j + } + lo = hi + } + ascending := total - tally.discordant - tally.tiedY + tally.concordant = ascending - (tally.tiedX - tiedBoth) + return tally +} + +// sortTally sorts vals by merge sort and returns the number of pairs +// i < j with vals[i] > vals[j] and the number of pairs i < j with +// vals[i] == vals[j]. Both are exact integers, so the counts are the +// ones a brute-force enumeration of the pairs produces, at O(n log n) +// cost: an element taken from the right run descends past every left +// element still waiting, and the equal pairs are the groups the +// finished sort leaves. vals is consumed as scratch and left sorted. +func sortTally(vals []float64) (descending, equal int64) { + n := len(vals) + if n < 2 { + return 0, 0 + } + buf := make([]float64, n) + src, dst := vals, buf + for width := 1; width < n; width *= 2 { + for lo := 0; lo < n; lo += 2 * width { + mid := min(lo+width, n) + hi := min(lo+2*width, n) + i, j, k := lo, mid, lo + for i < mid && j < hi { + if src[i] <= src[j] { + dst[k] = src[i] + i++ + } else { + descending += int64(mid - i) + dst[k] = src[j] + j++ + } + k++ + } + for i < mid { + dst[k] = src[i] + i++ + k++ + } + for j < hi { + dst[k] = src[j] + j++ + k++ + } + } + // Every position was written, so src holds the merged sequence. + src, dst = dst, src + } + for i := 0; i < n; { + j := i + 1 + for j < n && src[j] == src[i] { + j++ + } + c := int64(j - i) + equal += c * (c - 1) / 2 + i = j + } + return descending, equal +} diff --git a/stats/rankcorr_test.go b/stats/rankcorr_test.go new file mode 100644 index 0000000..e97084a --- /dev/null +++ b/stats/rankcorr_test.go @@ -0,0 +1,127 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" +) + +// TestSpearmanRho pins ρ on exact cases: a perfect monotone pairing is +// exactly ±1, the classic d² example x = (1..5) against +// y = (3, 1, 4, 2, 5) has rank differences (−2, 1, −1, 2, 0), so +// Σd² = 10 and ρ = 1 − 60/120 = 0.5, and the tied pairing +// x = (1, 2, 2, 4) against y = (1, 2, 3, 4) works out through the +// mid-ranks to √0.9. +func TestSpearmanRho(t *testing.T) { + x := mustFloats(t, []float64{1, 2, 3, 4, 5}) + y := mustFloats(t, []float64{10, 20, 30, 40, 50}) + rho, err := SpearmanRho(x, y) + if err != nil || rho != 1 { + t.Fatalf("perfect monotone ρ = %v, %v, want exactly 1", rho, err) + } + rho, err = SpearmanRho(x, mustFloats(t, []float64{50, 40, 30, 20, 10})) + if err != nil || rho != -1 { + t.Fatalf("reversed ρ = %v, %v, want exactly −1", rho, err) + } + rho, err = SpearmanRho(x, mustFloats(t, []float64{3, 1, 4, 2, 5})) + if err != nil || math.Abs(rho-0.5) > 1e-14 { + t.Fatalf("classic example ρ = %v, %v, want 0.5", rho, err) + } + rho, err = SpearmanRho( + mustFloats(t, []float64{1, 2, 2, 4}), + mustFloats(t, []float64{1, 2, 3, 4})) + if err != nil || math.Abs(rho-math.Sqrt(0.9)) > 1e-14 { + t.Fatalf("tied ρ = %.16g, %v, want √0.9 = %.16g", rho, err, math.Sqrt(0.9)) + } + // The tie correction is honest: the same pairing without the tie in + // x ranks higher. + untied, _ := SpearmanRho(mustFloats(t, []float64{1, 2, 3, 4}), + mustFloats(t, []float64{1, 2, 3, 4})) + if untied != 1 || rho >= 1 { + t.Fatalf("ties did not lower ρ: tied %v, untied %v", rho, untied) + } + if _, err := SpearmanRho(x, mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "the samples have") { + t.Fatalf("length mismatch: got %v, want the length refusal", err) + } + if _, err := SpearmanRho(mustFloats(t, []float64{1}), mustFloats(t, []float64{1})); err == nil || !strings.Contains(err.Error(), "at least two paired") { + t.Fatalf("one pair: got %v, want the pairing floor refusal", err) + } + if _, err := SpearmanRho(mustFloats(t, []float64{1, 2, 3}), mustFloats(t, []float64{5, 5, 5})); err == nil || !strings.Contains(err.Error(), "ranks all agree") { + t.Fatalf("a constant sample: got %v, want the constant-sample refusal", err) + } + if _, err := SpearmanRho(mustFloats(t, []float64{1, math.NaN(), 3}), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("a NaN observation: got %v, want the non-finite refusal", err) + } + if _, err := SpearmanRho(mustComplexes(t, []complex128{1, 2, 3}, 3), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("a complex sample: got %v, want the complex refusal", err) + } +} + +// TestKendallTau pins τ-b on hand-counted cases. For x = (1, 2, 3, 4) +// against y = (1, 3, 2, 4) only the pair (2, 3) against (3, 2) is +// discordant, so τ = (5−1)/6 = 2/3. With the tie x = (1, 2, 2, 4) the +// denominator shrinks to √((6−1)(6−0)) and τ-b = 5/√30; the naive n₀ +// denominator would report 5/6 and miss the attainable maximum. The +// all-tied pairing x = (1, 1, 2, 2) against y = (5, 5, 7, 7) is +// perfectly monotone and must give exactly ±1, where the naive +// normalisation would report 4/6. +func TestKendallTau(t *testing.T) { + tau, err := KendallTau( + mustFloats(t, []float64{1, 2, 3, 4}), + mustFloats(t, []float64{1, 3, 2, 4})) + if err != nil || math.Abs(tau-2.0/3) > 1e-15 { + t.Fatalf("hand case τ = %.16g, %v, want 2/3", tau, err) + } + tau, err = KendallTau( + mustFloats(t, []float64{1, 2, 2, 4}), + mustFloats(t, []float64{1, 2, 3, 4})) + if err != nil || math.Abs(tau-5/math.Sqrt(30)) > 1e-14 { + t.Fatalf("tied τ-b = %.16g, %v, want 5/√30 = %.16g", tau, err, 5/math.Sqrt(30)) + } + tau, err = KendallTau( + mustFloats(t, []float64{1, 1, 2, 2}), + mustFloats(t, []float64{5, 5, 7, 7})) + if err != nil || tau != 1 { + t.Fatalf("perfectly monotone with ties τ-b = %v, %v, want exactly 1", tau, err) + } + tau, err = KendallTau( + mustFloats(t, []float64{1, 1, 2, 2}), + mustFloats(t, []float64{7, 7, 5, 5})) + if err != nil || tau != -1 { + t.Fatalf("perfectly reversed with ties τ-b = %v, %v, want exactly −1", tau, err) + } + // A tie-heavy pairing: 3 concordant and 3 discordant cross pairs + // cancel exactly, so τ-b is exactly 0 with the denominator + // √((15−6)(15−3)) well defined. Shifting the second block up turns + // it into 6 concordant, 1 discordant, τ-b = 5/√117. + tau, err = KendallTau( + mustFloats(t, []float64{1, 1, 1, 2, 2, 2}), + mustFloats(t, []float64{1, 2, 3, 1, 2, 3})) + if err != nil || tau != 0 { + t.Fatalf("tie-heavy τ-b = %.16g, %v, want exactly 0", tau, err) + } + tau, err = KendallTau( + mustFloats(t, []float64{1, 1, 1, 2, 2, 2}), + mustFloats(t, []float64{1, 2, 3, 2, 3, 4})) + if err != nil || math.Abs(tau-5/math.Sqrt(117)) > 1e-14 { + t.Fatalf("tie-heavy shifted τ-b = %.16g, %v, want 5/√117", tau, err) + } + if _, err := KendallTau(mustFloats(t, []float64{1, 2, 3}), mustFloats(t, []float64{1, 2})); err == nil || !strings.Contains(err.Error(), "the samples have") { + t.Fatalf("length mismatch: got %v, want the length refusal", err) + } + if _, err := KendallTau(mustFloats(t, []float64{1}), mustFloats(t, []float64{1})); err == nil || !strings.Contains(err.Error(), "at least two paired") { + t.Fatalf("one pair: got %v, want the pairing floor refusal", err) + } + if _, err := KendallTau(mustFloats(t, []float64{4, 4, 4}), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "denominator zero") { + t.Fatalf("a constant sample: got %v, want the constant-sample refusal", err) + } + if _, err := KendallTau(mustFloats(t, []float64{1, math.Inf(1), 3}), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("a non-finite observation: got %v, want the non-finite refusal", err) + } + if _, err := KendallTau(mustFloats(t, []float64{1, 2, 3}), mustComplexes(t, []complex128{1, 2, 3}, 3)); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("a complex second sample: got %v, want the complex refusal", err) + } +} diff --git a/stats/regression.go b/stats/regression.go new file mode 100644 index 0000000..bc2a7d7 --- /dev/null +++ b/stats/regression.go @@ -0,0 +1,579 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Linear regression with classical inference: the ordinary +// least-squares fit together with the uncertainty statement every +// empirical paper needs: standard errors, t-tests on each +// coefficient, R², adjusted R² and the F-test of the model as a +// whole. The linear algebra is the normal equations solved by the +// shared LU: regression designs are small and well conditioned in +// practice, and a caller with a genuinely ill-conditioned design +// should regularise (or reach for linalg.SolveTruncated) rather than +// trust any black-box fit. + +// LinearRegressionResult carries the fit and its inference. Each +// slice is indexed by column of the design matrix, in order. +type LinearRegressionResult struct { + // Coefficients are the least-squares estimates β̂. + Coefficients []float64 + // StandardErrors are the estimated standard deviations of the + // coefficient estimators. + StandardErrors []float64 + // TStatistics are β̂/SE per coefficient. + TStatistics []float64 + // PValues are the two-sided p-values of the t-tests. + PValues []float64 + // ResidualVariance is σ̂² = RSS/(n − p). + ResidualVariance float64 + // RSquared and AdjustedRSquared measure the fit. + RSquared float64 + AdjustedRSquared float64 + // FStatistic with DModel and DResidual as its degrees of freedom + // and FPValue its tail probability. DModel is p − 1 for a design + // with a constant column and p without one; DResidual is n − p. + FStatistic float64 + DModel int + DResidual int + FPValue float64 + // Fitted and Residuals align with the rows of the design. + Fitted []float64 + Residuals []float64 +} + +// LinearRegression fits y = X·β by ordinary least squares over the +// design matrix X (n rows, p columns, the intercept included by the +// caller as a constant column when wanted) and reports the full +// classical inference. Inputs must be rank-2 / rank-1 of matching +// length, real-valued, with n > p and X of full column rank; a +// rank-deficient design is an error naming the condition. +func LinearRegression(x, y *core.Array) (*LinearRegressionResult, error) { + const name = "LinearRegression" + if x.NDim() != 2 { + return nil, base.Errf("%s: the design must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 { + return nil, base.Errf("%s: the response must be rank 1, got shape %s", name, base.ShapeText(y.Shape())) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + n, p := x.Shape()[0], x.Shape()[1] + if y.Len() != n { + return nil, base.Errf("%s: the design has %d rows but the response %d", name, n, y.Len()) + } + if n <= p { + return nil, base.Errf("%s: need n > p, got %d observations and %d columns", name, n, p) + } + // Non-finite input has no answer to report: a single NaN would + // propagate into every coefficient and every statistic, and the + // other tests in the package refuse it for the same reason. Both + // scans are bounded by the visible element counts: a rebased view's + // payload may run past them, and a non-finite slot there is nobody's + // observation. + nVis := x.Len() + fx := rawFloats(x) + fy := rawFloats(y) + if fx == nil { + for i := range x.Len() { + if v := x.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: the design holds the non-finite value %g", name, v) + } + } + } else { + for _, v := range fx[:nVis] { + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: the design holds the non-finite value %g", name, v) + } + } + } + if fy == nil { + for i := range n { + if v := y.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: the response holds the non-finite value %g", name, v) + } + } + } else { + for _, v := range fy[:n] { + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, base.Errf("%s: the response holds the non-finite value %g", name, v) + } + } + } + // Whether the caller supplied an intercept, as a constant column. + // It decides what the null model is: with a constant column it is + // the mean of y (centred total sum of squares), without one it is + // zero and Σy² plays that role. The distinction changes R², the + // model degrees of freedom and the F statistic. + hasConstant := hasConstantColumn(x, n, p) + + // Normal equations: (XᵀX)β = Xᵀy. The row-wise walk below visits + // the rows in the same order the column-wise walk did, so every + // entry sums identical products in identical order; the hoisted + // row value re-reads the same bits the inner loop re-read. XᵀX is + // symmetric and each lower-triangle entry equals its upper twin bit + // for bit (mirrorUpper: the row-wise product commutes bitwise and + // both entries sum the rows in the same order), so the accumulation + // runs the upper triangle alone and mirrors it once. + xtx := make([][]float64, p) + for i := range p { + xtx[i] = make([]float64, p) + } + xty := make([]float64, p) + if fx != nil && fy != nil { + for r := range n { + row := fx[r*p : r*p+p] + yv := fy[r] + for i, xi := range row { + xty[i] += xi * yv + // Upper triangle, both operands pre-sliced from i: the + // same products in the same order, bounds checks elided. + ai := xtx[i][i:] + for j, xj := range row[i:] { + ai[j] += xi * xj + } + } + } + } else { + for r := range n { + yv := y.FloatAt(r) + for i := range p { + xi := x.FloatAt(r*p + i) + xty[i] += xi * yv + for j := i; j < p; j++ { + xtx[i][j] += xi * x.FloatAt(r*p+j) + } + } + } + } + mirrorUpper(xtx) + // SolveSystem factors its matrix in place, so the covariance pass + // below needs a pristine copy of the normal equations. + xtxPristine := make([][]float64, p) + for i := range p { + xtxPristine[i] = append([]float64(nil), xtx[i]...) + } + // SolveSystem consumes columns: one column holding Xᵀy. + solved, err := base.SolveSystem(name, xtx, [][]float64{xty}) + if err != nil { + return nil, base.Errf("%s: the design is rank deficient (%w)", name, err) + } + beta := make([]float64, p) + for i := range p { + beta[i] = solved[0][i] + } + + out := &LinearRegressionResult{Coefficients: beta} + out.Fitted = make([]float64, n) + out.Residuals = make([]float64, n) + rss := 0.0 + tss := 0.0 + uncentred := 0.0 + mean := 0.0 + if fy != nil { + for _, v := range fy[:n] { + mean += v + } + } else { + for i := range n { + mean += y.FloatAt(i) + } + } + mean /= float64(n) + for r := range n { + f := 0.0 + if fx != nil { + row := fx[r*p : r*p+p] + for i, xi := range row { + f += beta[i] * xi + } + } else { + for i := range p { + f += beta[i] * x.FloatAt(r*p+i) + } + } + out.Fitted[r] = f + var res, yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + res = yv - f + out.Residuals[r] = res + rss += res * res + tss += (yv - mean) * (yv - mean) + uncentred += yv * yv + } + if !hasConstant { + // The null model is y = 0, so the uncentred total is what the + // model has to beat, and it carries n degrees of freedom. + tss = uncentred + } + dof := n - p + out.ResidualVariance = rss / float64(dof) + if tss == 0 { + // A constant response reproduced exactly: R² is 1 by the + // perfect-fit convention, not the 1 − 0/0 NaN every consumer + // would propagate. The same guard the F statistic below has. + out.RSquared = 1 + out.AdjustedRSquared = 1 + } else { + out.RSquared = 1 - rss/tss + tssDOF := n - 1 + if !hasConstant { + tssDOF = n + } + out.AdjustedRSquared = 1 - (rss/float64(dof))/(tss/float64(tssDOF)) + } + out.DModel = p - 1 + if !hasConstant { + out.DModel = p + } + out.DResidual = dof + + // Covariance of β̂: σ̂²(XᵀX)⁻¹, its diagonal read from one + // factorisation of the pristine normal equations against all p unit + // columns at once. Solving one unit vector per coefficient + // refactors the same matrix p times; the shared solve factors once + // and substitutes each column through the identical factor, so the + // diagonal is the one the p separate solves produced, bit for bit. + out.StandardErrors = make([]float64, p) + out.TStatistics = make([]float64, p) + out.PValues = make([]float64, p) + unit := make([][]float64, p) + for j := range p { + unit[j] = make([]float64, p) + unit[j][j] = 1 + } + inv, err := base.SolveSystem(name, xtxPristine, unit) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + for j := range p { + v := out.ResidualVariance * inv[j][j] // σ̂²·(XᵀX)⁻¹_jj + switch { + case v > 0: + se := math.Sqrt(v) + out.StandardErrors[j] = se + out.TStatistics[j] = beta[j] / se + pv, err := twoSidedT(out.TStatistics[j], dof) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + out.PValues[j] = pv + case v == 0: + // An exact fit: the coefficient is infinitely many standard + // errors from zero, and the evidence is total. Reporting + // t = 0 next to p = 0 would contradict itself. A zero + // coefficient beside the zero standard error has nothing + // to test and reports p = 1. + out.StandardErrors[j] = 0 + if beta[j] != 0 { + out.TStatistics[j] = math.Copysign(math.Inf(1), beta[j]) + out.PValues[j] = 0 + } else { + out.PValues[j] = 1 + } + default: + // A near-collinear design drives the solve's diagonal + // negative through rounding alone: the Wald variance would + // be a NaN beside a nil error, the same refusal GLM makes. + return nil, base.Errf("%s: the design is near-collinear: the variance of coefficient %d came out negative (%g)", name, j, v) + } + } + + // F-test of the model: H₀: every coefficient is zero. With a + // constant column this is the usual regression F against the mean; + // without one it is the test against the zero model (DModel = p). + if out.DModel > 0 { + explained := tss - rss + if explained < 0 { + explained = 0 // rounding only, and a negative F is meaningless + } + out.FStatistic = explained / float64(out.DModel) / out.ResidualVariance + if math.IsNaN(out.FStatistic) { + // 0/0: a response with no variation at all, reproduced + // exactly by the fit. There is no evidence of a model, so + // the statistic is the zero the test reads as p = 1, not a + // NaN that every consumer would propagate. + out.FStatistic = 0 + } + // Tail of F(d1, d2) at f: the regularised incomplete beta + // I_{d2/(d2+d1·f)}(d2/2, d1/2). The argument is clamped to + // [0,1]: the identity is only defined there, and an F of 0 + // (explained 0) or an infinite one would step outside through + // rounding alone. + d1, d2 := float64(out.DModel), float64(out.DResidual) + xi := d2 / (d2 + d1*out.FStatistic) + xi = min(max(xi, 0), 1) + tail, err := BetaIncomplete(xi, d2/2, d1/2) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + out.FPValue = tail + } else { + // An intercept-only design has no model term to test: the F + // stays at its zero value and the p value is 1, the same + // convention the F = 0 guard below uses. Leaving the zero + // value in FPValue would report the null model as maximally + // significant. + out.FPValue = 1 + } + return out, nil +} + +// twoSidedT returns P(|T| > |t|) for Student-t with df degrees of +// freedom, by the closed-form tail I_z(df/2, 1/2) with z = df/(df+t²). +// The identity is used rather than 2·(1 − T_cdf(t)): near t = 0 the +// subtraction cancels catastrophically, while the incomplete beta +// stays accurate into the far tail where p values matter most. +func twoSidedT(t float64, df int) (float64, error) { + z := float64(df) / (float64(df) + t*t) + p, err := BetaIncomplete(z, float64(df)/2, 0.5) + if err != nil { + return 0, err + } + return p, nil +} + +// hasConstantColumn reports whether an (n, p) design holds a column of +// one repeated value, the caller-supplied intercept. The detection +// must read the design as handed in: a column that is constant there +// can stop being constant under a further transformation (Weighted- +// LinearRegression's sqrt-weighted design is the case in point), and +// the caller's null model follows the design it actually supplied. +func hasConstantColumn(x *core.Array, n, p int) bool { + fx := rawFloats(x) + for j := range p { + first := x.FloatAt(j) + constant := true + for r := 1; r < n; r++ { + var v float64 + if fx != nil { + v = fx[r*p+j] + } else { + v = x.FloatAt(r*p + j) + } + if v != first { + constant = false + break + } + } + if constant { + return true + } + } + return false +} + +// WeightedLinearRegression fits y = X·β by weighted least squares, +// observation i carrying the positive weight w[i]: the normal +// equations run on the sqrt-weighted system, so every statistic is +// the classical weighted-theory one (σ̂² on Σw·r² with n − p degrees +// of freedom, SEs from σ̂²(XᵀWX)⁻¹), while Fitted and Residuals are +// reported in the original, unweighted units. The weights must be +// finite and positive; everything else validates as LinearRegression +// does. +// +// The model statistics are the weighted-theory ones as well: R², the +// adjusted R² and the model F test are the centred quantities against +// the weighted mean Σw·y/Σw whenever the design as supplied carries a +// constant column. That column is detected on the unweighted design, +// where it is still constant: sqrt(w) makes even the intercept column +// non-constant in the system that is actually solved. +func WeightedLinearRegression(x, y, w *core.Array) (*LinearRegressionResult, error) { + const name = "WeightedLinearRegression" + if x.NDim() != 2 { + return nil, base.Errf("%s: the design must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 || w.NDim() != 1 { + return nil, base.Errf("%s: the response and the weights must be rank 1", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex || w.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + n, p := x.Shape()[0], x.Shape()[1] + if y.Len() != n || w.Len() != n { + return nil, base.Errf("%s: the design has %d rows, the response %d and the weights %d", + name, n, y.Len(), w.Len()) + } + xw := core.New(core.Float, n, p) + yw := core.New(core.Float, n) + xwVals := xw.RawFloats() + ywVals := yw.RawFloats() + // The payload walks below read the dense slices directly where they + // exist: the elements are the ones FloatAt returns, so every product + // and every sum keeps its exact operand bits. + fw := rawFloats(w) + fx := rawFloats(x) + fy := rawFloats(y) + for r := range n { + var weight float64 + if fw != nil { + weight = fw[r] + } else { + weight = w.FloatAt(r) + } + if math.IsNaN(weight) || math.IsInf(weight, 0) || weight <= 0 { + return nil, base.Errf("%s: weight %d is %g, want a finite positive value", name, r, weight) + } + sqrtW := math.Sqrt(weight) + for j := range p { + var xj float64 + if fx != nil { + xj = fx[r*p+j] + } else { + xj = x.FloatAt(r*p + j) + } + xwVals[r*p+j] = xj * sqrtW + } + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + ywVals[r] = yv * sqrtW + } + out, err := LinearRegression(xw, yw) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + // hasConstant is read off the unweighted design, because that is the + // model the caller described; the sqrt-weighted system cannot answer + // the question, its intercept column is sqrt(w). + hasConstant := hasConstantColumn(x, n, p) + // Fitted and Residuals back in the original units, against the + // same coefficients. + for r := range n { + f := 0.0 + if fx != nil { + row := fx[r*p : r*p+p] + for i, xi := range row { + f += out.Coefficients[i] * xi + } + } else { + for i := range p { + f += out.Coefficients[i] * x.FloatAt(r*p+i) + } + } + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + out.Fitted[r] = f + out.Residuals[r] = yv - f + } + // The weighted model statistics, from the residuals just computed: + // RSS_w = Σw·r², the weighted mean ȳ_w = Σw·y/Σw, and the centred + // total Σw·(y − ȳ_w)². The delegated fit had to answer the same + // questions for the sqrt-weighted system, which is a different + // regression and reports the uncentred conventions whenever the + // weights vary, so the four model-level fields are overwritten here. + sumW, sumWY, rssW := 0.0, 0.0, 0.0 + for r := range n { + var wr float64 + if fw != nil { + wr = fw[r] + } else { + wr = w.FloatAt(r) + } + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + sumW += wr + sumWY += wr * yv + rssW += wr * out.Residuals[r] * out.Residuals[r] + } + tssW := 0.0 + if hasConstant { + meanW := sumWY / sumW + for r := range n { + var wr, yv float64 + if fw != nil { + wr = fw[r] + } else { + wr = w.FloatAt(r) + } + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + d := yv - meanW + tssW += wr * d * d + } + } else { + // Without an intercept the null model is zero, so Σw·y² is the + // total the model has to beat and it carries n degrees of freedom. + for r := range n { + var wr, yv float64 + if fw != nil { + wr = fw[r] + } else { + wr = w.FloatAt(r) + } + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + tssW += wr * yv * yv + } + } + tssDOF := n - 1 + if !hasConstant { + tssDOF = n + } + if tssW == 0 { + // Constant weighted response, exact fit: 1, as above. + out.RSquared = 1 + out.AdjustedRSquared = 1 + } else { + out.RSquared = 1 - rssW/tssW + out.AdjustedRSquared = 1 - (rssW/float64(out.DResidual))/(tssW/float64(tssDOF)) + } + out.DModel = p - 1 + if !hasConstant { + out.DModel = p + } + if out.DModel > 0 { + explained := tssW - rssW + if explained < 0 { + explained = 0 // rounding only, and a negative F is meaningless + } + out.FStatistic = explained / float64(out.DModel) / out.ResidualVariance + if math.IsNaN(out.FStatistic) { + // 0/0, as in the unweighted path: nothing to test, p = 1. + out.FStatistic = 0 + } + d1, d2 := float64(out.DModel), float64(out.DResidual) + xi := d2 / (d2 + d1*out.FStatistic) + xi = min(max(xi, 0), 1) + tail, err := BetaIncomplete(xi, d2/2, d1/2) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + out.FPValue = tail + } else { + // Intercept-only, as in the unweighted path: no model term to + // test, F 0 and p 1. + out.FStatistic = 0 + out.FPValue = 1 + } + return out, nil +} diff --git a/stats/regression_edge_pin_test.go b/stats/regression_edge_pin_test.go new file mode 100644 index 0000000..187d42e --- /dev/null +++ b/stats/regression_edge_pin_test.go @@ -0,0 +1,241 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Regression-edge pins for the linear model: the statistics of a +// design without an intercept column, the tail of the coefficient +// p-values, and the exact-fit report. + +// fact returns n!, small n only. +func fact(n int) float64 { + r := 1.0 + for i := 2; i <= n; i++ { + r *= float64(i) + } + return r +} + +// binom returns the binomial coefficient, small n only. +func binom(n, k int) float64 { return fact(n) / (fact(k) * fact(n-k)) } + +// tTailExactEven returns P(|T| > t) for even df in closed form. There, +// the incomplete beta has a = df/2, an integer, and b = 1/2, so the +// integral is elementary: nothing is approximated and nothing from the +// library is used. +func tTailExactEven(tv float64, df int) float64 { + m := df / 2 + z := float64(df) / (float64(df) + tv*tv) + // ∫_0^z u^(m−1)(1−u)^(−1/2) du, with u = 1 − s², is + // 2∫_{√(1−z)}^{1} (1−s²)^(m−1) ds. + s0 := math.Sqrt(1 - z) + antiderivative := func(s float64) float64 { + sum := 0.0 + for k := 0; k <= m-1; k++ { + sum += binom(m-1, k) * math.Pow(-1, float64(k)) * math.Pow(s, float64(2*k+1)) / float64(2*k+1) + } + return sum + } + num := 2 * (antiderivative(1) - antiderivative(s0)) + lg, _ := math.Lgamma(float64(m)) + lgb, _ := math.Lgamma(0.5) + lgs, _ := math.Lgamma(float64(m) + 0.5) + beta := math.Exp(lg + lgb - lgs) // B(m, 1/2) + return num / beta +} + +// TestTwoSidedTAccuracy pins the coefficient tail against exact closed +// forms and against the far tail, where the previous 2·(1 − T_cdf) +// form lost every digit and returned an exact zero. +func TestTwoSidedTAccuracy(t *testing.T) { + // Exact references: Cauchy (df = 1) and the df = 2 closed form. + for _, tc := range []struct { + t float64 + want float64 + }{{0.5, 2 * math.Atan(1/0.5) / math.Pi}, {2, 2 * math.Atan(1.0/2) / math.Pi}} { + got, err := twoSidedT(tc.t, 1) + if err != nil { + t.Fatal(err) + } + if math.Abs(got-tc.want) > 1e-13*tc.want { + t.Fatalf("twoSidedT(%v, 1) = %.17g, want %.17g", tc.t, got, tc.want) + } + } + for _, tc := range []struct { + t float64 + want float64 + }{{0.5, 1 - 0.5/math.Sqrt(2+0.25)}, {2, 1 - 2/math.Sqrt(6)}} { + got, err := twoSidedT(tc.t, 2) + if err != nil { + t.Fatal(err) + } + if math.Abs(got-tc.want) > 1e-13*tc.want { + t.Fatalf("twoSidedT(%v, 2) = %.17g, want %.17g", tc.t, got, tc.want) + } + } + // Even degrees of freedom: the elementary closed form. + for _, tc := range []struct { + t float64 + df int + }{{2, 6}, {8, 6}, {0.5, 6}, {2, 10}, {8, 10}, {3, 20}} { + got, err := twoSidedT(tc.t, tc.df) + if err != nil { + t.Fatalf("twoSidedT(%v, %d): %v", tc.t, tc.df, err) + } + want := tTailExactEven(tc.t, tc.df) + if math.Abs(got-want) > 1e-9*want { + t.Fatalf("twoSidedT(%v, %d) = %.17g, closed form says %.17g", tc.t, tc.df, got, want) + } + } + // Large df, against a reference computed outside the library to + // 60 digits: the tail of t(8 | df = 1e6) is + // 1.2455063433202503e-15. The cancelling form returned 1.33227e-15 + // here, 7 % high, so this pins the accuracy rather than the order. + const independentTail = 1.2455063433202503e-15 + got, err := twoSidedT(8, 1000000) + if err != nil { + t.Fatal(err) + } + if math.Abs(got-independentTail) > 1e-6*independentTail { + t.Fatalf("twoSidedT(8, 1e6) = %.17g, the independent reference is %.17g", got, independentTail) + } + // The far tail must stay positive: the cancelling form returned 0 + // for t = 30, df = 100, where the true tail is 8.4e-52. + far, err := twoSidedT(30, 100) + if err != nil { + t.Fatal(err) + } + if far <= 0 { + t.Fatalf("twoSidedT(30, 100) = %v, want a positive tail", far) + } + if far > 1e-40 { + t.Fatalf("twoSidedT(30, 100) = %v, want a tail near 8.4e-52", far) + } + // Exact value at the centre. + if p, err := twoSidedT(0, 7); err != nil || p != 1 { + t.Fatalf("twoSidedT(0, 7) = %v (err %v), want exactly 1", p, err) + } +} + +// TestLinearRegressionWithoutIntercept pins the statistics of a design +// with no constant column: the uncentred total sum of squares is the +// null model, and the model degrees of freedom are the column count. +func TestLinearRegressionWithoutIntercept(t *testing.T) { + t.Run("single column", func(t *testing.T) { + x := mustMatrix(t, []float64{1, 0, -1}, 3, 1) + y := mustFloats(t, []float64{3, 1, 2}, 3) + res, err := LinearRegression(x, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + // beta = Σxy/Σx² = 1/2, rss = 13.5, Σy² = 14. + if math.Abs(res.Coefficients[0]-0.5) > 1e-15 { + t.Fatalf("slope = %v, want 0.5", res.Coefficients[0]) + } + wantR2 := 1 - 13.5/14 + if math.Abs(res.RSquared-wantR2) > 1e-12 { + t.Fatalf("R² = %v, want %v (uncentred)", res.RSquared, wantR2) + } + if res.DModel != 1 { + t.Fatalf("DModel = %d, want 1", res.DModel) + } + if !(res.FStatistic > 0) || res.FPValue <= 0 || res.FPValue > 1 { + t.Fatalf("F = %v, p = %v, want a positive statistic and a probability", res.FStatistic, res.FPValue) + } + }) + t.Run("two columns", func(t *testing.T) { + x := mustMatrix(t, []float64{1, 0, 0, 1, -1, 1}, 3, 2) + y := mustFloats(t, []float64{3, 1, 2}, 3) + res, err := LinearRegression(x, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + // beta = [5/3, 7/3], rss = 16/3, Σy² = 14. + wantR2 := 1 - (16.0/3)/14 + if math.Abs(res.RSquared-wantR2) > 1e-12 { + t.Fatalf("R² = %v, want %v", res.RSquared, wantR2) + } + if res.DModel != 2 { + t.Fatalf("DModel = %d, want 2", res.DModel) + } + }) + t.Run("intercept unchanged", func(t *testing.T) { + // The centred form still applies when a constant column is + // present: an exact line fits perfectly. + x := mustMatrix(t, []float64{1, 1, 1, 2, 1, 3}, 3, 2) + y := mustFloats(t, []float64{3, 5, 7}, 3) + res, err := LinearRegression(x, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + if math.Abs(res.RSquared-1) > 1e-12 { + t.Fatalf("R² = %v, want 1", res.RSquared) + } + if res.DModel != 1 { + t.Fatalf("DModel = %d, want p−1 = 1", res.DModel) + } + }) +} + +// TestLinearRegressionExactFitReport pins the exact-fit report: zero +// standard errors mean an infinite statistic, not a zero one, and the +// p-value is zero rather than absent. +func TestLinearRegressionExactFitReport(t *testing.T) { + x := mustMatrix(t, []float64{1, 1, 1, 2, 1, 3, 1, 4}, 4, 2) + y := mustFloats(t, []float64{3, 5, 7, 9}, 4) + res, err := LinearRegression(x, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + if res.RSquared != 1 { + t.Fatalf("R² = %v, want exactly 1", res.RSquared) + } + for j := range 2 { + if res.StandardErrors[j] != 0 { + t.Fatalf("se[%d] = %v, want 0", j, res.StandardErrors[j]) + } + if !math.IsInf(res.TStatistics[j], 0) { + t.Fatalf("t[%d] = %v, want ±Inf", j, res.TStatistics[j]) + } + if res.PValues[j] != 0 { + t.Fatalf("p[%d] = %v, want 0", j, res.PValues[j]) + } + } +} + +// TestLinearRegressionRefusesNonFinite pins the input contract: the +// other tests in the package refuse non-finite data and this one must +// not answer with a silent column of NaN. +func TestLinearRegressionRefusesNonFinite(t *testing.T) { + x := mustMatrix(t, []float64{1, 1, 1, 2, 1, 3}, 3, 2) + for _, bad := range []float64{math.NaN(), math.Inf(1), math.Inf(-1)} { + y := mustFloats(t, []float64{3, bad, 7}, 3) + if _, err := LinearRegression(x, y); err == nil { + t.Fatalf("expected an error for the response holding %v", bad) + } else if !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("error = %v, want a non-finite refusal", err) + } + } + xb := mustMatrix(t, []float64{1, 1, 1, 2, 1, math.NaN()}, 3, 2) + if _, err := LinearRegression(xb, mustFloats(t, []float64{3, 5, 7}, 3)); err == nil { + t.Fatal("expected an error for a design holding NaN") + } +} + +// mustMatrix builds an (r, c) float64 array. +func mustMatrix(t *testing.T, vals []float64, r, c int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, r, c) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} diff --git a/stats/regression_test.go b/stats/regression_test.go new file mode 100644 index 0000000..0898235 --- /dev/null +++ b/stats/regression_test.go @@ -0,0 +1,187 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// mustDesign builds a design matrix from row-major values. +func mustDesign(t *testing.T, vals []float64, rows, cols int) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, rows, cols) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestLinearRegressionExactModel pins the inference on a noiseless +// linear model: the coefficients recover, R² is 1, σ̂² is 0 and the +// perfect fit is flagged as such (t infinite, p zero). +func TestLinearRegressionExactModel(t *testing.T) { + // y = 2 + 3x over x = 1..6. + design := mustDesign(t, []float64{ + 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, 1, 6, + }, 6, 2) + yv := make([]float64, 6) + for i := range 6 { + yv[i] = 2 + 3*float64(i+1) + } + y, err := core.FromFloats(yv, 6) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + res, err := LinearRegression(design, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + if math.Abs(res.Coefficients[0]-2) > 1e-12 || math.Abs(res.Coefficients[1]-3) > 1e-12 { + t.Fatalf("coefficients = (%.10f, %.10f), want (2, 3)", res.Coefficients[0], res.Coefficients[1]) + } + if math.Abs(res.RSquared-1) > 1e-12 { + t.Fatalf("R² = %.12f, want 1", res.RSquared) + } + if res.ResidualVariance > 1e-24 { + t.Fatalf("σ̂² = %.3g, want 0", res.ResidualVariance) + } + if res.FPValue > 1e-12 { + t.Fatalf("F p-value = %.3g, want 0", res.FPValue) + } +} + +// TestLinearRegressionNoisy pins the inference against hand-computed +// formulas on a small noisy dataset. +func TestLinearRegressionNoisy(t *testing.T) { + design := mustDesign(t, []float64{ + 1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, + }, 6, 2) + yv := []float64{1.2, 2.1, 2.9, 4.2, 4.8, 6.3} + y, _ := core.FromFloats(yv, 6) + res, err := LinearRegression(design, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + n, p := 6.0, 2.0 + // Reference computations from first principles. + xv := []float64{0, 1, 2, 3, 4, 5} + sx := 0.0 + sy := 0.0 + sxy := 0.0 + sxx := 0.0 + for i := range 6 { + sx += xv[i] + sy += yv[i] + sxy += xv[i] * yv[i] + sxx += xv[i] * xv[i] + } + slope := (n*sxy - sx*sy) / (n*sxx - sx*sx) + intercept := (sy - slope*sx) / n + if math.Abs(res.Coefficients[1]-slope) > 1e-12 || math.Abs(res.Coefficients[0]-intercept) > 1e-12 { + t.Fatalf("coefficients = (%.10f, %.10f), want (%.10f, %.10f)", + res.Coefficients[0], res.Coefficients[1], intercept, slope) + } + rss := 0.0 + for i := range 6 { + r := yv[i] - (intercept + slope*xv[i]) + rss += r * r + if math.Abs(res.Residuals[i]-r) > 1e-12 { + t.Fatalf("residual[%d] = %.10f, want %.10f", i, res.Residuals[i], r) + } + } + if math.Abs(res.ResidualVariance-rss/(n-p)) > 1e-12 { + t.Fatalf("σ̂² = %.10f, want %.10f", res.ResidualVariance, rss/(n-p)) + } + // SE of the slope: σ̂·sqrt(1/Sxx_c) with Sxx_c = Σ(x−x̄)². + xbar := sx / n + sxxc := 0.0 + for _, v := range xv { + sxxc += (v - xbar) * (v - xbar) + } + seSlope := math.Sqrt(rss/(n-p)) / math.Sqrt(sxxc) + // NaN would slip through a plain tolerance comparison, so it gets + // its own gate before the value check. + if math.IsNaN(res.StandardErrors[1]) || math.IsNaN(res.TStatistics[1]) || math.IsNaN(res.PValues[1]) { + t.Fatalf("non-finite inference: SE=%.6g t=%.6g p=%.6g", + res.StandardErrors[1], res.TStatistics[1], res.PValues[1]) + } + if math.Abs(res.StandardErrors[1]-seSlope) > 1e-12 { + t.Fatalf("SE(slope) = %.10f, want %.10f", res.StandardErrors[1], seSlope) + } + if math.Abs(res.TStatistics[1]-slope/seSlope) > 1e-10 { + t.Fatalf("t(slope) = %.10f, want %.10f", res.TStatistics[1], slope/seSlope) + } + // Two-sided p from the t CDF itself (independent re-entry). + pv, err := twoSidedT(res.TStatistics[1], 4) + if err != nil { + t.Fatalf("twoSidedT: %v", err) + } + if math.Abs(res.PValues[1]-pv) > 1e-14 { + t.Fatalf("p(slope) = %.10f, want %.10f", res.PValues[1], pv) + } + // F = t² for the single-slope model. + if math.Abs(res.FStatistic-res.TStatistics[1]*res.TStatistics[1]) > 1e-8 { + t.Fatalf("F = %.8f, want t² = %.8f", res.FStatistic, res.TStatistics[1]*res.TStatistics[1]) + } + // F p equals the slope's t p for p = 2. + if math.Abs(res.FPValue-res.PValues[1]) > 1e-10 { + t.Fatalf("F p = %.10f, want %.10f", res.FPValue, res.PValues[1]) + } +} + +// TestLinearRegressionTwoPredictors pins a multi-column design against +// a brute-force normal-equation solve in the test. +func TestLinearRegressionTwoPredictors(t *testing.T) { + g := core.NewGenerator(17) + const n = 40 + vals := make([]float64, n*3) + yv := make([]float64, n) + for i := range n { + a := g.NormalUnit() + b := g.NormalUnit() + vals[i*3] = 1 + vals[i*3+1] = a + vals[i*3+2] = b + yv[i] = -1.5 + 2*a - 0.7*b + 0.3*g.NormalUnit() + } + design := mustDesign(t, vals, n, 3) + y, _ := core.FromFloats(yv, n) + res, err := LinearRegression(design, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + if math.Abs(res.Coefficients[0]+1.5) > 0.2 || math.Abs(res.Coefficients[1]-2) > 0.2 || + math.Abs(res.Coefficients[2]+0.7) > 0.2 { + t.Fatalf("coefficients = %v, want roughly (-1.5, 2, -0.7)", res.Coefficients) + } + if res.RSquared < 0.9 { + t.Fatalf("R² = %.3f, the model explains more than that", res.RSquared) + } + if res.FPValue > 1e-10 { + t.Fatalf("F p = %.3g, the model is significant", res.FPValue) + } + if res.DModel != 2 || res.DResidual != 37 { + t.Fatalf("degrees of freedom (%d, %d), want (2, 37)", res.DModel, res.DResidual) + } +} + +// TestLinearRegressionErrors pins the input gates. +func TestLinearRegressionErrors(t *testing.T) { + x := mustDesign(t, []float64{1, 1, 1, 1}, 2, 2) + y, _ := core.FromFloats([]float64{1, 2}, 2) + if _, err := LinearRegression(x, y); err == nil { + t.Error("n == p accepted") + } + if _, err := LinearRegression(x, y); err == nil { + t.Error("square design accepted") + } + bad := mustDesign(t, []float64{1, 2, 2, 4, 1, 2}, 3, 2) + y3, _ := core.FromFloats([]float64{1, 2, 3}, 3) + if _, err := LinearRegression(bad, y3); err == nil { + t.Error("rank-deficient design accepted") + } +} diff --git a/stats/robust.go b/stats/robust.go new file mode 100644 index 0000000..8ca8446 --- /dev/null +++ b/stats/robust.go @@ -0,0 +1,102 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "slices" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Robust location and scale summaries: the pair that survives a few +// wild samples the mean and standard deviation would happily chase. + +// MedianAbsoluteDeviation returns the median of |x − median(x)|, the +// robust scale estimate that breaks down only when nearly half the +// sample is wild. Multiply by 1.4826 to read it as a standard +// deviation on Gaussian data. +func MedianAbsoluteDeviation(a *core.Array) (float64, error) { + median, err := Median(a) + if err != nil { + return 0, err + } + deviations := core.New(core.Float, a.Len()) + vals := deviations.RawFloats() + // The walk is bounded by the sample's own element count: a rebased + // view's payload may run past its visible elements, and writing one + // per payload slot would overrun the deviations buffer. + if fs := rawFloats(a); fs != nil { + for i, v := range fs[:a.Len()] { + vals[i] = math.Abs(v - median) + } + } else { + for i := range a.Len() { + vals[i] = math.Abs(a.FloatAt(i) - median) + } + } + return Median(deviations) +} + +// TrimmedMean averages the sample after dropping the fraction from +// each tail: fraction 0.1 discards the smallest and the largest ten +// percent (floored to whole samples) before averaging, which keeps +// the mean honest against one-sided contamination. fraction must lie +// in [0, 0.5) and leave at least one sample in the middle; a non-finite +// sample is an error. +func TrimmedMean(a *core.Array, fraction float64) (float64, error) { + if a.Dtype() == core.Complex { + return 0, base.Errf("TrimmedMean: complex samples have no ordering") + } + n := a.Len() + if n == 0 { + return 0, base.Errf("TrimmedMean: an empty sample has no mean") + } + if !(fraction >= 0) || fraction >= 0.5 { + return 0, base.Errf("TrimmedMean: the fraction must lie in [0, 0.5), got %g", fraction) + } + if err := checkFinite("TrimmedMean", "the sample", a); err != nil { + return 0, err + } + trim := int(fraction * float64(n)) + if n-2*trim < 1 { + return 0, base.Errf("TrimmedMean: trimming %d samples a side of %d leaves nothing", trim, n) + } + sortedVals := make([]float64, n) + if fs := rawFloats(a); fs != nil { + copy(sortedVals, fs) + } else { + for i := range n { + sortedVals[i] = a.FloatAt(i) + } + } + slices.Sort(sortedVals) + window := sortedVals[trim : n-trim] + // Two-pass scaled summation: dividing through by the largest + // magnitude first keeps every partial sum inside [−k, k] for a + // window of k samples, so a mean of huge values cannot overflow + // into an infinity the way the direct accumulation does. Kahan + // compensation only repairs rounding, never overflow, so the + // scaling is what carries the guarantee here. + mag := 0.0 + for _, v := range window { + if m := math.Abs(v); m > mag { + mag = m + } + } + if mag == 0 { + return 0, nil + } + total := 0.0 + for _, v := range window { + total += v / mag + } + // The division comes before the multiplication: the scaled total is + // at most the window's size in magnitude, so dividing by the count + // first keeps the final scaling inside the representable range, + // where multiplying the raw total back by the magnitude could + // overflow before the division runs. + return total / float64(n-2*trim) * mag, nil +} diff --git a/stats/robustregression.go b/stats/robustregression.go new file mode 100644 index 0000000..ff9ed9a --- /dev/null +++ b/stats/robustregression.go @@ -0,0 +1,652 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// Robust regression: the fits that survive the wild observations the +// classical least squares would chase. Two estimators live here. The +// Huber M-estimator bounds the influence of a gross outlier by +// replacing the squared loss with one that grows quadratically only +// inside a band around zero and linearly outside it, fitted by +// iteratively reweighted least squares. Theil-Sen replaces the whole +// least-squares machinery with the median of the pairwise slopes, +// which a minority of broken points cannot move. + +// DefaultHuberTuning is the Huber tuning constant the plain +// HuberRegression entry uses. The value 1.345 is the literature's +// standard choice: with the band measured in robust standard +// deviations it puts the estimator at 95 percent asymptotic efficiency +// at the Gaussian while keeping the influence of an outlier bounded at +// 1.345 times what a residual inside the band would have. +const DefaultHuberTuning = 1.345 + +// huberMaxIterations bounds the reweighting loop, and the tolerance is +// the largest coefficient movement one iteration may leave behind for +// the fit to call itself settled. Both follow the glm.go convention. +const ( + huberMaxIterations = 100 + huberTolerance = 1e-10 +) + +// HuberRegressionResult carries a Huber M-estimate of the linear +// model. +type HuberRegressionResult struct { + // Coefficients are the M-estimates β̂, one per design column, in + // the design's own order. + Coefficients []float64 + // StandardErrors are the classical asymptotic standard errors read + // from σ²·(XᵀWX)⁻¹, W the final weight vector: the same shape of + // statement WeightedLinearRegression makes, with the robust scale + // σ in place of the residual standard deviation. + StandardErrors []float64 + // Weights are the final IRLS weights, one per observation: exactly + // 1 inside the band |r| ≤ k·σ and tapering as k·σ/|r| outside it. + // They are the diagnostic a robust fit exists to produce: the + // contaminated observations are the ones at the bottom of the + // list. + Weights []float64 + // Scale is the final robust scale σ, 1.4826 times the median + // absolute deviation of the residuals, the estimate the band is + // measured in. + Scale float64 + // Fitted and Residuals align with the rows of the design. + Fitted []float64 + Residuals []float64 + // Iterations counts the reweighting steps taken; Converged reports + // whether the coefficient updates fell under the tolerance. + Iterations int + Converged bool +} + +// HuberRegression fits y = X·β with the Huber M-estimator at the +// default tuning constant. See HuberRegressionTuned for the full +// contract. +func HuberRegression(x, y *core.Array) (*HuberRegressionResult, error) { + return HuberRegressionTuned(x, y, DefaultHuberTuning) +} + +// HuberRegressionTuned fits y = X·β by Huber's M-estimation with the +// tuning constant tuning: the loss is r²/2 inside the band |r| ≤ +// tuning·σ and tuning·σ·(|r| − tuning·σ/2) outside it, so a wild +// observation pulls like a linear, not a quadratic, residual. The fit +// runs by iteratively reweighted least squares: ordinary least squares +// to start, then each round re-estimates the robust scale σ as 1.4826 +// times the median absolute deviation of the current residuals, +// weights each observation by 1 inside the band and tuning·σ/|r| +// outside it, and solves the weighted normal equations until no +// coefficient moves by more than 1e-10. +// +// The design carries n rows and p columns exactly as +// LinearRegression's, the intercept included by the caller as a +// constant column when wanted, and the same validations apply: n > p, +// a full-rank design, real finite input, and a positive finite tuning +// constant. A robust scale that collapses to zero, more than half the +// residuals landing on one value, stops the iteration as an exact fit +// and is reported rather than divided by. +func HuberRegressionTuned(x, y *core.Array, tuning float64) (*HuberRegressionResult, error) { + const name = "HuberRegression" + if x.NDim() != 2 { + return nil, base.Errf("%s: the design must be rank 2, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 { + return nil, base.Errf("%s: the response must be rank 1", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return nil, base.Errf("%s: complex inputs are not supported", name) + } + n, p := x.Shape()[0], x.Shape()[1] + if y.Len() != n { + return nil, base.Errf("%s: the design has %d rows but the response %d", name, n, y.Len()) + } + if n <= p { + return nil, base.Errf("%s: need n > p, got %d observations and %d columns", name, n, p) + } + if p == 0 { + return nil, base.Errf("%s: the design must carry at least one column", name) + } + if err := checkFinite(name, "the design", x); err != nil { + return nil, err + } + if err := checkFinite(name, "the response", y); err != nil { + return nil, err + } + if math.IsNaN(tuning) || math.IsInf(tuning, 0) || tuning <= 0 { + return nil, base.Errf("%s: the tuning constant must be finite and positive, got %g", name, tuning) + } + fx := rawFloats(x) + fy := rawFloats(y) + + // The start is the ordinary least squares answer: the M-estimator's + // own optimum is rarely far from it, and the reweighting does the + // rest. The normal equations go through the shared LU, exactly as + // LinearRegression solves them. + beta, err := leastSquaresSolve(name, n, p, x, y, fx, fy, nil) + if err != nil { + return nil, err + } + residuals := make([]float64, n) + regressionResiduals(n, p, x, y, fx, fy, beta, residuals) + scaleScratch := make([]float64, n) + // One workspace serves every reweighting round: the solve consumes + // its buffers in place and the next round clears them first. + ws := newLeastSquaresWorkspace(p) + out := &HuberRegressionResult{ + Coefficients: beta, + Weights: make([]float64, n), + } + converged := false + iterations := huberMaxIterations + for iter := 1; iter <= huberMaxIterations; iter++ { + scale := huberScale(residuals, scaleScratch) + if scale == 0 { + // The robust scale has collapsed: more than half the + // residuals sit on one value, the band has nothing to + // widen, and the fit is as settled as it will ever be. The + // weights follow the collapsed band, one inside it and + // zero outside, and the loop stops rather than divide by + // the zero the tapering would need. + for i, r := range residuals { + out.Weights[i] = huberWeight(r, 0) + } + out.Scale = 0 + converged = true + iterations = iter - 1 + break + } + band := tuning * scale + for i, r := range residuals { + out.Weights[i] = huberWeight(r, band) + } + updated, err := leastSquaresSolveWS(name, n, p, x, y, fx, fy, out.Weights, ws) + if err != nil { + return nil, err + } + worst := 0.0 + for j := range p { + if d := math.Abs(updated[j] - beta[j]); d > worst { + worst = d + } + beta[j] = updated[j] + } + regressionResiduals(n, p, x, y, fx, fy, beta, residuals) + if worst < huberTolerance { + out.Scale = huberScale(residuals, scaleScratch) + // The reported weights follow the answer, not the step + // that produced it: the collapsed-scale branch refreshes + // them, and the converged exit must agree with the Scale + // and Residuals it reports (huberWeight at band 0 is the + // collapsed rule, one inside and zero outside). + band := tuning * out.Scale + for i, r := range residuals { + out.Weights[i] = huberWeight(r, band) + } + converged = true + iterations = iter + break + } + } + if !converged { + return nil, base.Errf("%s: %d iterations did not converge", name, huberMaxIterations) + } + out.Iterations = iterations + out.Converged = converged + + // Fitted and Residuals from the settled coefficients, and the + // standard errors from σ²·(XᵀWX)⁻¹, its diagonal read from one + // factorisation of the pristine normal equations against all p unit + // columns at once, the way LinearRegression reads its own + // covariance diagonal. + out.Fitted = make([]float64, n) + out.Residuals = make([]float64, n) + copy(out.Residuals, residuals) + weighted := out.Scale * out.Scale + wxxPristine := weightedNormal(n, p, x, fx, out.Weights) + out.StandardErrors = make([]float64, p) + unit := make([][]float64, p) + for j := range p { + unit[j] = make([]float64, p) + unit[j][j] = 1 + } + inv, err := base.SolveSystem(name, wxxPristine, unit) + if err != nil { + return nil, base.Errf("%s: %w", name, err) + } + for j := range p { + v := weighted * inv[j][j] + switch { + case v > 0: + out.StandardErrors[j] = math.Sqrt(v) + case v == 0: + // An exact fit: nothing to estimate, a zero standard error + // beside the zero residual variance. + out.StandardErrors[j] = 0 + default: + return nil, base.Errf("%s: the design is near-collinear: the variance of coefficient %d came out negative (%g)", name, j, v) + } + } + return out, nil +} + +// huberWeight is the IRLS weight of one residual for a band: exactly +// one inside the band, the band over the magnitude outside it, which +// is the taper that turns the quadratic loss linear. A collapsed band +// leaves the zero residual at weight one and everything else at zero. +func huberWeight(residual, band float64) float64 { + a := math.Abs(residual) + if a <= band { + return 1 + } + return band / a +} + +// huberScale estimates the robust scale of a residual sample as +// 1.4826 times the median absolute deviation, the constant that makes +// the estimate consistent for the standard deviation of Gaussian +// residuals. It is scale-free against a shift because the median is +// taken twice, once at the centre and once over the deviations. +// +// scratch is a buffer as long as residuals, reused across the +// reweighting rounds: the centre pass sorts it, and the deviation pass +// then writes every slot of it before the second median reads any, so +// the rounds cannot see each other's values. +func huberScale(residuals, scratch []float64) float64 { + vals := scratch[:len(residuals)] + copy(vals, residuals) + centre := medianSlice(vals) + for i, r := range residuals { + vals[i] = math.Abs(r - centre) + } + return 1.4826 * medianSlice(vals) +} + +// medianSlice rearranges vals in place and returns their median, +// averaging the two middle values on even length exactly as Median does: +// the halved-magnitudes sum, immune to the overflow the literal average +// risks. Callers hand over a disposable slice. +// +// The two middle order statistics come from a selection, not from a sort: +// at the Theil-Sen observation cap the slope list holds millions of +// entries and the sort cost more than everything else in the fit put +// together. The selected values are the ones the sort left at the middle +// indices, with one honest difference: a sort is unstable, so among values +// that compare equal but differ in bits (the two zeros, a NaN payload) the +// sort's choice of which one lands at the middle is unspecified, and the +// selection may return the other. Every other input answers the identical +// number. +func medianSlice(vals []float64) float64 { + n := len(vals) + if n%2 == 1 { + selectNth(vals, n/2) + return vals[n/2] + } + selectNth(vals, n/2) + // Everything before index n/2 is no larger than the upper middle, so + // the lower middle is the largest of that prefix. + lo := vals[0] + for _, v := range vals[:n/2] { + if v > lo { + lo = v + } + } + return lo/2 + vals[n/2]/2 +} + +// selectNth rearranges vals so that the element at index k is the k-th +// smallest, every element before it is no larger, and every element after +// it is no smaller, using Hoare's partition with a median-of-three pivot +// and an insertion sort for short ranges. The pivot choice and the range +// cutoff decide speed only: whichever pivot is taken, the value that ends +// up at k is the k-th order statistic. +func selectNth(vals []float64, k int) { + lo, hi := 0, len(vals)-1 + for hi-lo > 24 { + mid := lo + (hi-lo)/2 + if vals[mid] < vals[lo] { + vals[mid], vals[lo] = vals[lo], vals[mid] + } + if vals[hi] < vals[lo] { + vals[hi], vals[lo] = vals[lo], vals[hi] + } + if vals[hi] < vals[mid] { + vals[hi], vals[mid] = vals[mid], vals[hi] + } + pivot := vals[mid] + i, j := lo, hi + for i <= j { + for vals[i] < pivot { + i++ + } + for pivot < vals[j] { + j-- + } + if i <= j { + vals[i], vals[j] = vals[j], vals[i] + i++ + j-- + } + } + if k <= j { + hi = j + continue + } + if k >= i { + lo = i + continue + } + return + } + for i := lo + 1; i <= hi; i++ { + v := vals[i] + j := i - 1 + for j >= lo && v < vals[j] { + vals[j+1] = vals[j] + j-- + } + vals[j+1] = v + } +} + +// regressionResiduals recomputes r = y − Xβ into residuals. +func regressionResiduals(n, p int, x, y *core.Array, fx, fy []float64, beta, residuals []float64) { + for i := range n { + f := 0.0 + if fx != nil { + row := fx[i*p : i*p+p] + for j, xj := range row { + f += beta[j] * xj + } + } else { + for j := range p { + f += beta[j] * x.FloatAt(i*p+j) + } + } + var yv float64 + if fy != nil { + yv = fy[i] + } else { + yv = y.FloatAt(i) + } + residuals[i] = yv - f + } +} + +// weightedNormal assembles the weighted normal equations matrix +// XᵀWX from a weight vector, or the unweighted XᵀX when w is nil. +// The caller consumes the matrix through the shared solve, which +// factors in place, so every call builds fresh. +func weightedNormal(n, p int, x *core.Array, fx []float64, w []float64) [][]float64 { + m := make([][]float64, p) + for i := range p { + m[i] = make([]float64, p) + } + weightedNormalInto(m, n, p, x, fx, w) + return m +} + +// weightedNormalInto assembles XᵀWX (or XᵀX when w is nil) into the +// provided matrix, which the caller has cleared: every entry +// accumulates from zero, so a cleared buffer answers exactly what a +// fresh allocation answered. The accumulation order is the row walk's +// own, unchanged. +func weightedNormalInto(m [][]float64, n, p int, x *core.Array, fx []float64, w []float64) { + for r := range n { + wr := 1.0 + if w != nil { + wr = w[r] + } + if fx != nil { + row := fx[r*p : r*p+p] + for a, xa := range row { + ga := wr * xa + ma := m[a] + for b, xb := range row { + ma[b] += ga * xb + } + } + } else { + for a := range p { + xa := x.FloatAt(r*p + a) + ga := wr * xa + for b := range p { + m[a][b] += ga * x.FloatAt(r*p+b) + } + } + } + } +} + +// leastSquaresWorkspace carries the normal-equation buffers one +// reweighting loop refills: the shared solve factors its matrix and +// overwrites its right-hand side in place, so every round clears the +// same backing arrays and the accumulation sees the zero state a fresh +// allocation carried. +type leastSquaresWorkspace struct { + mat [][]float64 + rhs []float64 +} + +func newLeastSquaresWorkspace(p int) *leastSquaresWorkspace { + m := make([][]float64, p) + for i := range p { + m[i] = make([]float64, p) + } + return &leastSquaresWorkspace{mat: m, rhs: make([]float64, p)} +} + +// leastSquaresSolve solves the (weighted) normal equations in one +// shot: XᵀWX β = XᵀWy with W the weights, or the plain XᵀX system +// when w is nil. The right-hand side is assembled alongside the +// matrix so both see the same weights. The returned slice is the +// workspace's right-hand side, overwritten in place by the solve: the +// caller must consume it before the workspace is refilled. +func leastSquaresSolve(name string, n, p int, x, y *core.Array, fx, fy []float64, w []float64) ([]float64, error) { + return leastSquaresSolveWS(name, n, p, x, y, fx, fy, w, newLeastSquaresWorkspace(p)) +} + +// leastSquaresSolveWS is leastSquaresSolve on a caller-owned +// workspace, for a reweighting loop that solves once per round. +func leastSquaresSolveWS(name string, n, p int, x, y *core.Array, fx, fy []float64, w []float64, ws *leastSquaresWorkspace) ([]float64, error) { + for i := range p { + clear(ws.mat[i]) + } + weightedNormalInto(ws.mat, n, p, x, fx, w) + rhs := ws.rhs + clear(rhs) + for r := range n { + wr := 1.0 + if w != nil { + wr = w[r] + } + var yv float64 + if fy != nil { + yv = fy[r] + } else { + yv = y.FloatAt(r) + } + g := wr * yv + if fx != nil { + row := fx[r*p : r*p+p] + for a, xa := range row { + rhs[a] += xa * g + } + } else { + for a := range p { + rhs[a] += x.FloatAt(r*p+a) * g + } + } + } + solved, err := base.SolveSystem(name, ws.mat, [][]float64{rhs}) + if err != nil { + return nil, base.Errf("%s: the design is rank deficient (%w)", name, err) + } + return solved[0], nil +} + +// TheilSenMaxObservations is the exactness contract of +// TheilSenRegression: the median of the pairwise slopes is computed +// over all n(n−1)/2 of them, which at 4096 observations is already +// some eight million slopes and a good fraction of a gigabyte of +// working memory. Beyond the cap the estimator refuses rather than +// silently degrade to a sample of itself; the cost is named in the +// error so the caller can subsample deliberately. +const TheilSenMaxObservations = 4096 + +// theilSenParallelPairs is the pairwise-slope count from which the walk +// splits across workers: below it a crew costs more to start than the +// walk it would carry, and above it a worker is handed this many +// slopes, so the number of blocks follows the pair count rather than +// the row count. +const theilSenParallelPairs = 1 << 18 + +// theilSenSlopePoolMax bounds the pair-slope buffer the pool keeps, in +// float64 entries. The cap-sized fit needs 8,386,560 of them, inside +// the bound; a longer list allocates fresh and is dropped on return, +// so one oversized call cannot pin a larger buffer on every +// processor, and sync.Pool forgets what it holds at each garbage +// collection besides. +const theilSenSlopePoolMax = 1 << 23 + +// theilSenSlopes is the pooled pair-slope buffer. The pointer form +// keeps the Put from boxing a slice header on every return. +type theilSenSlopes struct{ s []float64 } + +var theilSenSlopePool = sync.Pool{ + New: func() any { return new(theilSenSlopes) }, +} + +// TheilSenRegression fits the simple linear model y = a + b·x by the +// Theil-Sen estimator: the slope is the median of the pairwise slopes +// (y_j − y_i)/(x_j − x_i) over all pairs with distinct predictors, +// and the intercept is the median of y_i − b·x_i at that slope. Both +// medians are exact: the slope breaks down only when nearly half the +// points are broken, and one wild observation among hundreds cannot +// drag the answer at all. At least three observations with at least +// two distinct predictors are needed; every input must be finite, and +// samples beyond TheilSenMaxObservations are refused with the cost +// named rather than answered approximately. +func TheilSenRegression(x, y *core.Array) (intercept, slope float64, err error) { + const name = "TheilSenRegression" + if x.NDim() != 1 { + return 0, 0, base.Errf("%s: the predictor must be rank 1, got shape %s", name, base.ShapeText(x.Shape())) + } + if y.NDim() != 1 { + return 0, 0, base.Errf("%s: the response must be rank 1", name) + } + if x.Dtype() == core.Complex || y.Dtype() == core.Complex { + return 0, 0, base.Errf("%s: complex inputs are not supported", name) + } + n := x.Len() + if y.Len() != n { + return 0, 0, base.Errf("%s: the predictor has %d samples but the response %d", name, n, y.Len()) + } + if n < 3 { + return 0, 0, base.Errf("%s: at least three observations are needed, got %d", name, n) + } + if n > TheilSenMaxObservations { + return 0, 0, base.Errf("%s: %d observations would need the exact median over %d pairwise slopes; the exactness contract ends at %d, subsample deliberately instead", + name, n, n*(n-1)/2, TheilSenMaxObservations) + } + if err := checkFinite(name, "the predictor", x); err != nil { + return 0, 0, err + } + if err := checkFinite(name, "the response", y); err != nil { + return 0, 0, err + } + xs := make([]float64, n) + ys := make([]float64, n) + if fx := rawFloats(x); fx != nil { + copy(xs, fx) + } else { + for i := range n { + xs[i] = x.FloatAt(i) + } + } + if fy := rawFloats(y); fy != nil { + copy(ys, fy) + } else { + for i := range n { + ys[i] = y.FloatAt(i) + } + } + // All pairwise slopes over the pairs whose predictor actually + // differs: a repeated predictor carries no slope information, and + // dividing by its zero would poison the median. The slopes land in + // one buffer in the row walk's own order, so the rows can be filled + // by a crew and the collected slice is the serial walk's own, + // element for element; the per-row pair counts are what let the + // blocks be cut before the walk starts. + counts := make([]int, n) + seen := make(map[float64]int, n) + for i := n - 1; i >= 0; i-- { + equal := seen[xs[i]] + seen[xs[i]] = equal + 1 + counts[i] = n - 1 - i - equal + } + offsets := make([]int, n+1) + for i := range n { + offsets[i+1] = offsets[i] + counts[i] + } + kept := offsets[n] + if kept == 0 { + return 0, 0, base.Errf("%s: the predictor does not vary, no slope exists", name) + } + sb := theilSenSlopePool.Get().(*theilSenSlopes) + slopes := sb.s + if cap(slopes) < kept { + slopes = make([]float64, kept) + } + slopes = slopes[:kept] + fill := func(lo, hi int) { + for i := lo; i < hi; i++ { + off := offsets[i] + xi, yi := xs[i], ys[i] + for j := i + 1; j < n; j++ { + if dx := xs[j] - xi; dx != 0 { + slopes[off] = (ys[j] - yi) / dx + off++ + } + } + } + } + // The walk's cost falls with i, so an even split of the rows would + // leave the first worker with a quarter of the work: the blocks are + // cut where the pair count crosses an equal share instead. + parts := min(kept/theilSenParallelPairs, n) + if parts < 2 { + fill(0, n) + } else { + per := kept / parts + bounds := make([]int, 1, parts+1) + for i, cut := 1, per; i < n && len(bounds) < parts; i++ { + if offsets[i] >= cut { + bounds = append(bounds, i) + cut += per + } + } + bounds = append(bounds, n) + engine.Parallel(len(bounds)-1, func(start, end int) { + for k := start; k < end; k++ { + fill(bounds[k], bounds[k+1]) + } + }) + } + slope = medianSlice(slopes) + if cap(slopes) <= theilSenSlopePoolMax { + sb.s = slopes[:cap(slopes)] + theilSenSlopePool.Put(sb) + } + intercepts := make([]float64, n) + for i := range n { + intercepts[i] = ys[i] - slope*xs[i] + } + return medianSlice(intercepts), slope, nil +} diff --git a/stats/robustregression_test.go b/stats/robustregression_test.go new file mode 100644 index 0000000..c43ba00 --- /dev/null +++ b/stats/robustregression_test.go @@ -0,0 +1,345 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "slices" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// huberFixture builds the seeded contaminated fixture the Huber pins +// run on: sixty observations of a clean linear model, then ten percent +// of the responses pushed thirty units off the trend. It returns the +// design with the constant column first, the clean response and the +// contaminated response. +func huberFixture(t *testing.T) (*core.Array, *core.Array, *core.Array) { + t.Helper() + const n = 60 + g := core.NewGenerator(11) + designVals := make([]float64, 0, 2*n) + clean := make([]float64, n) + for i := range n { + x := g.NormalUnit() + clean[i] = 2 + 3*x + 0.5*g.NormalUnit() + designVals = append(designVals, 1, x) + } + design := mustFromFloats(t, designVals, n, 2) + contaminated := append([]float64(nil), clean...) + for i := 0; i < n; i += 10 { + contaminated[i] += 100 + } + return design, mustFromFloats(t, clean, n), mustFromFloats(t, contaminated, n) +} + +// TestHuberResistsOutliers is the pin that gives the estimator its +// reason to exist: on the contaminated fixture the Huber fit lands +// beside the clean-data ordinary least squares answer, while plain +// least squares on the same contaminated data is dragged visibly off +// it. Both margins are measured, not asserted as anecdotes. +func TestHuberResistsOutliers(t *testing.T) { + design, clean, contaminated := huberFixture(t) + cleanOLS, err := LinearRegression(design, clean) + if err != nil { + t.Fatalf("LinearRegression on the clean data: %v", err) + } + if math.Abs(cleanOLS.Coefficients[0]-2) > 0.2 || math.Abs(cleanOLS.Coefficients[1]-3) > 0.2 { + t.Fatalf("the clean-data OLS is off the truth: (%.4f, %.4f)", cleanOLS.Coefficients[0], cleanOLS.Coefficients[1]) + } + olsCont, err := LinearRegression(design, contaminated) + if err != nil { + t.Fatalf("LinearRegression on the contaminated data: %v", err) + } + huber, err := HuberRegression(design, contaminated) + if err != nil { + t.Fatalf("HuberRegression: %v", err) + } + if !huber.Converged { + t.Fatalf("the Huber fit did not converge") + } + dragged := math.Abs(olsCont.Coefficients[1] - cleanOLS.Coefficients[1]) + resisted := math.Abs(huber.Coefficients[1] - cleanOLS.Coefficients[1]) + t.Logf("slope: clean OLS %.4f, contaminated OLS %.4f (drag %.4f), contaminated Huber %.4f (drag %.4f)", + cleanOLS.Coefficients[1], olsCont.Coefficients[1], dragged, huber.Coefficients[1], resisted) + if dragged < 0.3 { + t.Fatalf("the contamination failed to drag OLS in the first place: drag %.4f", dragged) + } + if resisted > 0.15 { + t.Fatalf("the Huber slope sits %.4f from the clean-data answer, want at most 0.15", resisted) + } + if d := math.Abs(huber.Coefficients[0] - cleanOLS.Coefficients[0]); d > 0.15 { + t.Fatalf("the Huber intercept sits %.4f from the clean-data answer, want at most 0.15", d) + } +} + +// TestHuberWeightsBandTaper pins the weight function itself: exactly +// one strictly inside the band and on its boundary, the band over the +// magnitude outside it, and no weight anywhere else on the number +// line. The fixture's contaminated observations must then be the ones +// the fit tapered hardest. +func TestHuberWeightsBandTaper(t *testing.T) { + // The function, on exact values. + if huberWeight(0.4, 1) != 1 { + t.Fatalf("a residual inside the band does not carry weight one") + } + if huberWeight(1, 1) != 1 { + t.Fatalf("a residual on the band boundary does not carry weight one") + } + if huberWeight(-1, 1) != 1 { + t.Fatalf("the band is not symmetric: a negative boundary residual tapered") + } + if w := huberWeight(2, 1); w != 0.5 { + t.Fatalf("a residual at twice the band carries %g, want 0.5", w) + } + if w := huberWeight(-3, 1.5); w != 0.5 { + t.Fatalf("a negative residual at twice the band carries %g, want 0.5", w) + } + if huberWeight(0, 0) != 1 { + t.Fatalf("a collapsed band must keep a zero residual at weight one") + } + if huberWeight(1, 0) != 0 { + t.Fatalf("a collapsed band must give a nonzero residual zero weight") + } + // The taper decreases strictly beyond the band. + previous := 1.0 + for magnitude := 1.5; magnitude < 20; magnitude += 0.5 { + w := huberWeight(magnitude, 1) + if w >= previous { + t.Fatalf("the taper rose at magnitude %.1f: %g after %g", magnitude, w, previous) + } + previous = w + } + // The fit's weights: the six contaminated observations are the six + // lightest. + design, _, contaminated := huberFixture(t) + huber, err := HuberRegression(design, contaminated) + if err != nil { + t.Fatalf("HuberRegression: %v", err) + } + weights := append([]float64(nil), huber.Weights...) + sorted := append([]float64(nil), weights...) + slices.Sort(sorted) + lightest := map[float64]bool{} + for _, w := range sorted[:6] { + lightest[w] = true + } + for i := 0; i < 60; i += 10 { + if !lightest[huber.Weights[i]] { + t.Fatalf("contaminated observation %d carries weight %.4f, not among the six lightest", i, huber.Weights[i]) + } + } +} + +// TestHuberExactFit walks the collapsed-scale corner: a constant +// response is reproduced exactly by the unpenalised fit, every +// residual is zero, the robust scale is zero, and the fit reports the +// exactness instead of dividing by it. +func TestHuberExactFit(t *testing.T) { + design := mustFromFloats(t, []float64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4}, 5, 2) + y := mustFromFloats(t, []float64{2.5, 2.5, 2.5, 2.5, 2.5}, 5) + res, err := HuberRegression(design, y) + if err != nil { + t.Fatalf("HuberRegression: %v", err) + } + if !res.Converged { + t.Fatalf("the exact fit did not report convergence") + } + if res.Scale != 0 { + t.Fatalf("the exact fit reports scale %g, want 0", res.Scale) + } + if math.Abs(res.Coefficients[0]-2.5) > 1e-9 || math.Abs(res.Coefficients[1]) > 1e-9 { + t.Fatalf("the exact fit moved: (%.12g, %.12g), want (2.5, 0)", res.Coefficients[0], res.Coefficients[1]) + } + for i := range 5 { + if res.Weights[i] != 1 { + t.Fatalf("an exact fit tapered observation %d to weight %g", i, res.Weights[i]) + } + if math.Abs(res.Residuals[i]) > 1e-12 { + t.Fatalf("the exact fit left residual %g at observation %d", res.Residuals[i], i) + } + if res.StandardErrors[0] != 0 { + t.Fatalf("the exact fit reports standard error %g, want 0", res.StandardErrors[0]) + } + } +} + +// TestTheilSenExactLine pins the estimator on an outlier-free line: +// every pairwise slope of y = 2 + 3x is exactly 3 in floating point, +// so their median is exactly 3 and the intercept exactly 2, bit for +// bit. +func TestTheilSenExactLine(t *testing.T) { + const n = 10 + xs := make([]float64, n) + ys := make([]float64, n) + for i := range n { + xs[i] = float64(i) + ys[i] = 2 + 3*float64(i) + } + x := mustFromFloats(t, xs, n) + y := mustFromFloats(t, ys, n) + intercept, slope, err := TheilSenRegression(x, y) + if err != nil { + t.Fatalf("TheilSenRegression: %v", err) + } + if slope != 3 { + t.Fatalf("slope = %.17g, want exactly 3", slope) + } + if intercept != 2 { + t.Fatalf("intercept = %.17g, want exactly 2", intercept) + } +} + +// TestTheilSenSurvivesBrokenPoint breaks one observation of an exact +// line by a thousand units: the least squares slope is destroyed by +// the leverage, while the Theil-Sen slope and intercept stay exactly +// on the line, because the broken point's slopes sit outside the +// median window. +func TestTheilSenSurvivesBrokenPoint(t *testing.T) { + const n, broken = 25, 7 + xs := make([]float64, n) + ys := make([]float64, n) + for i := range n { + xs[i] = float64(i) + ys[i] = 1 + 0.5*float64(i) + } + ys[broken] += 1000 + x := mustFromFloats(t, xs, n) + y := mustFromFloats(t, ys, n) + design := mustFromFloats(t, []float64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, 1, 6, 1, 7, 1, 8, 1, 9, 1, 10, 1, 11, 1, 12, 1, 13, 1, 14, 1, 15, 1, 16, 1, 17, 1, 18, 1, 19, 1, 20, 1, 21, 1, 22, 1, 23, 1, 24}, n, 2) + ols, err := LinearRegression(design, y) + if err != nil { + t.Fatalf("LinearRegression: %v", err) + } + dragged := math.Abs(ols.Coefficients[1] - 0.5) + t.Logf("one broken point drags the OLS slope to %.4f; Theil-Sen holds it", ols.Coefficients[1]) + if dragged < 3 { + t.Fatalf("the broken point failed to destroy OLS in the first place: drag %.4f", dragged) + } + intercept, slope, err := TheilSenRegression(x, y) + if err != nil { + t.Fatalf("TheilSenRegression: %v", err) + } + if slope != 0.5 { + t.Fatalf("slope = %.17g under one broken point, want exactly 0.5", slope) + } + if intercept != 1 { + t.Fatalf("intercept = %.17g under one broken point, want exactly 1", intercept) + } +} + +// TestTheilSenValidation refuses the inputs the estimator cannot +// answer exactly: wrong shapes, non-finite samples, degenerate +// predictors and samples past the exactness contract. +func TestTheilSenValidation(t *testing.T) { + x := mustFromFloats(t, []float64{0, 1, 2, 3}, 4) + y := mustFromFloats(t, []float64{0, 1, 2, 3}, 4) + if _, _, err := TheilSenRegression(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 4, 2), y); err == nil || !strings.Contains(err.Error(), "the predictor must be rank 1") { + t.Fatalf("a rank 2 predictor: got %v, want the rank refusal", err) + } + if _, _, err := TheilSenRegression(x, mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2)); err == nil || !strings.Contains(err.Error(), "the response must be rank 1") { + t.Fatalf("a rank 2 response: got %v, want the rank refusal", err) + } + if _, _, err := TheilSenRegression(x, mustFromFloats(t, []float64{1, 2, 3}, 3)); err == nil || !strings.Contains(err.Error(), "samples but the response") { + t.Fatalf("a length mismatch: got %v, want the length refusal", err) + } + if _, _, err := TheilSenRegression(x, core.New(core.Complex, 4)); err == nil || !strings.Contains(err.Error(), "complex") { + t.Fatalf("a complex response: got %v, want the complex refusal", err) + } + // Two agreeing observations reach the three-observation floor the + // simple model needs for a median of pairwise slopes. + if _, _, err := TheilSenRegression(mustFromFloats(t, []float64{0, 1}, 2), mustFromFloats(t, []float64{0, 1}, 2)); err == nil || !strings.Contains(err.Error(), "at least three observations") { + t.Fatalf("two observations: got %v, want the three-observation floor", err) + } + if _, _, err := TheilSenRegression(mustFromFloats(t, []float64{2, 2, 2, 2}, 4), y); err == nil || !strings.Contains(err.Error(), "does not vary") { + t.Fatalf("a constant predictor: got %v, want the no-slope refusal", err) + } + if _, _, err := TheilSenRegression(mustFromFloats(t, []float64{0, 1, math.NaN(), 3}, 4), y); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("a non-finite predictor: got %v, want the non-finite refusal", err) + } + if _, _, err := TheilSenRegression(x, mustFromFloats(t, []float64{0, 1, 2, math.Inf(1)}, 4)); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("a non-finite response: got %v, want the non-finite refusal", err) + } + if _, err := HuberRegression(x, y); err == nil || !strings.Contains(err.Error(), "must be rank 2") { + t.Fatalf("a rank 1 design: got %v, want the rank refusal", err) + } + // One past the exactness contract: the refusal names the cost. + big := make([]float64, TheilSenMaxObservations+1) + for i := range big { + big[i] = float64(i) + } + _, _, err := TheilSenRegression(mustFromFloats(t, big, len(big)), mustFromFloats(t, big, len(big))) + if err == nil || !strings.Contains(err.Error(), "exactness contract") { + t.Fatalf("a sample past the exactness contract: got %v, want the exactness refusal", err) + } + // At the contract itself the fit runs. + if _, _, err := TheilSenRegression(mustFromFloats(t, big[:TheilSenMaxObservations], TheilSenMaxObservations), mustFromFloats(t, big[:TheilSenMaxObservations], TheilSenMaxObservations)); err != nil { + t.Fatalf("a sample at the contract was refused: %v", err) + } +} + +// TestRobustEstimatorsOnIntegerArrays exercises the widening +// accessor's fallback paths: an integer design and response reach the +// estimators through FloatAt rather than a raw float payload, and the +// fits must answer exactly as they do on the widened floats. +func TestRobustEstimatorsOnIntegerArrays(t *testing.T) { + design := mustFromInts(t, []int64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5, 1, 6}, 7, 2) + y := mustFromInts(t, []int64{2, 3, 4, 5, 6, 7, 8}, 7) + huber, err := HuberRegressionTuned(design, y, DefaultHuberTuning) + if err != nil { + t.Fatalf("HuberRegression on integer input: %v", err) + } + if math.Abs(huber.Coefficients[0]-2) > 1e-9 || math.Abs(huber.Coefficients[1]-1) > 1e-9 { + t.Fatalf("the integer-input Huber fit is (%.6f, %.6f), want (2, 1)", + huber.Coefficients[0], huber.Coefficients[1]) + } + x := mustFromInts(t, []int64{0, 1, 2, 3, 4, 5, 6}, 7) + intercept, slope, err := TheilSenRegression(x, y) + if err != nil { + t.Fatalf("TheilSenRegression on integer input: %v", err) + } + if slope != 1 || intercept != 2 { + t.Fatalf("the integer-input Theil-Sen fit is (%.17g, %.17g), want exactly (2, 1)", intercept, slope) + } +} + +// TestHuberValidation refuses the malformed inputs, tuning constants +// included: the tuning constant is the estimator's shape, and a zero +// or non-finite one has no Huber loss behind it. +func TestHuberValidation(t *testing.T) { + design := mustFromFloats(t, []float64{1, 0, 1, 1, 1, 2, 1, 3, 1, 4, 1, 5}, 6, 2) + y := mustFromFloats(t, []float64{1, 3, 2, 5, 4, 7}, 6) + if _, err := HuberRegressionTuned(design, y, 0); err == nil { + t.Fatalf("a zero tuning constant was accepted") + } + if _, err := HuberRegressionTuned(design, y, -1); err == nil { + t.Fatalf("a negative tuning constant was accepted") + } + if _, err := HuberRegressionTuned(design, y, math.NaN()); err == nil { + t.Fatalf("a NaN tuning constant was accepted") + } + if _, err := HuberRegressionTuned(design, y, math.Inf(1)); err == nil { + t.Fatalf("an infinite tuning constant was accepted") + } + if _, err := HuberRegressionTuned(design, mustFromFloats(t, []float64{1, 2, 3, math.NaN(), 5, 7}, 6), 1.345); err == nil { + t.Fatalf("a non-finite response was accepted") + } + if _, err := HuberRegressionTuned(mustFromFloats(t, []float64{1, 0, 1, 1}, 2, 2), mustFromFloats(t, []float64{1, 2}, 2), 1.345); err == nil { + t.Fatalf("n = p was accepted") + } + if _, err := HuberRegressionTuned(mustFromFloats(t, []float64{1, 1, 1, 2, 1, 3}, 3, 2), mustFromFloats(t, []float64{1, 2, 3, 4}, 4), 1.345); err == nil { + t.Fatalf("a length mismatch was accepted") + } + complexDesign := core.New(core.Complex, 6, 2) + if _, err := HuberRegressionTuned(complexDesign, y, 1.345); err == nil { + t.Fatalf("complex input was accepted") + } + // A rank-deficient design is refused by the shared solve. + singular := mustFromFloats(t, []float64{1, 2, 1, 2, 1, 2, 1, 2, 1, 2, 1, 2}, 6, 2) + if _, err := HuberRegressionTuned(singular, y, 1.345); err == nil { + t.Fatalf("a rank-deficient design was accepted") + } +} diff --git a/stats/rolling.go b/stats/rolling.go new file mode 100644 index 0000000..0fd79c8 --- /dev/null +++ b/stats/rolling.go @@ -0,0 +1,253 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Windowed (rolling) reductions over a series: each output element +// summarises one window of consecutive samples. The result holds +// n − window + 1 elements, one per full window, aligned so element i +// summarises samples [i, i+window). Series of every real dtype are +// accepted, and the extrema follow the package's NaN rule: a NaN never +// wins a comparison, and a window that holds nothing but NaN answers +// NaN, exactly as core.Min and core.Max do. + +func rollingCheck(a *core.Array, window int) ([]float64, int, error) { + const name = "Rolling" + if a.NDim() != 1 { + return nil, 0, base.Errf("%s: the series must be a vector, got shape %s", name, base.ShapeText(a.Shape())) + } + if a.Dtype() == core.Complex { + return nil, 0, base.Errf("%s: complex series have no ordering to reduce", name) + } + n := a.Len() + if window < 1 || window > n { + return nil, 0, base.Errf("%s: the window must lie in [1, %d], got %d", name, n, window) + } + // Promote through FloatAt, never through RawFloats alone: the + // float64 payload is empty for an int or float32 array, and a + // strided view would be read at the wrong stride. A dense float64 + // array copies its payload directly, the same elements the + // accessor walk returned. The copy also means the fold below + // costs no per-element bounds check. + src := make([]float64, n) + if fs := rawFloats(a); fs != nil { + copy(src, fs) + } else { + for i := range n { + src[i] = a.FloatAt(i) + } + } + return src, n - window + 1, nil +} + +// RollingMean averages each window of the series. +func RollingMean(a *core.Array, window int) (*core.Array, error) { + src, outLen, err := rollingCheck(a, window) + if err != nil { + return nil, err + } + out := core.New(core.Float, outLen) + rollingTotals(src, out.RawFloats(), outLen, window, true) + return out, nil +} + +// RollingSum totals each window of the series. +func RollingSum(a *core.Array, window int) (*core.Array, error) { + src, outLen, err := rollingCheck(a, window) + if err != nil { + return nil, err + } + out := core.New(core.Float, outLen) + rollingTotals(src, out.RawFloats(), outLen, window, false) + return out, nil +} + +// rollingTotals fills the first outLen entries of vals with each full +// window's total of src: the sum, or the mean for mean. +// +// Past a short window the total moves through the series instead of +// rescanning it: it is carried between positions as a two-float pair, +// so one position costs the entering element plus the negated leaving +// one, whatever the window's length, where the rescan costs a fold +// over the whole window at every position. The pair, not a bare +// running sum, is what keeps that honest: a subtractive update sheds +// the low bits of every add and subtract, and on a long window of +// mixed magnitudes the drift ends up wider than what a per-window +// rescan loses, while the pair carries every bit the format holds. +// +// Short windows stay with the rescan, and that is a measured choice, +// not a concession: the rescan's fold costs a couple of cycles a +// sample while the carried walk pays its compensation chain at every +// position whatever the window, so below the crossover the rescan is +// the faster walk by two to three times, and equally accurate. +// +// A window that answers a non-finite total, through a non-finite +// sample or an overflowed sum, answers exactly what the per-window +// rescan answered: the carried pair turns non-finite with it, the +// window is folded from scratch, and the refolded pair becomes the +// carried state for the windows that follow. +func rollingTotals(src []float64, vals []float64, outLen, window int, mean bool) { + if window < rollingIncrementalWindow { + for i := range outLen { + total := 0.0 + for _, v := range src[i : i+window] { + total += v + } + if mean { + total /= float64(window) + } + vals[i] = total + } + return + } + hi, lo := 0.0, 0.0 + for _, v := range src[:window] { + hi, lo = rollingAdd(hi, lo, v) + } + total := hi + lo + if !rollingFinite(total) { + for _, v := range src[:window] { + total += v + } + } + if mean { + total /= float64(window) + } + vals[0] = total + for i := 1; i < outLen; i++ { + hi, lo = rollingAdd(hi, lo, src[i+window-1]) + hi, lo = rollingAdd(hi, lo, -src[i-1]) + total = hi + lo + if !rollingFinite(total) { + // The window answers what the per-window rescan answered, + // and the carried state restarts from the window itself, + // refolded with the same two-float care the walk carries: + // seeding the state from the rescan's rounded total would + // leak that rounding into every window after. + hi, lo, total = 0, 0, 0 + for _, v := range src[i : i+window] { + total += v + hi, lo = rollingAdd(hi, lo, v) + } + } + if mean { + total /= float64(window) + } + vals[i] = total + } +} + +// rollingIncrementalWindow is the window length the rolling totals +// switch from the per-window rescan to the carried update at. The +// rescan's cost per position grows with the window while the carried +// walk's is flat, and the crossover sits near a window of forty; +// sixty-four is the power of two above it, where the carried walk +// already answers twice as fast. +const rollingIncrementalWindow = 64 + +// rollingAdd returns the two-float pair for hi + lo + b. Knuth's +// TwoSum catches the rounding error of the wide add, and the +// renormalisation folds the pair back so the low word stays at the +// rounding level of the high one; the pair then represents the carried +// total to double-double precision across an unbounded walk of adds +// and subtracts. +func rollingAdd(hi, lo, b float64) (float64, float64) { + s := hi + b + bb := s - hi + lo += (hi - (s - bb)) + (b - bb) + t := s + lo + return t, lo - (t - s) +} + +// rollingFinite reports whether v is a finite number: the carried +// total is trusted only while it stays one. +func rollingFinite(v float64) bool { + return math.Abs(v) <= math.MaxFloat64 +} + +// RollingMax tracks each window's largest sample. A NaN never wins a +// comparison and an all-NaN window answers NaN, the package's rule. +func RollingMax(a *core.Array, window int) (*core.Array, error) { + src, outLen, err := rollingCheck(a, window) + if err != nil { + return nil, err + } + return rollingExtreme(src, outLen, window, true), nil +} + +// RollingMin tracks each window's smallest sample, with the same NaN +// rule as RollingMax. +func RollingMin(a *core.Array, window int) (*core.Array, error) { + src, outLen, err := rollingCheck(a, window) + if err != nil { + return nil, err + } + return rollingExtreme(src, outLen, window, false), nil +} + +// rollingExtreme folds each window with a monotonic deque: an index +// leaves the back of the deque only when a later sample is strictly +// better, so the front always holds the window's extreme, and each +// index enters and leaves the deque once. The scan is O(n) where a +// rescan of every window is O(n·window), and it selects the sample the +// rescan selected: a comparison is strict in both, so a tie, the ±0 +// pair included, keeps the earlier index, and a window whose samples +// are all NaN answers its last element, exactly the value the rescan's +// own seed walk ends on. +func rollingExtreme(src []float64, outLen, window int, greater bool) *core.Array { + out := core.New(core.Float, outLen) + vals := out.RawFloats() + // The deque holds indices in ascending order, improving towards the + // back; head is its front, and every entry before head has expired. + // The buffer is compacted once the dead prefix outgrows the live + // region, which keeps it proportional to the deque's depth rather + // than to the series length: each compaction copies at most the + // entries it drops, and every entry is dropped once. + deque := make([]int, 0, min(window, rollingDequeCompact)) + head := 0 + for i, v := range src { + if head >= rollingDequeCompact && head >= len(deque)-head { + deque = deque[:copy(deque, deque[head:])] + head = 0 + } + if !math.IsNaN(v) { + if greater { + for len(deque) > head && src[deque[len(deque)-1]] < v { + deque = deque[:len(deque)-1] + } + } else { + for len(deque) > head && src[deque[len(deque)-1]] > v { + deque = deque[:len(deque)-1] + } + } + deque = append(deque, i) + } + if i+1 < window { + continue + } + oldest := i - window + 1 + for head < len(deque) && deque[head] < oldest { + head++ + } + if head == len(deque) { + // Every sample of the window is NaN; the rescan answers the + // window's last element, and so does this. + vals[oldest] = v + continue + } + vals[oldest] = src[deque[head]] + } + return out +} + +// rollingDequeCompact is the dead-prefix length at which the windowed +// extrema compact their deque buffer: eight kilobytes of indices, past +// which the copy pays for itself on any series long enough to reach it. +const rollingDequeCompact = 1 << 10 diff --git a/stats/rolling_prec_test.go b/stats/rolling_prec_test.go new file mode 100644 index 0000000..6772d51 --- /dev/null +++ b/stats/rolling_prec_test.go @@ -0,0 +1,230 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Precision pins for the incremental rolling totals: the carried +// two-float update is held against an exact big.Float reference on the +// adversarial data the drift shows on, beside the per-window rescan it +// replaced and the uncompensated incremental walk it must beat. + +package stats + +import ( + "math" + "math/big" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// rollingPrecReference folds one window of vals exactly, at 200 bits. +func rollingPrecReference(vals []float64, start, window int) *big.Float { + sum := new(big.Float).SetPrec(200) + for _, v := range vals[start : start+window] { + sum.Add(sum, new(big.Float).SetPrec(200).SetFloat64(v)) + } + return sum +} + +// rollingPrecWalk carries the exact total across the whole series with +// the entering-minus-leaving update at 200 bits: one reference value +// per window position, at a precision the float64 answers cannot see. +func rollingPrecWalk(vals []float64, window int) []*big.Float { + refs := make([]*big.Float, len(vals)-window+1) + refs[0] = rollingPrecReference(vals, 0, window) + for i := 1; i < len(refs); i++ { + ref := new(big.Float).SetPrec(200).Set(refs[i-1]) + ref.Add(ref, new(big.Float).SetPrec(200).SetFloat64(vals[i+window-1])) + ref.Sub(ref, new(big.Float).SetPrec(200).SetFloat64(vals[i-1])) + refs[i] = ref + } + return refs +} + +// rollingPrecWorst returns the worst relative deviation of got from +// the reference walk, with the position it sat at. +func rollingPrecWorst(t *testing.T, label string, got []float64, refs []*big.Float) (float64, int) { + t.Helper() + worst, worstAt := 0.0, 0 + for i, ref := range refs { + refF, _ := ref.Float64() + dev := math.Abs(got[i]-refF) / math.Abs(refF) + if dev > worst { + worst, worstAt = dev, i + } + } + t.Logf("%s: worst relative deviation %.3g at position %d", label, worst, worstAt) + return worst, worstAt +} + +// rollingPrecOldFold is the algorithm the incremental totals replaced: +// every window folded from scratch, exactly as rolling.go carried it +// before. It is the accuracy bar the carried update may not drop +// below. +func rollingPrecOldFold(vals []float64, window int, mean bool) []float64 { + outLen := len(vals) - window + 1 + out := make([]float64, outLen) + for i := range outLen { + total := 0.0 + for _, v := range vals[i : i+window] { + total += v + } + if mean { + total /= float64(window) + } + out[i] = total + } + return out +} + +// rollingPrecPureIncremental is the uncompensated carried update, the +// rejected alternative: enter minus leave on a bare running sum. It is +// kept here, measured, because the drift it suffers is the reason the +// shipped walk carries a two-float pair. +func rollingPrecPureIncremental(vals []float64, window int) []float64 { + outLen := len(vals) - window + 1 + out := make([]float64, outLen) + total := 0.0 + for _, v := range vals[:window] { + total += v + } + out[0] = total + for i := 1; i < outLen; i++ { + total += vals[i+window-1] - vals[i-1] + out[i] = total + } + return out +} + +// TestRollingSumMeanPrecisionAgainstBigFloat holds the incremental +// rolling totals against the exact reference on data chosen for the +// drift: values near 1e8 beside values near 1, signs mixed, windows +// long. The acceptance bar is the accuracy of the per-window rescan: +// the carried pair must sit no further from the reference than the +// rescan does, worst position against worst position. The +// uncompensated walk is measured beside them and reported, and the +// reference itself is spot-checked against direct per-window folds. +func TestRollingSumMeanPrecisionAgainstBigFloat(t *testing.T) { + const n = 40000 + // Alternating magnitudes: the running total sits near 1e11 while + // single units enter and leave it. + alternating := make([]float64, n) + // Mostly small with sparse large spikes, so the carried total is + // dominated by samples long gone from the window. + spiked := make([]float64, n) + // A slow ramp of near-equal large values: every window sum is a + // wide cancellation the plain update is worst at. + ramp := make([]float64, n) + for i := range n { + sign := 1.0 + if (i/64)%2 == 1 { + sign = -1 + } + alternating[i] = sign * 1e8 * (1 + float64(i%3)*0.25) + if i%2 == 1 { + alternating[i] = float64(i%7) - 3 + } + spiked[i] = float64(i%11) - 5 + if i%97 == 0 { + spiked[i] = 1e8 * sign + } + ramp[i] = 1e8 + 0.001*float64(i) + } + for _, tc := range []struct { + name string + vals []float64 + }{ + {"alternating magnitudes", alternating}, + {"sparse spikes", spiked}, + {"near-equal ramp", ramp}, + } { + t.Run(tc.name, func(t *testing.T) { + for _, window := range []int{8, 64, 4096} { + t.Run("window", func(t *testing.T) { + refs := rollingPrecWalk(tc.vals, window) + // The reference walk is itself incremental; prove it + // carries no drift by refolding sampled windows whole. + for start := 0; start < len(refs); start += 2048 { + if diff := rollingPrecReference(tc.vals, start, window); diff.Cmp(refs[start]) != 0 { + t.Fatalf("reference drift at %d: %v vs %v", start, diff, refs[start]) + } + } + sumArr, err := core.FromFloats(tc.vals, len(tc.vals)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + newSum, err := RollingSum(sumArr, window) + if err != nil { + t.Fatalf("RollingSum: %v", err) + } + newMean, err := RollingMean(sumArr, window) + if err != nil { + t.Fatalf("RollingMean: %v", err) + } + gotSum := newSum.RawFloats()[:newSum.Len()] + gotMean := newMean.RawFloats()[:newMean.Len()] + oldSum := rollingPrecOldFold(tc.vals, window, false) + oldMean := rollingPrecOldFold(tc.vals, window, true) + pureSum := rollingPrecPureIncremental(tc.vals, window) + worstSum, _ := rollingPrecWorst(t, "new sum", gotSum, refs) + worstOldSum, _ := rollingPrecWorst(t, "rescan sum", oldSum, refs) + rollingPrecWorst(t, "uncompensated sum", pureSum, refs) + refMeans := make([]*big.Float, len(refs)) + for i, ref := range refs { + refMeans[i] = new(big.Float).SetPrec(200).Quo(ref, big.NewFloat(float64(window))) + } + worstMean, _ := rollingPrecWorst(t, "new mean", gotMean, refMeans) + worstOldMean, _ := rollingPrecWorst(t, "rescan mean", oldMean, refMeans) + if worstSum > worstOldSum { + t.Fatalf("RollingSum window %d: worst relative deviation %.3g is past the rescan's %.3g", + window, worstSum, worstOldSum) + } + if worstMean > worstOldMean { + t.Fatalf("RollingMean window %d: worst relative deviation %.3g is past the rescan's %.3g", + window, worstMean, worstOldMean) + } + }) + } + }) + } +} + +// TestRollingSumMeanNonFiniteMatchesRescan pins the non-finite route: +// a window carrying a NaN or an infinity, and the windows that recover +// after one, answer the per-window rescan's value bit for bit, and the +// carried walk resumes cleanly once every non-finite sample has left. +func TestRollingSumMeanNonFiniteMatchesRescan(t *testing.T) { + vals := make([]float64, 512) + for i := range vals { + vals[i] = float64(i%13) - 6 + 0.125*float64(i%7) + } + vals[5] = math.NaN() + vals[100] = math.Inf(1) + vals[201] = math.Inf(-1) + vals[300] = 1.5e308 + vals[301] = 1.5e308 + a, err := core.FromFloats(vals, len(vals)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + for _, window := range []int{1, 2, 4, 64} { + for _, mean := range []bool{false, true} { + var got *core.Array + if mean { + got, err = RollingMean(a, window) + } else { + got, err = RollingSum(a, window) + } + if err != nil { + t.Fatalf("window %d: %v", window, err) + } + want := rollingPrecOldFold(vals, window, mean) + for i := range want { + gotV, wantV := got.RawFloats()[i], want[i] + if gotV != wantV && !(math.IsNaN(gotV) && math.IsNaN(wantV)) { + t.Fatalf("window %d mean %v: position %d = %.17g, the rescan answers %.17g", + window, mean, i, gotV, wantV) + } + } + } + } +} diff --git a/stats/smallops_test.go b/stats/smallops_test.go new file mode 100644 index 0000000..cf492cd --- /dev/null +++ b/stats/smallops_test.go @@ -0,0 +1,474 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "slices" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +func smallVector(t *testing.T, vals []float64) *core.Array { + t.Helper() + a, err := core.FromFloats(vals, len(vals)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// TestRollingFamily checks the window reductions on a known series: +// alignment, the shrinking length and each reduction's value. +func TestRollingFamily(t *testing.T) { + a := smallVector(t, []float64{1, 4, 2, 8, 5, 7}) + mean, err := RollingMean(a, 3) + if err != nil { + t.Fatalf("RollingMean: %v", err) + } + if mean.Len() != 4 { + t.Fatalf("length %d, want 4", mean.Len()) + } + wantMean := []float64{7.0 / 3, 14.0 / 3, 5, 20.0 / 3} + for i := range wantMean { + if math.Abs(mean.FloatAt(i)-wantMean[i]) > 1e-12 { + t.Fatalf("mean[%d] = %.12g, want %.12g", i, mean.FloatAt(i), wantMean[i]) + } + } + sum, _ := RollingSum(a, 2) + if sum.FloatAt(0) != 5 || sum.FloatAt(4) != 12 { + t.Fatalf("sum ends = %g, %g, want 5, 12", sum.FloatAt(0), sum.FloatAt(4)) + } + max, _ := RollingMax(a, 3) + if max.FloatAt(0) != 4 || max.FloatAt(3) != 8 { + t.Fatalf("max ends = %g, %g, want 4, 8", max.FloatAt(0), max.FloatAt(3)) + } + min, _ := RollingMin(a, 2) + if min.FloatAt(2) != 2 || min.FloatAt(1) != 2 { + t.Fatalf("min = %g, %g, want 2, 2", min.FloatAt(1), min.FloatAt(2)) + } + if _, err := RollingMean(a, 0); err == nil { + t.Fatal("zero window accepted") + } + if _, err := RollingMean(a, 7); err == nil { + t.Fatal("window past the length accepted") + } +} + +// TestMedianAbsoluteDeviation checks the definition on a sample with +// one wild point: the breakdown-proof scale stays at 1, the middle +// deviation, where the standard deviation would explode. +func TestMedianAbsoluteDeviation(t *testing.T) { + a := smallVector(t, []float64{1, 2, 3, 4, 100}) + mad, err := MedianAbsoluteDeviation(a) + if err != nil { + t.Fatalf("MedianAbsoluteDeviation: %v", err) + } + if mad != 1 { + t.Fatalf("MAD = %g, want 1", mad) + } +} + +// TestTrimmedMean checks the tail fraction: dropping one sample a +// side (ten percent) of the contaminated sample discards both tails, +// leaving the clean mean of 2 to 9. +func TestTrimmedMean(t *testing.T) { + a := smallVector(t, []float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 100}) + mean, err := TrimmedMean(a, 0.1) + if err != nil { + t.Fatalf("TrimmedMean: %v", err) + } + if math.Abs(mean-5.5) > 1e-12 { + t.Fatalf("trimmed mean = %.12g, want 5.5", mean) + } + if _, err := TrimmedMean(a, 0.5); err == nil { + t.Fatal("half fraction accepted") + } + if _, err := TrimmedMean(a, -0.1); err == nil { + t.Fatal("negative fraction accepted") + } +} + +// TestHistogram2D checks the count matrix on a constructed pair of +// samples, including the last-bin absorption of the maximum. +func TestHistogram2D(t *testing.T) { + x := smallVector(t, []float64{0.5, 1.5, 2.5, 3.5, 2.9}) + y := smallVector(t, []float64{10, 20, 30, 40, 39}) + counts, xEdges, yEdges, err := Histogram2D(x, y, 4, 2) + if err != nil { + t.Fatalf("Histogram2D: %v", err) + } + if counts.Shape()[0] != 4 || counts.Shape()[1] != 2 { + t.Fatalf("shape %v, want [4 2]", counts.Shape()) + } + if len(xEdges) != 5 || len(yEdges) != 3 { + t.Fatalf("edges %d and %d, want 5 and 3", len(xEdges), len(yEdges)) + } + // x range [0.5, 3.5] over 4 bins of width 0.75; y range [10, 40] + // over 2 bins of width 15. + want := map[[2]int]int{ + {0, 0}: 1, // (0.5, 10) + {1, 0}: 1, // (1.5, 20) + {2, 1}: 1, // (2.5, 30) + {3, 1}: 2, // (2.9, 39) and (3.5, 40) in the last bins + } + ints := counts.RawInts() + for xi := range 4 { + for yi := range 2 { + got := int(ints[xi*2+yi]) + if want[[2]int{xi, yi}] != got { + t.Fatalf("count[%d][%d] = %d, want %d", xi, yi, got, want[[2]int{xi, yi}]) + } + } + } + if _, _, _, err := Histogram2D(smallVector(t, []float64{1}), y, 2, 2); err == nil { + t.Fatal("mismatched lengths accepted") + } + if _, _, _, err := Histogram2D(x, y, 0, 2); err == nil { + t.Fatal("zero bins accepted") + } +} + +// TestWeightedLinearRegression checks the weighted fit on an exact +// line and on a sample whose one wild point the weights silence: the +// unweighted fit chases the outlier, the weighted one recovers the +// line the other three points share. +func TestWeightedLinearRegression(t *testing.T) { + design, err := core.FromFloats([]float64{1, 1, 1, 2, 1, 3, 1, 4}, 4, 2) + if err != nil { + t.Fatalf("design: %v", err) + } + y := smallVector(t, []float64{3, 5, 7, 9}) // y = 2x + 1 exactly + w := smallVector(t, []float64{1, 1, 1, 1}) + res, err := WeightedLinearRegression(design, y, w) + if err != nil { + t.Fatalf("WeightedLinearRegression: %v", err) + } + if math.Abs(res.Coefficients[0]-1) > 1e-9 || math.Abs(res.Coefficients[1]-2) > 1e-9 { + t.Fatalf("coefficients = %g, %g, want 1, 2", res.Coefficients[0], res.Coefficients[1]) + } + // The fourth observation is the outlier now; heavy weights on the + // clean points pull the fit back onto y = 2x + 1. + y2 := smallVector(t, []float64{3, 5, 7, 50}) + w2 := smallVector(t, []float64{1000, 1000, 1000, 0.001}) + res2, err := WeightedLinearRegression(design, y2, w2) + if err != nil { + t.Fatalf("WeightedLinearRegression: %v", err) + } + if math.Abs(res2.Coefficients[1]-2) > 1e-3 { + t.Fatalf("slope = %.6g, want 2 within the weights' pull", res2.Coefficients[1]) + } + // Residuals read in the original units: the clean points sit on + // the line again. + if math.Abs(res2.Residuals[0]) > 1e-2 { + t.Fatalf("first residual = %.6g, want a clean-point zero", res2.Residuals[0]) + } + if _, err := WeightedLinearRegression(design, y, smallVector(t, []float64{1, 1, 1, -1})); err == nil { + t.Fatal("negative weight accepted") + } +} + +// TestHistogram2DParallelCounting pins the paired sweep's split: past +// the per-worker floor the sample is cut into chunks, each worker +// counts into a private matrix and the merge adds every chunk into the +// total. The counts are exact integers, so the split must produce the +// matrix the single-worker sweep produces, cell for cell. +func TestHistogram2DParallelCounting(t *testing.T) { + const n, xBins, yBins = 1 << 16, 16, 16 + xv := make([]float64, n) + yv := make([]float64, n) + for i := range n { + xv[i] = float64((i*7919)%1000)/1000*4 - 2 + yv[i] = float64((i*104729)%1000)/1000*6 - 3 + } + xv[0], xv[1] = -2, 2 + yv[0], yv[1] = -3, 3 + x := smallVector(t, xv) + y := smallVector(t, yv) + + prev := engine.SetNumWorkers(4) + defer engine.SetNumWorkers(prev) + if w := engine.WorkersFor(n); w < 2 { + t.Fatalf("the sweep did not split across workers: %d", w) + } + split, xEdges, yEdges, err := Histogram2D(x, y, xBins, yBins) + if err != nil { + t.Fatalf("Histogram2D across workers: %v", err) + } + engine.SetNumWorkers(1) + serial, _, _, err := Histogram2D(x, y, xBins, yBins) + if err != nil { + t.Fatalf("Histogram2D on one worker: %v", err) + } + if split.Len() != xBins*yBins { + t.Fatalf("the matrix holds %d cells, want %d", split.Len(), xBins*yBins) + } + got, want := split.RawInts(), serial.RawInts() + total := int64(0) + for i := range got { + if got[i] != want[i] { + t.Fatalf("cell %d = %d across workers, %d on one, want the same", i, got[i], want[i]) + } + total += got[i] + } + if total != n { + t.Fatalf("the counts total %d, want %d", total, n) + } + // The binning arithmetic, recomputed from the returned edges. + xLo, xWidth := xEdges[0], xEdges[1]-xEdges[0] + yLo, yWidth := yEdges[0], yEdges[1]-yEdges[0] + ref := make([]int64, xBins*yBins) + for i := range n { + xi := int((xv[i] - xLo) / xWidth) + if xi >= xBins { + xi = xBins - 1 + } + if xi < 0 { + xi = 0 + } + yi := int((yv[i] - yLo) / yWidth) + if yi >= yBins { + yi = yBins - 1 + } + if yi < 0 { + yi = 0 + } + ref[xi*yBins+yi]++ + } + for i := range xBins * yBins { + if got[i] != ref[i] { + t.Fatalf("cell %d = %d, the binning arithmetic gives %d", i, got[i], ref[i]) + } + } +} + +// refWindowExtreme is the retired window rescan: the seed walks to the +// first non-NaN sample and the rest folds with a strict comparison, so +// a tie keeps the earlier sample, the ±0 pair included, and an all-NaN +// window answers its last element. The monotonic deque must select the +// same sample, bit for bit. +func refWindowExtreme(src []float64, window int, greater bool) []float64 { + out := make([]float64, len(src)-window+1) + for i := range out { + w := src[i : i+window] + best, k := w[0], 1 + for math.IsNaN(best) && k < len(w) { + best = w[k] + k++ + } + for _, v := range w[k:] { + if greater { + if v > best { + best = v + } + } else { + if v < best { + best = v + } + } + } + out[i] = best + } + return out +} + +// checkRollingRescan folds both window extrema over the series at every +// window and compares every output bit against the rescan. +func checkRollingRescan(t *testing.T, series []float64) { + t.Helper() + windows := make([]int, len(series)) + for i := range windows { + windows[i] = i + 1 + } + checkRollingRescanAt(t, series, windows...) +} + +// checkRollingRescanAt folds both window extrema over the series at the +// given windows and compares every output bit against the rescan. +func checkRollingRescanAt(t *testing.T, series []float64, windows ...int) { + t.Helper() + for _, window := range windows { + for _, greater := range []bool{true, false} { + var got *core.Array + var err error + if greater { + got, err = RollingMax(smallVector(t, series), window) + } else { + got, err = RollingMin(smallVector(t, series), window) + } + if err != nil { + t.Fatalf("window %d over %v: %v", window, series, err) + } + want := refWindowExtreme(series, window, greater) + if got.Len() != len(want) { + t.Fatalf("window %d: %d outputs, want %d", window, got.Len(), len(want)) + } + for i := range want { + if gb, wb := math.Float64bits(got.FloatAt(i)), math.Float64bits(want[i]); gb != wb { + // A long series is not worth printing: the index, the + // two values and the window say which sample the fold + // picked instead. + if len(series) > 16 { + t.Fatalf("window %d: [%d] = %v (%#x), the rescan gives %v (%#x) over a %d-sample series", + window, i, got.FloatAt(i), gb, want[i], wb, len(series)) + } + t.Fatalf("window %d over %v: [%d] = %v (%#x), the rescan gives %v (%#x)", + window, series, i, got.FloatAt(i), gb, want[i], wb) + } + } + } + } +} + +// TestRollingRescanTiesAndNaN pins the deque's selection rule against +// the retired window rescan on the samples where the rule is visible: +// a tie between +0 and −0 (the earlier sample wins, so the sign of the +// zero is part of the answer), NaN sharing a window with a number, an +// all-NaN window and the infinities. +func TestRollingRescanTiesAndNaN(t *testing.T) { + negZero, posZero := math.Copysign(0, -1), 0.0 + checkRollingRescan(t, []float64{negZero, posZero}) + checkRollingRescan(t, []float64{posZero, negZero}) + checkRollingRescan(t, []float64{negZero, posZero, negZero, posZero}) + checkRollingRescan(t, []float64{negZero, posZero, math.NaN()}) + checkRollingRescan(t, []float64{math.NaN(), negZero, posZero}) + checkRollingRescan(t, []float64{posZero, math.NaN(), negZero}) + checkRollingRescan(t, []float64{math.NaN(), math.NaN(), math.NaN()}) + checkRollingRescan(t, []float64{math.Inf(-1), math.Inf(1), negZero, posZero, math.NaN()}) +} + +// TestRollingRescanCompaction pins the deque's buffer compaction: once +// the dead prefix reaches rollingDequeCompact the buffer is compacted, +// and the compaction must drop dead indices only, never the live entry +// at the deque's front. The series is monotone with a NaN every +// threshold indices: the NaN leaves the front one index ahead of the +// window's own oldest, which is exactly the state a compaction that +// copies from the front instead of from the dead prefix destroys, and +// the window's extreme then comes back from the wrong sample. +func TestRollingRescanCompaction(t *testing.T) { + const n = 4 * rollingDequeCompact + series := make([]float64, n) + for i := range n { + if i > 0 && i%rollingDequeCompact == 0 { + series[i] = math.NaN() + } else { + series[i] = -float64(i) + } + } + checkRollingRescanAt(t, series, 2, 3, rollingDequeCompact) +} + +func TestMedianSliceMatchesSort(t *testing.T) { + // The selection must answer what the sort answered. The pool holds no + // NaN: Go orders NaN nowhere, so "the sorted middle" is not a + // meaningful reference once one is present, and the medians that + // matter (slopes, residuals, deviations) are finite by the time they + // reach here. The zeros are in, because the two zeros are the one + // pair that compares equal while differing in bits, and that is the + // case the selection cannot pin. + pool := []float64{0, math.Copysign(0, -1), -1, 1, 2, -2, 0.5, -0.5, 3, -3, 7, 7, 7} + lens := []int{1, 2, 3, 4, 5, 9, 14, 25, 26, 27, 50, 100, 1000, 4097} + for _, n := range lens { + for shift := range len(pool) { + vals := make([]float64, n) + for i := range vals { + vals[i] = pool[(i+shift)%len(pool)] + } + want := append([]float64(nil), vals...) + slices.Sort(want) + var expected float64 + if n%2 == 1 { + expected = want[n/2] + } else { + expected = want[n/2-1]/2 + want[n/2]/2 + } + got := medianSlice(vals) + if n%2 == 1 { + // An odd length has no tie ambiguity at the middle: the + // value is unique, so the bits must agree outright. + if math.Float64bits(got) != math.Float64bits(expected) { + t.Errorf("n=%d shift=%d: median %v, sorted %v", n, shift, got, expected) + } + } else if got != expected { + t.Errorf("n=%d shift=%d: median %v, sorted %v", n, shift, got, expected) + } + // The selection leaves a permutation of the input. The two + // sorts used for the check are themselves unstable, so the + // comparison is by value and the negative zeros are counted + // separately: that is the one distinction a value comparison + // cannot see. + gotSorted := append([]float64(nil), vals...) + slices.Sort(gotSorted) + for i := range gotSorted { + if gotSorted[i] != want[i] { + t.Fatalf("n=%d shift=%d: elements changed at %d: %v against %v", n, shift, i, gotSorted[i], want[i]) + } + } + if negZero(gotSorted) != negZero(want) { + t.Fatalf("n=%d shift=%d: negative zeros changed: %d against %d", n, shift, negZero(gotSorted), negZero(want)) + } + } + } +} + +func TestMedianSliceNaNStaysInTheInput(t *testing.T) { + // A NaN in the input is not the selection's business: the multiset + // must survive it and the median must be a value that was there. + vals := []float64{1, math.NaN(), 2, 3, 4} + got := medianSlice(vals) + if got == got { + found := false + for _, v := range []float64{1, 2, 3, 4} { + if got == v { + found = true + } + } + if !found { + t.Errorf("median %v is not an input value", got) + } + } + count := 0 + for _, v := range vals { + if math.IsNaN(v) { + count++ + } + } + if count != 1 { + t.Errorf("the selection lost the NaN: %d of them remain", count) + } +} + +func TestMedianSliceOnDistinctValues(t *testing.T) { + // Distinct values have no tie ambiguity at all: the median must be the + // exact order statistic on every length. + for n := 1; n <= 200; n++ { + vals := make([]float64, n) + for i := range vals { + vals[i] = float64((i*37)%101) - 50 + } + want := append([]float64(nil), vals...) + slices.Sort(want) + expected := want[n/2] + if n%2 == 0 { + expected = want[n/2-1]/2 + want[n/2]/2 + } + if got := medianSlice(vals); got != expected { + t.Fatalf("n=%d: median %v, sorted %v", n, got, expected) + } + } +} + +// negZero counts the negative zeros in a slice: the one distinction a +// value comparison cannot see. +func negZero(vals []float64) int { + n := 0 + for _, v := range vals { + if v == 0 && math.Signbit(v) { + n++ + } + } + return n +} diff --git a/stats/stats.go b/stats/stats.go new file mode 100644 index 0000000..13e7725 --- /dev/null +++ b/stats/stats.go @@ -0,0 +1,543 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "sourcedock.dev/petrbalvin/tensor/internal/base" + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +import ( + "math" + "math/bits" + "slices" + "sync" + + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +// The descriptive summaries of the package. Both entry points follow +// the standard conventions: Median always returns float (averaging the +// two middle values on even length) and Std is the population standard +// deviation. + +// rawFloats returns the array's float64 payload when a is a dense +// float64 array and nil otherwise: hot loops branch once on the result +// and sweep the payload directly, falling back to the widening +// accessor for views and other dtypes. The elements are identical +// either way, so every raw sweep computes the same bits as the +// accessor walk it replaces. +func rawFloats(a *core.Array) []float64 { + if !a.Strided() && a.Dtype() == core.Float { + return a.RawFloats() + } + return nil +} + +// Median returns the median as a float64, averaging the two middle values +// when the length is even; an empty, non-finite or complex array is an +// error. +func Median(a *core.Array) (float64, error) { + if a.Dtype() == core.Complex { + return 0, base.Errf("Median: complex arrays have no median") + } + if a.Len() == 0 { + return 0, base.Errf("Median: an empty array has no median") + } + if err := checkFinite("Median", "the sample", a); err != nil { + return 0, err + } + vals := make([]float64, a.Len()) + if fs := rawFloats(a); fs != nil { + copy(vals, fs) + } else if a.Dtype() == core.Int { + // An int sample sorts and averages in its own type: widening + // first rounds every value above 2^53, and the median of + // {2^53, 2^53+1} came back as 2^53 instead of 2^53+0.5. + iv := make([]int64, a.Len()) + for i := range iv { + v, ierr := core.IntAt(a, i) + if ierr != nil { + return 0, base.Errf("Median: %w", ierr) + } + iv[i] = v + } + slices.Sort(iv) + n := len(iv) + if n%2 == 1 { + return float64(iv[n/2]), nil + } + // (x+y)/2 as halved magnitudes plus the carried halves: exact + // whenever the average is representable, correctly rounded + // beyond, and immune to the int64 sum overflow. + x, y := iv[n/2-1], iv[n/2] + return float64((x>>1)+(y>>1)) + float64((x&1)+(y&1))/2, nil + } else { + for i := range vals { + vals[i] = a.FloatAt(i) + } + } + slices.Sort(vals) + n := len(vals) + if n%2 == 1 { + return vals[n/2], nil + } + // (lo+hi)/2 as halved magnitudes summed: each division by two is + // exact, so the sum cannot overflow where the average itself is + // representable, Median([MaxFloat64, MaxFloat64]) being MaxFloat64 + // rather than the +Inf the literal sum produces. The int64 branch + // above is the same averaging shape in integer arithmetic. + lo, hi := vals[n/2-1], vals[n/2] + return lo/2 + hi/2, nil +} + +// sqDeviationBlock folds one block's (v − mean)² over vals[lo:hi] in +// the shape the core block fold keeps: four interleaved chains, so the +// adds of a long block overlap instead of queueing on one adder, +// combined as ((s0+s1)+(s2+s3)). +func sqDeviationBlock(vals []float64, mean float64, lo, hi int) float64 { + var s0, s1, s2, s3 float64 + i := lo + for ; i+4 <= hi; i += 4 { + d0 := vals[i] - mean + d1 := vals[i+1] - mean + d2 := vals[i+2] - mean + d3 := vals[i+3] - mean + s0 += d0 * d0 + s1 += d1 * d1 + s2 += d2 * d2 + s3 += d3 * d3 + } + for ; i < hi; i++ { + d := vals[i] - mean + s0 += d * d + } + return (s0 + s1) + (s2 + s3) +} + +// sqDeviationBlockAt is sqDeviationBlock over an accessor walk: the +// elements are the ones FloatAt returns, so the block answers the same +// bits the payload walk answers. +func sqDeviationBlockAt(a *core.Array, mean float64, lo, hi int) float64 { + var s0, s1, s2, s3 float64 + i := lo + for ; i+4 <= hi; i += 4 { + d0 := a.FloatAt(i) - mean + d1 := a.FloatAt(i+1) - mean + d2 := a.FloatAt(i+2) - mean + d3 := a.FloatAt(i+3) - mean + s0 += d0 * d0 + s1 += d1 * d1 + s2 += d2 * d2 + s3 += d3 * d3 + } + for ; i < hi; i++ { + d := a.FloatAt(i) - mean + s0 += d * d + } + return (s0 + s1) + (s2 + s3) +} + +// sqDeviations sums (v − mean)² over vals through the canonical +// partition the core reductions keep: fixed blocks of the length alone, +// one partial per block, the partials combined through the balanced +// tree. The squared deviations are all non-negative, but a single chain +// a million long still sheds the rounding of every add against an +// accumulator already near the total, and the partition shortens each +// chain by the block count: measured against an exact referent at +// n = 2^20 the error falls by roughly an order of magnitude on +// adversarial magnitude orders, and the fold answers the same bits +// whatever the worker count, the partition being a function of the +// length alone. +func sqDeviations(vals []float64, mean float64) float64 { + n := len(vals) + parts := core.FoldParts(n) + if parts == 1 { + return sqDeviationBlock(vals, mean, 0, n) + } + partials := make([]float64, parts) + for c := range parts { + partials[c] = sqDeviationBlock(vals, mean, core.FoldBoundary(n, c), core.FoldBoundary(n, c+1)) + } + return core.TreeSum(partials) +} + +// sqDeviationsAt is sqDeviations over an accessor walk, the same +// partition and the same block shape, so both routes answer identical +// bits for identical elements. +func sqDeviationsAt(a *core.Array, mean float64) float64 { + n := a.Len() + parts := core.FoldParts(n) + if parts == 1 { + return sqDeviationBlockAt(a, mean, 0, n) + } + partials := make([]float64, parts) + for c := range parts { + partials[c] = sqDeviationBlockAt(a, mean, core.FoldBoundary(n, c), core.FoldBoundary(n, c+1)) + } + return core.TreeSum(partials) +} + +// Std returns the population standard deviation (ddof = 0); an empty or +// complex array is an error. +func Std(a *core.Array) (float64, error) { + if a.Dtype() == core.Complex { + return 0, base.Errf("Std: complex arrays have no float standard deviation") + } + if a.Len() == 0 { + return 0, base.Errf("Std: an empty array has no standard deviation") + } + mean, _ := core.Mean(a) + var sum float64 + if fs := rawFloats(a); fs != nil { + sum = sqDeviations(fs[:a.Len()], mean) + } else { + sum = sqDeviationsAt(a, mean) + } + return math.Sqrt(sum / float64(a.Len())), nil +} + +// Var returns the population variance (ddof = 0, Std squared); an empty +// or complex array is an error. +func Var(a *core.Array) (float64, error) { + return variance(a, 0) +} + +// VarSample returns the unbiased sample variance (ddof = 1); fewer than +// two elements, or a complex array, is an error. +func VarSample(a *core.Array) (float64, error) { + return variance(a, 1) +} + +// variance computes the squared deviation from the mean with the given +// ddof. +func variance(a *core.Array, ddof int) (float64, error) { + name := "Var" + if ddof == 1 { + name = "VarSample" + } + if a.Dtype() == core.Complex { + return 0, base.Errf("%s: complex arrays have no float variance", name) + } + if a.Len() <= ddof { + return 0, base.Errf("%s: needs more than %d element(s)", name, ddof) + } + mean, _ := core.Mean(a) + var sum float64 + if fs := rawFloats(a); fs != nil { + sum = sqDeviations(fs[:a.Len()], mean) + } else { + sum = sqDeviationsAt(a, mean) + } + return sum / float64(a.Len()-ddof), nil +} + +// maxHistBins bounds the bin count a histogram may request. The edges +// and the counts together cost sixteen bytes per bin, so a million bins +// is already a sixteen-megabyte answer to a question no histogram plot +// asks; a larger request is refused instead of handed to the allocator, +// which also keeps every xBins*yBins product far inside an int. +const maxHistBins = 1 << 20 + +// histSerialSamples is the sample count one counting worker must carry +// before the sweep splits across goroutines, and histBinsPerWorker the +// number of samples it must carry per bin on top of that: a worker +// counts into a private array of one cell per bin, and the split only +// pays where the samples counted outweigh the cells that array costs. +// Together they bound the total private scratch by the sample's own +// size. +const ( + histSerialSamples = 1 << 10 + histBinsPerWorker = 8 +) + +// histCountBins adds one slice of the sample to counts: the bin is the +// sample's position on [lo, lo + bins·width], the maximum folds into the +// top bin and anything below the range into the first. +func histCountBins(vals []float64, lo, width float64, bins int, counts []int64) { + for _, v := range vals { + bin := int((v - lo) / width) + if bin >= bins { + bin = bins - 1 // the maximum lands in the top bin + } + if bin < 0 { + bin = 0 + } + counts[bin]++ + } +} + +// Histogram bins the values over [min, max] into bins equal-width bins. +// It returns int counts of length bins and float edges of length bins+1. +// The top bin includes the maximum; an all-equal sample widens to +// [v-0.5, v+0.5]; bins < 1, more than maxHistBins bins, an empty array, +// or a non-finite sample is an error. +func Histogram(a *core.Array, bins int) (*core.Array, *core.Array, error) { + if a.Dtype() == core.Complex { + return nil, nil, base.Errf("Histogram: complex arrays have no histogram") + } + if a.Len() == 0 { + return nil, nil, base.Errf("Histogram: an empty array has no histogram") + } + if bins < 1 { + return nil, nil, base.Errf("Histogram: needs at least one bin, got %d", bins) + } + if bins > maxHistBins { + return nil, nil, base.Errf("Histogram: %d bins exceed the %d-bin limit", bins, maxHistBins) + } + // Dense float64 payloads are swept through the raw slice: the + // elements are the ones FloatAt returns, without the per-element + // dtype dispatch. There the finiteness scan and the range come out of + // one pass; every other layout keeps the accessor scan and the two + // reduction passes. The int branch of the counting below is reached + // only when fs is nil, the branch that fills these two scalars. + fs := rawFloats(a) + var minS, maxS core.Scalar + var lo, hi float64 + if fs != nil { + // Element i of a dense array sits at payload index i, and a + // rebased view's payload may run past its own count: the walk is + // bounded by Len, the elements a caller can see. + fs = fs[:a.Len()] + lo, hi = fs[0], fs[0] + for i, v := range fs { + if math.IsNaN(v) || math.IsInf(v, 0) { + return nil, nil, base.Errf("Histogram: sample %d is not finite (%g)", i, v) + } + if v < lo { + lo = v + } + if v > hi { + hi = v + } + } + } else { + for i := range a.Len() { + if v := a.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return nil, nil, base.Errf("Histogram: sample %d is not finite (%g)", i, v) + } + } + var err error + if minS, err = core.Min(a); err != nil { + return nil, nil, err + } + if maxS, err = core.Max(a); err != nil { + return nil, nil, err + } + lo, hi = minS.Float(), maxS.Float() + } + if lo == hi { + lo -= 0.5 + hi += 0.5 + } + width := (hi - lo) / float64(bins) + if math.IsInf(width, 0) || math.IsNaN(width) { + // A sample holding both float extremes spans more than the + // float64 range: no finite edges exist, and the int((v−lo)/width) + // binning below would clamp every sample into bin 0 over ±Inf + // edges. Refuse rather than publish that. + return nil, nil, base.Errf("Histogram: the sample spans more than the float64 range (%g to %g)", lo, hi) + } + + edges := make([]float64, bins+1) + for i := range bins + 1 { + edges[i] = lo + float64(i)*width + } + + counts := make([]int64, bins) + if fs != nil { + // Counting a bin is exact integer arithmetic and every sample + // carries exactly one increment, so the sweep splits over disjoint + // slices of the payload and the private counters merge in any + // order into the totals a single pass produces. + minPerWorker := max(histSerialSamples, histBinsPerWorker*bins) + var mu sync.Mutex + engine.ParallelMin(len(fs), minPerWorker, func(start, end int) { + if start == 0 && end == len(fs) { + // The whole payload runs inline, the worker policy having + // found the split uneconomical: count straight into the + // result. + histCountBins(fs[start:end], lo, width, bins, counts) + return + } + local := make([]int64, bins) + histCountBins(fs[start:end], lo, width, bins, local) + mu.Lock() + for i, c := range local { + counts[i] += c + } + mu.Unlock() + }) + } else if a.Dtype() == core.Int && maxS.Int() > minS.Int() { + // An int sample is binned on exact integer arithmetic: + // widening v to float64 rounds above 2^53 and has misbinned + // legal samples (nanosecond timestamps live there). The bin is + // floor((v−lo)·bins/(hi−lo)) over the uint64 modular distance, + // the quotient taken from float64 and corrected against exact + // 128-bit products, which matches the rational equal-width + // edges the float path approximates. An all-equal sample took + // the widened float range above and stays on the float path. + loI, hiI := minS.Int(), maxS.Int() + span := uint64(hiI) - uint64(loI) + for i := range a.Len() { + v, verr := core.IntAt(a, i) + if verr != nil { + return nil, nil, base.Errf("Histogram: %w", verr) + } + d := uint64(v) - uint64(loI) + bin := int(float64(d) * float64(bins) / float64(span)) + // One exact correction step each way: the float quotient is + // within a few 2^-32 of the true index, so at most one + // boundary is crossed. + dh, dl := bits.Mul64(d, uint64(bins)) + for { + qh, ql := bits.Mul64(uint64(bin+1), span) + if qh < dh || (qh == dh && ql <= dl) { + bin++ + continue + } + qh, ql = bits.Mul64(uint64(bin), span) + if bin > 0 && (qh > dh || (qh == dh && ql > dl)) { + bin-- + continue + } + break + } + if bin >= bins { + bin = bins - 1 // the maximum lands in the top bin + } + if bin < 0 { + bin = 0 + } + counts[bin]++ + } + } else { + for i := range a.Len() { + v := a.FloatAt(i) + bin := int((v - lo) / width) + if bin >= bins { + bin = bins - 1 // the maximum lands in the top bin + } + if bin < 0 { + bin = 0 + } + counts[bin]++ + } + } + countsArr, err := core.FromInts(counts, bins) + if err != nil { + return nil, nil, err + } + edgesArr, err := core.FromFloats(edges, bins+1) + if err != nil { + return nil, nil, err + } + return countsArr, edgesArr, nil +} + +// BinCounts counts values falling into uniformly spaced bins over +// [min, max]. +func BinCounts(a *core.Array, bins int) (*core.Array, error) { + countsArr, _, err := Histogram(a, bins) + return countsArr, err +} + +// floatsToArray copies data into a new float array of the given shape. +func floatsToArray(data []float64, shape []int) *core.Array { + out := core.New(core.Float, shape...) + copy(out.RawFloats(), data) + return out +} + +// checkFinite scans a real-valued array for a non-finite entry and +// names it, the refusal every estimation entry point of the package +// shares: a single NaN or ±Inf would otherwise spread silently through +// the whole result. The label describes the array in the caller's own +// vocabulary ("the sample", "the covariance"). Callers must have +// refused complex input first, because the scan reads through FloatAt. +func checkFinite(name, label string, a *core.Array) error { + if fs := rawFloats(a); fs != nil { + // A rebased view's payload runs past its own count: only the + // visible elements are scanned. + for _, v := range fs[:a.Len()] { + if math.IsNaN(v) || math.IsInf(v, 0) { + return base.Errf("%s: %s holds the non-finite value %g", name, label, v) + } + } + return nil + } + for i := range a.Len() { + if v := a.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { + return base.Errf("%s: %s holds the non-finite value %g", name, label, v) + } + } + return nil +} + +// Quantile computes q quantiles (each in [0, 1]) of the sample using +// linear interpolation on the sorted values; an empty, non-finite or +// complex array is an error. +func Quantile(a *core.Array, qs []float64) (*core.Array, error) { + if a.Dtype() == core.Complex { + return nil, base.Errf("Quantile: complex arrays have no ordering") + } + n := a.Len() + if n == 0 { + return nil, base.Errf("Quantile: empty array has no quantiles") + } + // A NaN sorts into an arbitrary position and an infinity drags the + // interpolation across it, so the refusal comes before the sort, + // as in every other estimation entry point. + if err := checkFinite("Quantile", "the sample", a); err != nil { + return nil, err + } + sortedVals, err := core.Sort(a) + if err != nil { + return nil, err + } + intSample := a.Dtype() == core.Int + vals := make([]float64, len(qs)) + for qi, q := range qs { + // NaN-rejecting on purpose: NaN compares false against both + // bounds, so `q < 0 || q > 1` would let it through to the index + // conversion, where int(NaN) is a meaningless index. + if !(q >= 0 && q <= 1) { + return nil, base.Errf("Quantile: q must be in [0, 1], got %v", q) + } + pos := q * float64(n-1) + lo := int(math.Floor(pos)) + frac := pos - float64(lo) + if intSample { + // Interpolate on the exact integer difference: widening the + // endpoints first rounds both above 2^53, where a rounded + // pair can even collapse to equal floats and flatten the + // interpolation entirely. + x, xerr := core.IntAt(sortedVals, lo) + if xerr != nil { + return nil, base.Errf("Quantile: %w", xerr) + } + if lo+1 >= n || frac == 0 { + vals[qi] = float64(x) + continue + } + y, yerr := core.IntAt(sortedVals, lo+1) + if yerr != nil { + return nil, base.Errf("Quantile: %w", yerr) + } + d := y - x // exact unless the sorted pair spans more than 2^63 + if (x < 0) != (y < 0) && d < 0 { + vals[qi] = float64(x) + frac*(float64(y)-float64(x)) + continue + } + vals[qi] = float64(x) + frac*float64(d) + continue + } + v := sortedVals.FloatAt(lo) + if lo+1 < n { + v += frac * (sortedVals.FloatAt(lo+1) - v) + } + vals[qi] = v + } + return core.FromFloats(vals, len(qs)) +} diff --git a/stats/stats_random_test.go b/stats/stats_random_test.go new file mode 100644 index 0000000..cac9ea3 --- /dev/null +++ b/stats/stats_random_test.go @@ -0,0 +1,141 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "sourcedock.dev/petrbalvin/tensor/internal/core" + "strings" + "testing" +) + +func TestVar(t *testing.T) { + // The classic set: mean 5, population variance 4. + a := mustFromInts(t, []int64{2, 4, 4, 4, 5, 5, 7, 9}, 8) + v, err := Var(a) + if err != nil { + t.Fatalf("Var: %v", err) + } + if math.Abs(v-4) > 1e-9 { + t.Fatalf("Var: %v", v) + } + + // Sample variance of the same set: 32/7. + vs, err := VarSample(a) + if err != nil { + t.Fatalf("VarSample: %v", err) + } + if math.Abs(vs-32.0/7.0) > 1e-9 { + t.Fatalf("VarSample: %v", vs) + } + + // Std squared is Var. + std, _ := Std(a) + if math.Abs(std*std-v) > 1e-12 { + t.Fatalf("Std² vs Var: %v %v", std*std, v) + } + + one := mustFromFloats(t, []float64{1.5}, 1) + if _, err := VarSample(one); err == nil || !strings.Contains(err.Error(), "more than 1") { + t.Fatalf("VarSample single: %v", err) + } + empty := mustFromInts(t, nil, 0) + if _, err := Var(empty); err == nil || !strings.Contains(err.Error(), "more than 0") { + t.Fatalf("Var empty: %v", err) + } + c := mustFromComplexes(t, []complex128{1}, 1) + if _, err := Var(c); err == nil || !strings.Contains(err.Error(), "no float variance") { + t.Fatalf("Var complex: %v", err) + } +} + +func TestGeneratorNormal(t *testing.T) { + g := core.NewGenerator(11) + draws, err := core.Normal(g, 2000, 10, 2) + if err != nil { + t.Fatalf("Normal: %v", err) + } + if draws.Dtype() != core.Float || draws.Shape()[0] != 2000 { + t.Fatalf("Normal shape: %s", draws) + } + mean, _ := core.Mean(draws) + std, _ := Std(draws) + if mean < 9.8 || mean > 10.2 { + t.Fatalf("Normal mean: %v", mean) + } + if std < 1.8 || std > 2.2 { + t.Fatalf("Normal std: %v", std) + } + + // Determinism: the same seed gives bit-identical draws. + n1, _ := core.Normal(core.NewGenerator(5), 16, 0, 1) + n2, _ := core.Normal(core.NewGenerator(5), 16, 0, 1) + if !core.Equal(n1, n2) { + t.Fatalf("Normal must be reproducible") + } + + if _, err := core.Normal(g, 1, 0, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") { + t.Fatalf("Normal negative std: %v", err) + } + if _, err := core.Normal(g, -1, 0, 1); err == nil || !strings.Contains(err.Error(), "zero or greater") { + t.Fatalf("Normal negative n: %v", err) + } +} + +func TestGeneratorPermutation(t *testing.T) { + g := core.NewGenerator(13) + p, err := core.Permutation(g, 50) + if err != nil { + t.Fatalf("Permutation: %v", err) + } + seen := make([]bool, 50) + for i := range 50 { + v, _ := core.IntAt(p, i) + if v < 0 || v >= 50 || seen[v] { + t.Fatalf("Permutation not a permutation at %d: %d", i, v) + } + seen[v] = true + } + + // Determinism and difference between seeds. + p2, _ := core.Permutation(core.NewGenerator(13), 50) + if !core.Equal(p, p2) { + t.Fatalf("Permutation must be reproducible") + } + p3, _ := core.Permutation(core.NewGenerator(14), 50) + if core.Equal(p, p3) { + t.Fatalf("different seeds must give different permutations") + } + + if _, err := core.Permutation(g, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") { + t.Fatalf("Permutation negative: %v", err) + } +} + +func TestGeneratorShuffle(t *testing.T) { + g := core.NewGenerator(17) + a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 8) + + sh := core.Shuffle(g, a) + if sh.Len() != 8 || sh.Dtype() != core.Int { + t.Fatalf("Shuffle shape: %s", sh) + } + // A shuffle is a permutation of the same multiset. + sorted, _ := core.Sort(sh) + want, _ := core.Sort(a) + if !core.Equal(want, sorted) { + t.Fatalf("Shuffle is not a permutation: %s", sh) + } + // The receiver is untouched. + if v, _ := core.IntAt(a, 0); v != 1 { + t.Fatalf("Shuffle mutated the receiver: %d", v) + } + + // core.Complex arrays shuffle too. + c := mustFromComplexes(t, []complex128{1, complex(2, 2)}, 2) + cs := core.Shuffle(g, c) + if cs.Dtype() != core.Complex || cs.Len() != 2 { + t.Fatalf("Shuffle complex: %s", cs) + } +} diff --git a/stats/stats_test.go b/stats/stats_test.go new file mode 100644 index 0000000..2dd128e --- /dev/null +++ b/stats/stats_test.go @@ -0,0 +1,232 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" + "sourcedock.dev/petrbalvin/tensor/internal/engine" +) + +func TestMedian(t *testing.T) { + odd := mustFromInts(t, []int64{3, 1, 2}, 3) + m, err := Median(odd) + if err != nil || m != 2 { + t.Fatalf("Median odd: %v %v", m, err) + } + + even := mustFromInts(t, []int64{1, 2, 3, 4}, 4) + m, err = Median(even) + if err != nil || m != 2.5 { + t.Fatalf("Median even: %v %v", m, err) + } + + f := mustFromFloats(t, []float64{7.5, 2.5, 5.0}, 3) + m, err = Median(f) + if err != nil || m != 5.0 { + t.Fatalf("Median float: %v %v", m, err) + } + + empty := mustFromInts(t, nil, 0) + if _, err := Median(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("Median empty: %v", err) + } +} + +func TestStd(t *testing.T) { + // The classic set: mean 5, population variance 4, std 2. + a := mustFromInts(t, []int64{2, 4, 4, 4, 5, 5, 7, 9}, 8) + s, err := Std(a) + if err != nil { + t.Fatalf("Std: %v", err) + } + if math.Abs(s-2) > 1e-9 { + t.Fatalf("Std: %v", s) + } + + single := mustFromFloats(t, []float64{3.5}, 1) + s, err = Std(single) + if err != nil || s != 0 { + t.Fatalf("Std single: %v %v", s, err) + } + + empty := mustFromInts(t, nil, 0) + if _, err := Std(empty); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("Std empty: %v", err) + } +} + +func TestHistogram(t *testing.T) { + a := mustFromFloats(t, []float64{1, 1.5, 2, 3, 3.5, 4}, 6) + counts, edges, err := Histogram(a, 3) + if err != nil { + t.Fatalf("Histogram: %v", err) + } + // Bins over [1, 4]: [1,2), [2,3), [3,4]; 4 lands in the top bin. + wantCounts := mustFromInts(t, []int64{2, 1, 3}, 3) + if !core.Equal(wantCounts, counts) { + t.Fatalf("Histogram counts: %s", counts) + } + if edges.Len() != 4 { + t.Fatalf("Histogram edges len: %d", edges.Len()) + } + e0, _ := core.FloatAt(edges, 0) + e3, _ := core.FloatAt(edges, 3) + if e0 != 1 || e3 != 4 { + t.Fatalf("Histogram edges: %s", edges) + } + + // An all-equal sample widens to [v-0.5, v+0.5]. + same := mustFromInts(t, []int64{5, 5}, 2) + _, edges, err = Histogram(same, 2) + if err != nil { + t.Fatalf("Histogram same: %v", err) + } + e0, _ = core.FloatAt(edges, 0) + e2, _ := core.FloatAt(edges, 2) + if e0 != 4.5 || e2 != 5.5 { + t.Fatalf("Histogram same edges: %s", edges) + } + + empty := mustFromInts(t, nil, 0) + if _, _, err := Histogram(empty, 3); err == nil || !strings.Contains(err.Error(), "empty array") { + t.Fatalf("Histogram empty: %v", err) + } + if _, _, err := Histogram(a, 0); err == nil || !strings.Contains(err.Error(), "at least one bin") { + t.Fatalf("Histogram bins: %v", err) + } +} + +// TestHistogramNonFinite pins the finite-sample contract: a NaN or +// infinite sample makes Histogram error instead of producing a +// corrupted or empty binning. +func TestHistogramNonFinite(t *testing.T) { + withNaN := mustFromFloats(t, []float64{1, 2, math.NaN()}, 3) + if _, _, err := Histogram(withNaN, 2); err == nil { + t.Error("expected an error for a NaN sample") + } + withInf := mustFromFloats(t, []float64{1, 2, math.Inf(1)}, 3) + if _, _, err := Histogram(withInf, 2); err == nil { + t.Error("expected an error for an infinite sample") + } +} + +func TestHistogramViewCountsVisibleElements(t *testing.T) { + // A rebased view shares the parent's payload, which runs past the + // view's own count: the sweep must cover the visible elements only, + // exactly as the range reductions and the edges already do. + dense := mustFromFloats(t, []float64{100, 200, 1, 2, 3, 4, 500, 600}, 8) + view, err := core.Slice(dense, 0, 2, 6) + if err != nil { + t.Fatalf("Slice: %v", err) + } + counts, edges, err := Histogram(view, 4) + if err != nil { + t.Fatalf("Histogram: %v", err) + } + got := counts.RawInts()[:counts.Len()] + wantCounts := []int64{1, 1, 1, 1} + for i, w := range wantCounts { + if got[i] != w { + t.Errorf("counts[%d] = %d, want %d (all %v)", i, got[i], w, got) + } + } + gotEdges := edges.RawFloats()[:edges.Len()] + wantEdges := []float64{1, 1.75, 2.5, 3.25, 4} + for i, w := range wantEdges { + if gotEdges[i] != w { + t.Errorf("edges[%d] = %g, want %g (all %v)", i, gotEdges[i], w, gotEdges) + } + } +} + +func TestViewNonFiniteTailIsNotScanned(t *testing.T) { + // The tail of the parent's payload is invisible to the view, so a + // non-finite value there must not refuse an estimation on the view. + dense := mustFromFloats(t, []float64{1, 2, 3, 4, math.NaN(), math.Inf(1)}, 6) + view, err := core.Slice(dense, 0, 0, 4) + if err != nil { + t.Fatalf("Slice: %v", err) + } + if _, err := Median(view); err != nil { + t.Errorf("Median over a view with a non-finite tail: %v", err) + } + if _, _, err := Histogram(view, 2); err != nil { + t.Errorf("Histogram over a view with a non-finite tail: %v", err) + } +} + +// TestHistogramParallelCounting pins the split counting sweep: past the +// per-worker floor the payload is cut into chunks, each worker counts +// into a private array of one cell per bin and the merge adds every +// chunk into the total. The counts are exact integers and every sample +// carries exactly one increment, so the split must produce the very +// counts the single-worker sweep produces, bin for bin. +func TestHistogramParallelCounting(t *testing.T) { + const n, bins = 1 << 15, 256 + vals := make([]float64, n) + for i := range vals { + // Repeated values over a wide range, with the sample's own + // extremes at both ends so the first and the top bin are hit. + vals[i] = float64((i*7919)%1000)/1000*4 - 2 + } + vals[0], vals[1] = -2, 2 + a := mustFloats(t, vals, n) + + prev := engine.SetNumWorkers(4) + defer engine.SetNumWorkers(prev) + if w := engine.WorkersFor(n); w < 2 { + t.Fatalf("the sweep did not split across workers: %d", w) + } + split, splitEdges, err := Histogram(a, bins) + if err != nil { + t.Fatalf("Histogram across workers: %v", err) + } + engine.SetNumWorkers(1) + serial, serialEdges, err := Histogram(a, bins) + if err != nil { + t.Fatalf("Histogram on one worker: %v", err) + } + if split.Len() != bins || splitEdges.Len() != bins+1 { + t.Fatalf("lengths %d and %d, want %d and %d", split.Len(), splitEdges.Len(), bins, bins+1) + } + got, want := split.RawInts(), serial.RawInts() + total := int64(0) + for i := range bins { + if got[i] != want[i] { + t.Fatalf("count[%d] = %d across workers, %d on one, want the same", i, got[i], want[i]) + } + if splitEdges.RawFloats()[i] != serialEdges.RawFloats()[i] { + t.Fatalf("edge[%d] = %g across workers, %g on one", i, + splitEdges.RawFloats()[i], serialEdges.RawFloats()[i]) + } + total += got[i] + } + if total != n { + t.Fatalf("the counts total %d, want %d", total, n) + } + // The binning arithmetic, recomputed from the returned edges: a + // chunk boundary must not move a sample between bins. + lo := splitEdges.RawFloats()[0] + width := splitEdges.RawFloats()[1] - lo + ref := make([]int64, bins) + for _, v := range vals { + b := int((v - lo) / width) + if b >= bins { + b = bins - 1 + } + if b < 0 { + b = 0 + } + ref[b]++ + } + for i := range bins { + if got[i] != ref[i] { + t.Fatalf("count[%d] = %d, the binning arithmetic gives %d", i, got[i], ref[i]) + } + } +} diff --git a/stats/tail_and_extreme_pin_test.go b/stats/tail_and_extreme_pin_test.go new file mode 100644 index 0000000..5a0bd4d --- /dev/null +++ b/stats/tail_and_extreme_pin_test.go @@ -0,0 +1,394 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "math" + "strings" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// Far-tail and extreme-value pins: Wald inference that must not +// report NaN standard errors beside a nil error nor cancel its own +// p-values to zero, a median and a trimmed mean that must not +// overflow on representable samples, NaN samples that must not flow +// through the location summaries, and the guards around them. + +// gaussTailReference returns the two-sided standard normal tail +// 2·(1−Φ(z)) by composite Simpson integration of the Gaussian density +// over [z, z+40]. Every summand is positive, so the sum carries no +// cancellation, and the computation shares nothing with NormalCDF or +// math.Erfc, the pair the tail formulas under test are built on. At +// 2^22 intervals the truncation error sits below the float64 rounding +// of the sum; verified against a 220-bit big.Float quadrature with +// Richardson extrapolation, which reproduces the anchored constant in +// TestGaussTailReferenceAtNine and agrees with math.Erfc to its last +// ulp. +func gaussTailReference(z float64) float64 { + const n = 1 << 22 + h := 40.0 / n + norm := 2 / math.Sqrt(2*math.Pi) + total := 0.0 + for i := 0; i <= n; i++ { + t := z + float64(i)*h + w := 2.0 + if i == 0 || i == n { + w = 1 + } else if i%2 == 1 { + w = 4 + } + total += w * norm * math.Exp(-t*t/2) + } + return total * h / 3 +} + +// TestGaussTailReferenceAtNine anchors the quadrature helper at +// z = 9, where the two-sided tail is 2.2571768119076817e-19. The +// constant comes from an independent 220-bit Simpson quadrature with +// Richardson extrapolation; the cancelled form 2·(1−Φ(9)) answers an +// exact 0, and any fit carrying a z of 9 reports that 0 as its +// p-value before the Erfc repair. +func TestGaussTailReferenceAtNine(t *testing.T) { + const want = 2.2571768119076817e-19 + got := gaussTailReference(9) + if math.Abs(got-want) > 1e-6*want { + t.Fatalf("quadrature reference at z = 9 = %.15g, want %.15g", got, want) + } +} + +// TestPoissonRegressionFarTailPValue drives a fit whose slope carries +// z ≈ 16.4: the old algebraic tail returned an exact 0 there, while +// the true p-value is 2.4e-60, far inside the float64 range. The +// reported p-value must be positive and must match the independent +// quadrature at the achieved z. +func TestPoissonRegressionFarTailPValue(t *testing.T) { + const n = 4000 + design := core.New(core.Float, n, 2) + y := core.New(core.Float, n) + g := core.NewGenerator(3) + for i := range n { + xv := -1 + 2*g.Unit() + design.RawFloats()[i*2] = 1 + design.RawFloats()[i*2+1] = xv + y.RawFloats()[i] = math.Round(math.Exp(0.2 + 0.35*xv)) + } + res, err := PoissonRegression(design, y) + if err != nil { + t.Fatalf("PoissonRegression: %v", err) + } + z := res.ZStatistics[1] + if z < 8.3 { + t.Fatalf("slope z = %g, want a case past the z ≈ 8.3 cancellation cliff", z) + } + p := res.PValues[1] + if p <= 0 { + t.Fatalf("p-value = %g at z = %g, want the representable tail", p, z) + } + wantP := gaussTailReference(z) + if math.Abs(p-wantP) > 1e-6*wantP { + t.Fatalf("p-value = %.15g at z = %.15g, want the quadrature %.15g", p, z, wantP) + } +} + +// TestMannWhitneyUFarTailWithTies pushes the tie-corrected normal +// approximation past the z ≈ 8.3 cliff: a = 1..55 against b = +// 56..109 with 80 duplicated, so u = 0, one tie block of two feeds +// the corrected variance, and the hand-computed z is +// +// z = (1512.5 − 0.5) / sqrt(55·55/12·(111 − 6/(110·109))) ≈ 9.04. +// +// The old cancelled tail returned 0; the tail here is ~1.6e-19. +func TestMannWhitneyUFarTailWithTies(t *testing.T) { + aVals := make([]float64, 0, 55) + for v := 1; v <= 55; v++ { + aVals = append(aVals, float64(v)) + } + bVals := make([]float64, 0, 55) + for v := 56; v <= 109; v++ { + bVals = append(bVals, float64(v)) + } + bVals = append(bVals, 80) + u, p, err := MannWhitneyU(mustFloats(t, aVals), mustFloats(t, bVals)) + if err != nil { + t.Fatalf("MannWhitneyU: %v", err) + } + if u != 0 { + t.Fatalf("u = %g, want 0 for fully separated samples", u) + } + // The same z the test statistic walks, recomputed by hand from the + // known ranks and the single tie block of two. + variance := 55 * 55 / 12.0 * (111 - 6/(110.0*109.0)) + z := (math.Abs(u-55*55/2.0) - 0.5) / math.Sqrt(variance) + if z < 8.3 { + t.Fatalf("z = %g, want a case past the z ≈ 8.3 cancellation cliff", z) + } + if p <= 0 { + t.Fatalf("p-value = %g at z = %g, want the representable tail", p, z) + } + wantP := gaussTailReference(z) + if math.Abs(p-wantP) > 1e-6*wantP { + t.Fatalf("p-value = %.15g at z = %.15g, want the quadrature %.15g", p, z, wantP) + } +} + +// TestGLMWaldNearCollinearDesigns pins the Wald inference contract on +// near-collinear designs: the per-coefficient solve of the inverse +// Fisher information can land a diagonal entry a rounding step below +// zero, where the bare square root produced a NaN standard error and +// NaN p-values beside a nil error. A negative entry is now a named +// error; should a rounding difference keep it positive, the fit is +// still required to answer finite, non-negative standard errors. +// A comfortably identifiable design must fit exactly as before. +func TestGLMWaldNearCollinearDesigns(t *testing.T) { + // Poisson: seed 2, eps 1e-10 converges and then refuses. + const n = 200 + buildPoisson := func(eps float64) (*core.Array, *core.Array) { + x := core.New(core.Float, n, 3) + y := core.New(core.Float, n) + g := core.NewGenerator(2) + for i := range n { + xv := -1 + 2*g.Unit() + x.RawFloats()[i*3] = 1 + x.RawFloats()[i*3+1] = xv + x.RawFloats()[i*3+2] = xv * (1 + eps) + mu := math.Exp(0.2 + 0.5*xv) + y.RawFloats()[i] = math.Round(mu * (1 + (g.Unit()-0.5)*0.1)) + } + return x, y + } + x, y := buildPoisson(1e-10) + res, err := PoissonRegression(x, y) + if err != nil { + if !strings.Contains(err.Error(), "near-collinear") { + t.Fatalf("PoissonRegression on a near-collinear design: %v", err) + } + } else { + for j, se := range res.StandardErrors { + if math.IsNaN(se) || math.IsInf(se, 0) || se < 0 { + t.Fatalf("PoissonRegression standard error %d = %g, want a finite non-negative value", j, se) + } + } + for j, p := range res.PValues { + if math.IsNaN(p) { + t.Fatalf("PoissonRegression p-value %d = NaN on a near-collinear design", j) + } + } + } + // Logistic: seed 3, eps 1e-13 converges and then refuses. + buildLogistic := func(eps float64) (*core.Array, *core.Array) { + x := core.New(core.Float, n, 3) + y := core.New(core.Float, n) + g := core.NewGenerator(3) + for i := range n { + xv := -1 + 2*g.Unit() + x.RawFloats()[i*3] = 1 + x.RawFloats()[i*3+1] = xv + x.RawFloats()[i*3+2] = xv * (1 + eps) + pr := 1 / (1 + math.Exp(-(0.2 + 1.0*xv))) + bit := 0.0 + if g.Unit() < pr { + bit = 1 + } + y.RawFloats()[i] = bit + } + return x, y + } + xl, yl := buildLogistic(1e-13) + resl, errl := LogisticRegression(xl, yl) + if errl != nil { + if !strings.Contains(errl.Error(), "near-collinear") { + t.Fatalf("LogisticRegression on a near-collinear design: %v", errl) + } + } else { + for j, se := range resl.StandardErrors { + if math.IsNaN(se) || math.IsInf(se, 0) || se < 0 { + t.Fatalf("LogisticRegression standard error %d = %g, want a finite non-negative value", j, se) + } + } + for j, p := range resl.PValues { + if math.IsNaN(p) { + t.Fatalf("LogisticRegression p-value %d = NaN on a near-collinear design", j) + } + } + } + // A genuinely identifiable design fits as before, with finite + // inference throughout. The third column is quadratic on purpose: + // x(1+eps) is a scalar multiple of x for every eps, so any such + // design is exactly rank-deficient rather than a healthy contrast. + xh := core.New(core.Float, n, 3) + yh := core.New(core.Float, n) + gh := core.NewGenerator(2) + for i := range n { + xv := -1 + 2*gh.Unit() + xh.RawFloats()[i*3] = 1 + xh.RawFloats()[i*3+1] = xv + xh.RawFloats()[i*3+2] = xv * xv + yh.RawFloats()[i] = math.Round(math.Exp(0.2 + 0.5*xv + 0.3*xv*xv)) + } + resh, errh := PoissonRegression(xh, yh) + if errh != nil { + t.Fatalf("PoissonRegression on a healthy design: %v", errh) + } + for j, se := range resh.StandardErrors { + if !(se > 0) || math.IsInf(se, 0) { + t.Fatalf("healthy PoissonRegression standard error %d = %g", j, se) + } + } + xlh := core.New(core.Float, n, 3) + ylh := core.New(core.Float, n) + glh := core.NewGenerator(3) + for i := range n { + xv := -1 + 2*glh.Unit() + xlh.RawFloats()[i*3] = 1 + xlh.RawFloats()[i*3+1] = xv + xlh.RawFloats()[i*3+2] = xv * xv + pr := 1 / (1 + math.Exp(-(0.2 + 1.0*xv + 0.5*xv*xv))) + bit := 0.0 + if glh.Unit() < pr { + bit = 1 + } + ylh.RawFloats()[i] = bit + } + reslh, errlh := LogisticRegression(xlh, ylh) + if errlh != nil { + t.Fatalf("LogisticRegression on a healthy design: %v", errlh) + } + for j, se := range reslh.StandardErrors { + if !(se > 0) || math.IsInf(se, 0) { + t.Fatalf("healthy LogisticRegression standard error %d = %g", j, se) + } + } +} + +// TestPoissonRegressionAllZeroResponseDoesNotConverge: an all-zero +// count response has its maximum likelihood at minus infinity, the +// iteration can only march towards it, and the fit must report the +// exhausted budget as an error rather than hand back a diverged fit. +// The branch existed without coverage. +func TestPoissonRegressionAllZeroResponseDoesNotConverge(t *testing.T) { + const n = 60 + design := core.New(core.Float, n, 2) + y := core.New(core.Float, n) + for i := range n { + design.RawFloats()[i*2] = 1 + design.RawFloats()[i*2+1] = float64(i % 10) + } + res, err := PoissonRegression(design, y) + if err == nil { + t.Fatalf("PoissonRegression on an all-zero response returned the fit %+v", res) + } + if !strings.Contains(err.Error(), "did not converge") { + t.Fatalf("PoissonRegression on an all-zero response: %v", err) + } +} + +// TestRegressionDesignWithoutColumns: a design with rows but no +// columns passed validation and came back as an empty fit with every +// fitted value at 1. A design must carry at least one column. +func TestRegressionDesignWithoutColumns(t *testing.T) { + design := core.New(core.Float, 5, 0) + y := core.New(core.Float, 5) + if res, err := PoissonRegression(design, y); err == nil { + t.Fatalf("PoissonRegression accepted a column-free design: %+v", res) + } else if !strings.Contains(err.Error(), "at least one column") { + t.Fatalf("PoissonRegression on a column-free design: %v", err) + } + if res, err := LogisticRegression(design, y); err == nil { + t.Fatalf("LogisticRegression accepted a column-free design: %+v", res) + } else if !strings.Contains(err.Error(), "at least one column") { + t.Fatalf("LogisticRegression on a column-free design: %v", err) + } +} + +// TestExponentialCDFLeftTail pins the CDF against the Taylor series +// t − t²/2 in the far left tail, where 1 − e^{−rate·x} cancels: the +// literal form was 11 % off already at rate·x = 1e-16, and answers +// exactly zero not far below. +func TestExponentialCDFLeftTail(t *testing.T) { + for _, rate := range []float64{1, 2} { + for _, tv := range []float64{1e-16, 1e-12, 1e-8, 1e-6} { + x := tv / rate + got, err := ExponentialCDF(x, rate) + if err != nil { + t.Fatalf("ExponentialCDF(%g, %g): %v", x, rate, err) + } + want := tv - tv*tv/2 + if math.Abs(got-want) > 1e-9*want { + t.Fatalf("ExponentialCDF(%g, %g) = %.17g, want the Taylor %.17g", x, rate, got, want) + } + } + } +} + +// TestMedianEvenExtremeValues: the even-length average overflowed on +// magnitudes whose sum leaves the float64 range while the average +// stays inside it. +func TestMedianEvenExtremeValues(t *testing.T) { + big := math.MaxFloat64 + cases := []struct { + vals []float64 + want float64 + }{ + {[]float64{big, big}, big}, + {[]float64{-big, -big}, -big}, + {[]float64{-big, big}, 0}, + // Ordinary even samples keep their averages bit for bit. + {[]float64{1, 2}, 1.5}, + {[]float64{1, 4}, 2.5}, + } + for _, c := range cases { + got, err := Median(mustFloats(t, c.vals)) + if err != nil { + t.Fatalf("Median(%v): %v", c.vals, err) + } + if got != c.want { + t.Fatalf("Median(%v) = %g, want %g", c.vals, got, c.want) + } + } +} + +// TestTrimmedMeanExtremeValues: the direct accumulation overflowed to +// an infinite mean on a window whose true mean is representable. +func TestTrimmedMeanExtremeValues(t *testing.T) { + big := math.MaxFloat64 + got, err := TrimmedMean(mustFloats(t, []float64{big, 1, 2, big}), 0) + if err != nil { + t.Fatalf("TrimmedMean: %v", err) + } + if math.IsInf(got, 0) { + t.Fatalf("TrimmedMean([MaxFloat64, 1, 2, MaxFloat64]) = %g, want a finite mean", got) + } + // The exact mean is MaxFloat64/2 + 0.75, which rounds back to + // MaxFloat64/2: the correction is hundreds of orders below the + // spacing of the answer. + if want := big / 2; got != want { + t.Fatalf("TrimmedMean([MaxFloat64, 1, 2, MaxFloat64]) = %.17g, want %.17g", got, want) + } + if got, err := TrimmedMean(mustFloats(t, []float64{big, -big}), 0); err != nil || got != 0 { + t.Fatalf("TrimmedMean([MaxFloat64, -MaxFloat64]) = %g, %v; want 0, nil", got, err) + } + if got, err := TrimmedMean(mustFloats(t, []float64{1, 2, 3, 4}), 0); err != nil || got != 2.5 { + t.Fatalf("TrimmedMean([1, 2, 3, 4]) = %g, %v; want 2.5, nil", got, err) + } +} + +// TestLocationRefusesNonFiniteSamples: a NaN observation used to flow +// through the location summaries as a plausible number, +// Median([NaN, 1, 2, 3]) being 1.5. +func TestLocationRefusesNonFiniteSamples(t *testing.T) { + if _, err := Median(mustFloats(t, []float64{math.NaN(), 1, 2, 3})); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("Median on a NaN sample: %v", err) + } + if _, err := Median(mustFloats(t, []float64{1, math.Inf(1)})); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("Median on an infinite sample: %v", err) + } + if _, err := Quantile(mustFloats(t, []float64{1, math.NaN(), 2}), []float64{0.5}); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("Quantile on a NaN sample: %v", err) + } + if _, err := TrimmedMean(mustFloats(t, []float64{1, math.NaN(), 3}), 0); err == nil || !strings.Contains(err.Error(), "non-finite") { + t.Fatalf("TrimmedMean on a NaN sample: %v", err) + } +} diff --git a/stats/variance_prec_test.go b/stats/variance_prec_test.go new file mode 100644 index 0000000..90b6c25 --- /dev/null +++ b/stats/variance_prec_test.go @@ -0,0 +1,182 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Precision pins for the variance fold: the squared deviations sum +// through the canonical partition the core reductions keep, and the pin +// holds the fold against an exact big.Rat referent on the data the +// single chain loses on, beside the chain it replaced. The acceptance +// bar is the chain's own error: on every dataset the partition must sit +// no further from the referent than the chain it replaced, and on the +// descending-magnitude data, where every later add rounds against an +// accumulator already near the total, strictly closer. + +package stats + +import ( + "math" + "math/big" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// variancePrecDatasets builds the samples the fold is held on, at +// n = 2^20: a wide-magnitude square wave on a large offset, and a +// geometric spread of deviations walked in descending and ascending +// order of their squares, the shapes the chain loses the most and the +// least on. +func variancePrecDatasets() ([][]float64, []string) { + const n = 1 << 20 + descending := make([]float64, n) + ascending := make([]float64, n) + squareWave := make([]float64, n) + for i := range squareWave { + w := 1.0 + if (i/64)%2 == 1 { + w = -1 + } + squareWave[i] = 1e9 + w + 0.125*float64(i%7-3) + devD := math.Pow(2, 10-20*float64(i)/n) + devA := math.Pow(2, -10+20*float64(i)/n) + descending[i] = 1e6 + devD + ascending[i] = 1e6 + devA + } + return [][]float64{squareWave, descending, ascending}, + []string{"wide-magnitude square wave", "descending squares", "ascending squares"} +} + +// variancePrecExact sums the squared deviations from the exact mean +// exactly in big.Rat, the quantity both folds answer. +func variancePrecExact(vals []float64) *big.Rat { + mean := new(big.Rat) + for _, v := range vals { + mean.Add(mean, new(big.Rat).SetFloat64(v)) + } + mean.Quo(mean, new(big.Rat).SetInt64(int64(len(vals)))) + total := new(big.Rat) + for _, v := range vals { + d := new(big.Rat).SetFloat64(v) + d.Sub(d, mean) + total.Add(total, new(big.Rat).Mul(d, d)) + } + return total +} + +// variancePrecChainFold is the single chain the partition replaced: the +// accuracy bar the fold may not drop below. +func variancePrecChainFold(vals []float64, mean float64) float64 { + sum := 0.0 + for _, v := range vals { + d := v - mean + sum += d * d + } + return sum +} + +// variancePrecFoldRelErr reports the folded sum of squares' relative +// error against the exact referent. +func variancePrecFoldRelErr(t *testing.T, label string, folded float64, exact *big.Rat) *big.Rat { + t.Helper() + got := new(big.Rat).SetFloat64(folded) + got.Sub(got, exact) + got.Abs(got) + if exact.Sign() != 0 { + got.Quo(got, new(big.Rat).Abs(exact)) + } + f, _ := got.Float64() + t.Logf("%s: relative error %.3g", label, f) + return got +} + +// TestVarianceFoldPrecisionAgainstBigRat holds the canonical partition +// against the exact referent on the three datasets, at the same mean +// both folds read, so the comparison isolates the fold. The partition +// must never lose to the chain, and must beat it on the descending data. +func TestVarianceFoldPrecisionAgainstBigRat(t *testing.T) { + datasets, names := variancePrecDatasets() + for k, vals := range datasets { + t.Run(names[k], func(t *testing.T) { + exact := variancePrecExact(vals) + mean := 0.0 + for _, v := range vals { + mean += v + } + mean /= float64(len(vals)) + chain := variancePrecFoldRelErr(t, "single chain", variancePrecChainFold(vals, mean), exact) + partition := variancePrecFoldRelErr(t, "canonical partition", sqDeviations(vals, mean), exact) + if partition.Cmp(chain) > 0 { + t.Fatalf("the canonical partition's error %.3g exceeds the chain's %.3g", + ratFloat(partition), ratFloat(chain)) + } + if k == 1 && partition.Cmp(chain) == 0 { + t.Fatalf("the canonical partition failed to beat the chain on the descending data") + } + }) + } +} + +func ratFloat(r *big.Rat) float64 { + f, _ := r.Float64() + return f +} + +// TestVarianceFoldPathAgreement pins the two routes to one answer: a +// float64 sample rides the raw payload, a float32 sample of the same +// dyadic values rides the accessor, and every widening and every mean +// fold sees identical operands, so the two routes must answer identical +// bits, below and above the one-block partition length alike. +func TestVarianceFoldPathAgreement(t *testing.T) { + build := func(n int) []float64 { + vals := make([]float64, n) + for i := range vals { + // Exact in float32 and in float64 alike, with a magnitude + // spread so the fold has something to chew. + vals[i] = float64(i%97-48) / 4 * math.Pow(2, float64(i%11-5)) + } + return vals + } + for _, n := range []int{3000, 1<<16 + 7} { + vals := build(n) + wide, err := core.FromFloats(vals, n) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + narrow := core.New(core.Float32, n) + for i, v := range vals { + narrow.RawFloat32s()[i] = float32(v) + } + varW, err := Var(wide) + if err != nil { + t.Fatalf("Var: %v", err) + } + varN, err := Var(narrow) + if err != nil { + t.Fatalf("Var narrow: %v", err) + } + stdW, err := Std(wide) + if err != nil { + t.Fatalf("Std: %v", err) + } + stdN, err := Std(narrow) + if err != nil { + t.Fatalf("Std narrow: %v", err) + } + vsW, err := VarSample(wide) + if err != nil { + t.Fatalf("VarSample: %v", err) + } + vsN, err := VarSample(narrow) + if err != nil { + t.Fatalf("VarSample narrow: %v", err) + } + if math.Float64bits(varW) != math.Float64bits(varN) { + t.Fatalf("n = %d: Var answers %.17g on the accessor path and %.17g on the payload path", n, varN, varW) + } + if math.Float64bits(stdW) != math.Float64bits(stdN) { + t.Fatalf("n = %d: Std answers %.17g on the accessor path and %.17g on the payload path", n, stdN, stdW) + } + if math.Float64bits(vsW) != math.Float64bits(vsN) { + t.Fatalf("n = %d: VarSample answers %.17g on the accessor path and %.17g on the payload path", n, vsN, vsW) + } + } +} diff --git a/stats/view_tail_pin_test.go b/stats/view_tail_pin_test.go new file mode 100644 index 0000000..d6e40fd --- /dev/null +++ b/stats/view_tail_pin_test.go @@ -0,0 +1,191 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +// Pins for the rebased-view walks: a contiguous Slice shares its +// source's storage, so its payload runs past the view's own element +// count, and every raw-payload walk in the package is bounded by the +// visible elements. Each pin holds a view's answer against the same +// computation on fresh arrays carrying nothing but the visible +// elements, with junk parked in the invisible tail: a walk that reads +// one payload slot past its count fails the pin, and a walk that +// writes one panics outright. + +package stats + +import ( + "math" + "testing" + + "sourcedock.dev/petrbalvin/tensor/internal/core" +) + +// viewTailPool builds a 60-sample pool whose invisible tail, from +// position 40 on, holds values no caller can see, one of them a NaN. +func viewTailPool(t *testing.T) *core.Array { + t.Helper() + pool := make([]float64, 60) + for i := range pool { + pool[i] = float64(i%9) + 1 + } + pool[55] = math.NaN() // inside the invisible tail: must be ignored + a, err := core.FromFloats(pool, len(pool)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return a +} + +// viewTailHalves slices the pool's head into two disjoint views of ten, +// beside fresh arrays holding exactly the visible elements. +func viewTailHalves(t *testing.T) (va, vb, fa, fb *core.Array) { + t.Helper() + pool := viewTailPool(t) + va, err := core.Slice(pool, 0, 0, 10) + if err != nil { + t.Fatalf("Slice: %v", err) + } + vb, err = core.Slice(pool, 0, 10, 20) + if err != nil { + t.Fatalf("Slice: %v", err) + } + visible := pool.RawFloats()[:20] + fa, err = core.FromFloats(visible[:10], 10) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + fb, err = core.FromFloats(visible[10:], 10) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + return va, vb, fa, fb +} + +// TestViewTailANOVAOneWay pins the analysis of variance on rebased +// views: the group sums, the finite scan and the within sum of squares +// read the visible elements only. +func TestViewTailANOVAOneWay(t *testing.T) { + va, vb, fa, fb := viewTailHalves(t) + f1, p1, err1 := ANOVAOneWay([]*core.Array{va, vb}) + f2, p2, err2 := ANOVAOneWay([]*core.Array{fa, fb}) + if (err1 == nil) != (err2 == nil) { + t.Fatalf("ANOVAOneWay on views errored %v, on fresh arrays %v", err1, err2) + } + if err1 == nil && (f1 != f2 || p1 != p2) { + t.Fatalf("ANOVAOneWay on views = (%v, %v), fresh arrays answer (%v, %v)", f1, p1, f2, p2) + } +} + +// TestViewTailWelchTTest pins Welch's t-test on rebased views: the +// sample mean and variance read the visible elements only. +func TestViewTailWelchTTest(t *testing.T) { + va, vb, fa, fb := viewTailHalves(t) + t1, df1, p1, err1 := WelchTTest(va, vb) + t2, df2, p2, err2 := WelchTTest(fa, fb) + if (err1 == nil) != (err2 == nil) { + t.Fatalf("WelchTTest on views errored %v, on fresh arrays %v", err1, err2) + } + if err1 == nil && (t1 != t2 || df1 != df2 || p1 != p2) { + t.Fatalf("WelchTTest on views = (%v, %v, %v), fresh arrays answer (%v, %v, %v)", + t1, df1, p1, t2, df2, p2) + } +} + +// TestViewTailMAD pins the median absolute deviation on a rebased +// view: the deviation walk writes one cell per visible element, never +// one per payload slot. +func TestViewTailMAD(t *testing.T) { + pool := viewTailPool(t) + va, err := core.Slice(pool, 0, 0, 10) + if err != nil { + t.Fatalf("Slice: %v", err) + } + visible := pool.RawFloats()[:10] + fa, err := core.FromFloats(visible, 10) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + got, err1 := MedianAbsoluteDeviation(va) + want, err2 := MedianAbsoluteDeviation(fa) + if err1 != nil || err2 != nil { + t.Fatalf("MedianAbsoluteDeviation: %v against %v", err1, err2) + } + if got != want { + t.Fatalf("MedianAbsoluteDeviation on a view = %v, fresh array answers %v", got, want) + } +} + +// TestViewTailChiSquareGoodnessOfFit pins the goodness-of-fit test on +// rebased views: both frequency walks are bounded by the bin count, +// where an unbounded index into the other payload was a panic before it +// was a wrong statistic. +func TestViewTailChiSquareGoodnessOfFit(t *testing.T) { + pool := viewTailPool(t) + obsV, err := core.Slice(pool, 0, 20, 30) + if err != nil { + t.Fatalf("Slice: %v", err) + } + expV, err := core.Slice(pool, 0, 30, 40) + if err != nil { + t.Fatalf("Slice: %v", err) + } + visible := pool.RawFloats()[:40] + fObs, err := core.FromFloats(visible[20:30], 10) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + fExp, err := core.FromFloats(visible[30:40], 10) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + c1, df1, p1, errV := ChiSquareGoodnessOfFit(obsV, expV) + c2, df2, p2, errF := ChiSquareGoodnessOfFit(fObs, fExp) + if (errV == nil) != (errF == nil) { + t.Fatalf("ChiSquareGoodnessOfFit on views errored %v, on fresh arrays %v", errV, errF) + } + if errV == nil && (c1 != c2 || df1 != df2 || p1 != p2) { + t.Fatalf("ChiSquareGoodnessOfFit on views = (%v, %d, %v), fresh arrays answer (%v, %d, %v)", + c1, df1, p1, c2, df2, p2) + } +} + +// TestViewTailLinearRegression pins the fit on rebased views: the +// design and response scans refuse only what the visible rows hold, so +// junk in the invisible tail cannot reject a clean fit. +func TestViewTailLinearRegression(t *testing.T) { + design := make([]float64, 40) + for i := range design { + design[i] = float64(i % 7) + } + resp := make([]float64, 20) + for i := range resp { + resp[i] = float64(i%5) * 0.5 + } + // Both junk slots sit past their views' visible rows: design row 15 + // is invisible to the ten-row design view, response element 15 to + // the ten-row response view. + design[30] = math.NaN() + resp[15] = math.NaN() + designArr, err := core.FromFloats(design, len(design)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + respArr, err := core.FromFloats(resp, len(resp)) + if err != nil { + t.Fatalf("FromFloats: %v", err) + } + design2d, err := core.Reshape(designArr, 20, 2) + if err != nil { + t.Fatalf("Reshape: %v", err) + } + designView, err := core.Slice(design2d, 0, 0, 10) + if err != nil { + t.Fatalf("Slice: %v", err) + } + respView, err := core.Slice(respArr, 0, 0, 10) + if err != nil { + t.Fatalf("Slice: %v", err) + } + if _, err := LinearRegression(designView, respView); err != nil { + t.Fatalf("LinearRegression refused clean visible rows: %v", err) + } +} diff --git a/stats/zz_guard_test.go b/stats/zz_guard_test.go new file mode 100644 index 0000000..8fd8c21 --- /dev/null +++ b/stats/zz_guard_test.go @@ -0,0 +1,40 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: MIT + +package stats + +import ( + "os" + "runtime" + "testing" + "time" +) + +// TestMain installs a heap watchdog for the whole package run: the +// review probes include bin-count and quantile cases that can ask for +// tens of gigabytes when a guard regresses, and an editor session was +// already lost to a runaway test allocation once. Any runaway now panics +// and unwinds instead of taking the host with it. +func TestMain(m *testing.M) { + const cap = 3 << 30 // race instrumentation inflates the heap; a runaway is far above this + stop := make(chan struct{}) + go func() { + ticker := time.NewTicker(20 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-stop: + return + case <-ticker.C: + var m runtime.MemStats + runtime.ReadMemStats(&m) + if m.HeapAlloc > cap { + panic("stats test binary heap above 3 GiB: an entry point is allocating without bound") + } + } + } + }() + code := m.Run() + close(stop) + os.Exit(code) +}