feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,357 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user