201 lines
4.6 KiB
Go
201 lines
4.6 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"
|
|
)
|
|
|
|
func TestTensorSqueezeUnsqueezeClip(t *testing.T) {
|
|
// Squeeze/Unsqueeze round-trip with gradient.
|
|
x, _ := core.FromFloats([]float64{1, 2, 3, 4}, 1, 4, 1)
|
|
xt := FromArray(x, true)
|
|
sq, err := xt.Squeeze(2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if sq.Data().NDim() != 2 {
|
|
t.Fatalf("Squeeze ndim: %d", sq.Data().NDim())
|
|
}
|
|
back, err := sq.Unsqueeze(2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
s, _ := back.Sum()
|
|
if err := s.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for i := range 4 {
|
|
if g := xt.Grad().FloatAt(i); g != 1 {
|
|
t.Errorf("Squeeze/Unsqueeze grad[%d]: %v, want 1", i, g)
|
|
}
|
|
}
|
|
|
|
// Clip gradient: 1 inside [lo, hi], 0 outside.
|
|
c, _ := core.FromFloats([]float64{-1, 0.5, 2}, 3)
|
|
ct := FromArray(c, true)
|
|
cl, err := ct.Clip(0, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
s2, _ := cl.Sum()
|
|
if err := s2.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
want := []float64{0, 1, 0}
|
|
for i := range 3 {
|
|
if g := ct.Grad().FloatAt(i); g != want[i] {
|
|
t.Errorf("Clip grad[%d]: %v, want %v", i, g, want[i])
|
|
}
|
|
}
|
|
|
|
}
|
|
func TestAxisReductionAutograd(t *testing.T) {
|
|
x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
xt := FromArray(x, true)
|
|
|
|
s, err := xt.SumAxis(1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if s.Data().Len() != 2 {
|
|
t.Fatalf("SumAxis len: %d", s.Data().Len())
|
|
}
|
|
if err := s.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for i := range x.Len() {
|
|
if g := xt.Grad().FloatAt(i); g != 1 {
|
|
t.Errorf("SumAxis grad[%d]: %v, want 1", i, g)
|
|
}
|
|
}
|
|
|
|
mean, err := FromArray(x, true).MeanAxis(1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := mean.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func TestL2NormAxisAutogradGradient(t *testing.T) {
|
|
xv := []float64{3, 4, 0.5, 0.5}
|
|
x, _ := core.FromFloats(xv, 1, 1, 2, 2)
|
|
xt := FromArray(x, true)
|
|
out, err := xt.L2NormAxis(1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
s, _ := out.Sum()
|
|
if err := s.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
analytic := make([]float64, x.Len())
|
|
for i := range x.Len() {
|
|
analytic[i] = xt.Grad().FloatAt(i)
|
|
}
|
|
ref := numericGrad(func(a *core.Array) float64 {
|
|
o, err := FromArray(a, false).L2NormAxis(1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ss, _ := o.Sum()
|
|
return ss.Data().FloatAt(0)
|
|
}, x)
|
|
if d := maxAbsDiff(analytic, ref); d > 1e-6 {
|
|
t.Errorf("L2NormAxis grad: max diff %v", d)
|
|
}
|
|
}
|
|
|
|
func TestBroadcastToAutograd(t *testing.T) {
|
|
x, _ := core.FromFloats([]float64{1, 2, 3}, 1, 3)
|
|
xt := FromArray(x, true)
|
|
out, err := xt.BroadcastTo(2, 3)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if out.Data().Shape()[0] != 2 {
|
|
t.Fatalf("BroadcastTo shape: %v", out.Data().Shape())
|
|
}
|
|
onesArr, _ := core.Ones(core.Float, 2, 3)
|
|
loss, _ := out.Mul(FromArray(onesArr, false))
|
|
s, _ := loss.Sum()
|
|
if err := s.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
// Gradient sums over replicated rows.
|
|
for i := range 3 {
|
|
if g := xt.Grad().FloatAt(i); g != 2 {
|
|
t.Errorf("BroadcastTo grad[%d]: %v, want 2", i, g)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestPowAbsSqrtFloorAutogradGradient(t *testing.T) {
|
|
cases := []struct {
|
|
name string
|
|
vals []float64
|
|
fn func(*Tensor) (*Tensor, error)
|
|
}{
|
|
{"Pow3", []float64{0.5, 1.5}, func(x *Tensor) (*Tensor, error) { return x.Pow(3) }},
|
|
{"Abs", []float64{0.5, -1.5}, func(x *Tensor) (*Tensor, error) { return x.Abs() }},
|
|
{"Sqrt", []float64{0.25, 2.25}, func(x *Tensor) (*Tensor, error) { return x.Sqrt() }},
|
|
}
|
|
for _, tc := range cases {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
a, _ := core.FromFloats(tc.vals, 2)
|
|
at := FromArray(a, true)
|
|
out, err := tc.fn(at)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
s, err := out.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := s.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
analytic := make([]float64, a.Len())
|
|
for i := range a.Len() {
|
|
analytic[i] = at.Grad().FloatAt(i)
|
|
}
|
|
ref := numericGrad(func(v *core.Array) float64 {
|
|
o, err := tc.fn(FromArray(v, false))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ss, err := o.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return ss.Data().FloatAt(0)
|
|
}, a)
|
|
if d := maxAbsDiff(analytic, ref); d > 1e-6 {
|
|
t.Errorf("%s grad: max diff %v", tc.name, d)
|
|
}
|
|
})
|
|
}
|
|
|
|
// Floor contributes no gradient.
|
|
a, _ := core.FromFloats([]float64{1.4, 2.6}, 2)
|
|
at := FromArray(a, true)
|
|
fl, err := at.Floor()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
s, _ := fl.Sum()
|
|
if err := s.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for i := range a.Len() {
|
|
if g := at.Grad().FloatAt(i); g != 0 {
|
|
t.Errorf("Floor grad[%d]: %v, want 0", i, g)
|
|
}
|
|
}
|
|
}
|