101 lines
2.5 KiB
Go
101 lines
2.5 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 TestTensorTransposeAxesValues(t *testing.T) {
|
|
x, _ := core.FromFloats([]float64{
|
|
1, 2, 3,
|
|
4, 5, 6,
|
|
}, 2, 3)
|
|
|
|
out, err := FromArray(x, false).TransposeAxes(1, 0)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := out.Data().Shape(); got[0] != 3 || got[1] != 2 {
|
|
t.Fatalf("shape: %v", got)
|
|
}
|
|
want := []float64{1, 4, 2, 5, 3, 6}
|
|
for i := range want {
|
|
if g := out.Data().FloatAt(i); g != want[i] {
|
|
t.Fatalf("[%d] = %v, want %v", i, g, want[i])
|
|
}
|
|
}
|
|
|
|
// A rank-3 rotation moves the trailing axis to the front.
|
|
y, _ := core.FromFloats([]float64{
|
|
1, 2, 3, 4,
|
|
5, 6, 7, 8,
|
|
9, 10, 11, 12,
|
|
}, 2, 3, 2)
|
|
rotated, err := FromArray(y, false).TransposeAxes(2, 0, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := rotated.Data().Shape(); got[0] != 2 || got[1] != 2 || got[2] != 3 {
|
|
t.Fatalf("rank-3 shape: %v", got)
|
|
}
|
|
|
|
// Invalid permutations error before any graph work.
|
|
if _, err := FromArray(x, false).TransposeAxes(0, 0); err == nil {
|
|
t.Fatal("duplicate axis accepted")
|
|
}
|
|
if _, err := FromArray(x, false).TransposeAxes(0); err == nil {
|
|
t.Fatal("short permutation accepted")
|
|
}
|
|
}
|
|
|
|
// TestTensorTransposeAxesGradient routes a weighted sum through the
|
|
// permutation: the analytic input gradient is exactly the weight tensor
|
|
// played back through the inverse permutation.
|
|
func TestTensorTransposeAxesGradient(t *testing.T) {
|
|
x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
w, _ := core.FromFloats([]float64{0.5, -1, 2, 0.25, -0.75, 1.5}, 3, 2)
|
|
|
|
xt := FromArray(x, true)
|
|
joint, err := xt.TransposeAxes(1, 0) // gives (3, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
scaled, err := joint.Mul(FromArray(w, false))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
loss, err := scaled.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
g := xt.Grad()
|
|
if g == nil || g.Dtype() != core.Float {
|
|
t.Fatalf("gradient missing or wrong dtype: %v", g)
|
|
}
|
|
for i := range 6 {
|
|
row, col := i/3, i%3
|
|
if got := g.FloatAt(i); got != w.FloatAt(col*2+row) {
|
|
t.Errorf("grad[%d] = %v, want %v", i, got, w.FloatAt(col*2+row))
|
|
}
|
|
}
|
|
|
|
// Round-trip: permuting by (1,0) then back restores the values.
|
|
back, err := joint.TransposeAxes(1, 0)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for i := range 6 {
|
|
if back.Data().FloatAt(i) != x.FloatAt(i) {
|
|
t.Fatalf("round-trip[%d] = %v, want %v", i, back.Data().FloatAt(i), x.FloatAt(i))
|
|
}
|
|
}
|
|
}
|