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