Files
tensor/grad/fuzz_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

188 lines
4.7 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Property fuzz targets over index-moving operations. `go test` runs
// the seed corpus on every commit; longer campaigns run under
// -fuzz=Fuzz<Name> when a shape-handling change lands.
// FuzzTransposeAxesRoundTrip drives random permutations through the
// axis move and its inverse: whatever valid permutation arrives, the
// double transpose must restore the exact element order.
func FuzzTransposeAxesRoundTrip(f *testing.F) {
f.Add([]byte{0, 1}, 6)
f.Add([]byte{1, 0}, 6)
f.Add([]byte{2, 0, 1}, 8)
f.Fuzz(func(t *testing.T, permBytes []byte, total int) {
if total <= 0 || total > 4096 {
t.Skip()
}
rank := len(permBytes)
switch rank {
case 2:
total -= total % 2
case 3:
total -= total % 4
default:
t.Skip()
}
if total == 0 {
t.Skip()
}
vals := make([]float64, total)
for i := range vals {
vals[i] = float64(i)
}
var shape []int
if rank == 2 {
shape = []int{total / 2, 2}
} else {
shape = []int{total / 4, 2, 2}
}
a, _ := core.FromFloats(vals, shape...)
xt := FromArray(a, false)
dims := make([]int, rank)
for i, pb := range permBytes {
dims[i] = int(pb) % rank
}
moved, err := xt.TransposeAxes(dims...)
if err != nil {
return // duplicate axes rejected by validation, fine
}
back, err := moved.TransposeAxes(inversePerm(dims)...)
if err != nil {
t.Fatalf("inverse of %v failed: %v", dims, err)
}
for i := range vals {
if back.Data().FloatAt(i) != vals[i] {
t.Fatalf("round trip lost element %d", i)
}
}
})
}
// FuzzOneHotContracts checks both sides of the encoder contract for
// arbitrary code sets: in-range codes yield exactly one hot cell per
// row, any out-of-range code is a loud error. One input byte splits
// into a high bit forcing negativity plus a low-bit class selector.
func FuzzOneHotContracts(f *testing.F) {
f.Add([]byte{0, 1, 2}, uint8(3))
f.Add([]byte{5}, uint8(8))
f.Add([]byte{200, 201}, uint8(3))
f.Fuzz(func(t *testing.T, raw []byte, classByte uint8) {
classes := int(classByte)%9 + 1
codes := make([]int64, len(raw))
valid := true
for i, b := range raw {
c := int64(b)
if b >= 128 { // force some negative probes
c = -int64(b - 127)
} else {
c %= int64(classes)
}
if c < 0 || c >= int64(classes) {
valid = false
}
codes[i] = c
}
arr, _ := core.FromInts(codes, len(codes))
hot, err := core.OneHot(arr, classes)
if !valid {
if err == nil {
t.Fatalf("invalid codes accepted for %d classes", classes)
}
return
}
if err != nil {
t.Fatalf("valid codes rejected: %v", err)
}
for i := range len(codes) {
sum := 0.0
for j := range classes {
sum += float64(hot.FloatAt(i*classes + j))
}
if sum != 1 {
t.Fatalf("row %d sums to %v, want one hot cell", i, sum)
}
}
})
}
// FuzzConcatSplitGradientConserves mass: splitting a concatenated
// output's gradient must hand every element back to its own side with
// coefficient exactly one, for whatever layout the corpus invents.
func FuzzConcatSplitGradientConserves(f *testing.F) {
f.Add([]byte{1, 2, 3, 4}, uint8(2))
f.Add([]byte{9, 7, 5}, uint8(1))
f.Add([]byte{10, 20, 30, 40, 50, 60, 70, 80}, uint8(0))
f.Add([]byte{11, 21, 31, 41, 51, 61, 71, 81, 91, 101}, uint8(4))
f.Fuzz(func(t *testing.T, raw []byte, rowsByte uint8) {
rows := int(rowsByte)%7 + 1 // left matrix rows, 1..7
leftLen := rows * 2
if len(raw) <= leftLen {
t.Skip()
}
rightRows := (len(raw) - leftLen) / 2
leftVals := make([]float64, leftLen)
for i := range leftVals {
leftVals[i] = float64(raw[i])
}
rightVals := make([]float64, rightRows*2)
for i := range rightVals {
rightVals[i] = float64(raw[leftLen+i])
}
a, err := core.FromFloats(leftVals, rows, 2)
if err != nil {
t.Skip()
}
b, err := core.FromFloats(rightVals, rightRows, 2)
if err != nil {
t.Skip()
}
at := FromArray(a, true)
bt := FromArray(b, true)
joint, cerr := at.Concat(bt, 0)
if cerr != nil {
t.Fatal(cerr)
}
loss, serr := joint.Sum()
if serr != nil {
t.Fatal(serr)
}
if berr := loss.Backward(); berr != nil {
t.Fatal(berr)
}
ga, gb := at.Grad(), bt.Grad()
if ga.Len() != a.Len() || gb.Len() != b.Len() {
t.Fatalf("gradient spans drifted: %d+%d vs %d+%d",
ga.Len(), gb.Len(), a.Len(), b.Len())
}
for i := range ga.Len() {
if ga.FloatAt(i) != 1 {
t.Fatalf("left span slot %d = %v", i, ga.FloatAt(i))
}
}
for i := range gb.Len() {
if gb.FloatAt(i) != 1 {
t.Fatalf("right span slot %d = %v", i, gb.FloatAt(i))
}
}
})
}