188 lines
4.7 KiB
Go
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))
|
|
}
|
|
}
|
|
})
|
|
}
|