358 lines
10 KiB
Go
358 lines
10 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
|
}
|