// Copyright (c) 2026 Petr BalvĂ­n (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 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)) } } }) }