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