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