// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "math" "sourcedock.dev/petrbalvin/tensor/internal/core" "testing" ) // Regression pins for the broadcast and shape backwards: a broadcast // gradient must collapse to its source's shape, and every widened op // must hand each operand its own correctly shaped gradient buffer. // TestBroadcastRank1Gradient pins the rank-1 broadcast backward: the // gradient of a size-1 source broadcast to length m must collapse back // to a single sum, not arrive with the broadcast shape. func TestBroadcastRank1Gradient(t *testing.T) { x, err := core.FromFloats([]float64{0}, 1) if err != nil { t.Fatalf("FromFloats: %v", err) } xt := FromArray(x, true) y, err := xt.BroadcastTo(3) if err != nil { t.Fatalf("BroadcastTo: %v", err) } w, err := core.FromFloats([]float64{1, 2, 3}, 3) if err != nil { t.Fatalf("FromFloats: %v", err) } wt := FromArray(w, false) prod, err := y.Mul(wt) if err != nil { t.Fatalf("Mul: %v", err) } loss, err := prod.Sum() if err != nil { t.Fatalf("Sum: %v", err) } if err := loss.Backward(); err != nil { t.Fatalf("Backward: %v", err) } g := xt.Grad() if g.NDim() != 1 || g.Shape()[0] != 1 { t.Fatalf("gradient shape %v, want [1]", g.Shape()) } if math.Abs(g.FloatAt(0)-6) > 1e-12 { t.Fatalf("gradient = %g, want 6", g.FloatAt(0)) } } // TestBroadcastRank1ToMatrix pins the (1,) to (m, n) broadcast backward // against central differences. func TestBroadcastRank1ToMatrix(t *testing.T) { x, err := core.FromFloats([]float64{2}, 1) if err != nil { t.Fatalf("FromFloats: %v", err) } xt := FromArray(x, true) y, err := xt.BroadcastTo(2, 3) if err != nil { t.Fatalf("BroadcastTo: %v", err) } w, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) if err != nil { t.Fatalf("FromFloats: %v", err) } wt := FromArray(w, false) prod, err := y.Mul(wt) if err != nil { t.Fatalf("Mul: %v", err) } loss, err := prod.Sum() if err != nil { t.Fatalf("Sum: %v", err) } if err := loss.Backward(); err != nil { t.Fatalf("Backward: %v", err) } g := xt.Grad() if g.NDim() != 1 || g.Shape()[0] != 1 { t.Fatalf("gradient shape %v, want [1]", g.Shape()) } if math.Abs(g.FloatAt(0)-21) > 1e-12 { t.Fatalf("gradient = %g, want 21", g.FloatAt(0)) } } // TestL2NormAxisEmptyDim pins the backward against the integer division // by zero an empty reduced dimension used to hit. func TestL2NormAxisEmptyDim(t *testing.T) { x, err := core.FromFloats([]float64{}, 3, 0) if err != nil { t.Fatalf("FromFloats: %v", err) } xt := FromArray(x, true) n, err := xt.L2NormAxis(1) if err != nil { t.Fatalf("L2NormAxis: %v", err) } loss, err := n.Sum() if err != nil { t.Fatalf("Sum: %v", err) } if err := loss.Backward(); err != nil { t.Fatalf("Backward: %v", err) } g := xt.Grad() if g.Len() != 0 { t.Fatalf("gradient length %d, want 0", g.Len()) } } // TestShapeOpsRejectNonFloat pins the dtype contract on the shape ops // that used to record graph nodes without validating the dtype. func TestShapeOpsRejectNonFloat(t *testing.T) { i, err := core.FromInts([]int64{1, 2, 3, 4}, 2, 2) if err != nil { t.Fatalf("FromInts: %v", err) } it := FromArray(i, true) if _, err := it.Transpose(); err == nil { t.Error("Transpose accepted an int tensor") } if _, err := it.Squeeze(0); err == nil { t.Error("Squeeze accepted an int tensor") } if _, err := it.Unsqueeze(0); err == nil { t.Error("Unsqueeze accepted an int tensor") } if _, err := it.Reshape(4); err == nil { t.Error("Reshape accepted an int tensor") } if _, err := it.TransposeAxes(1, 0); err == nil { t.Error("TransposeAxes accepted an int tensor") } if _, err := it.Floor(); err == nil { t.Error("Floor accepted an int tensor") } if _, err := it.BroadcastTo(2, 3); err == nil { t.Error("BroadcastTo accepted an int tensor") } }