feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,100 @@
|
||||
// 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))
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user