Files
tensor/grad/permute_test.go
T

101 lines
2.5 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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))
}
}
}