Files
tensor/internal/core/reshape_pad_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

545 lines
15 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"strings"
"testing"
)
func TestFlatten(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)
got, err := Flatten(a, 0, -1)
if err != nil {
t.Fatal(err)
}
want, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 6)
if got.String() != want.String() {
t.Errorf("flatten all: got %s, want %s", got, want)
}
// Flatten only middle dimension.
got, err = Flatten(a, 1, 1)
if err != nil {
t.Fatal(err)
}
if got.Shape()[0] != 2 || got.Shape()[1] != 3 {
t.Errorf("flatten 1: shape = %v", got.Shape())
}
// Out-of-range.
if _, err := Flatten(a, 0, 5); err == nil {
t.Error("flatten: expected error for out-of-range endDim")
}
}
func TestSqueeze(t *testing.T) {
a, _ := FromFloats([]float64{1, 2, 3, 4}, 1, 2, 1, 2)
got, err := Squeeze(a, 0)
if err != nil {
t.Fatal(err)
}
if len(got.Shape()) != 3 || got.Shape()[0] != 2 {
t.Errorf("squeeze dim 0: shape = %v", got.Shape())
}
// Squeeze all (-1).
got, err = Squeeze(a, -1)
if err != nil {
t.Fatal(err)
}
if len(got.Shape()) != 2 || got.Shape()[0] != 2 {
t.Errorf("squeeze all: shape = %v", got.Shape())
}
// Cannot squeeze non-1 dimension.
if _, err := Squeeze(a, 1); err == nil {
t.Error("squeeze: expected error for non-1 dim")
}
}
func TestUnsqueeze(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3}, 3)
got, err := Unsqueeze(a, 0)
if err != nil {
t.Fatal(err)
}
if len(got.Shape()) != 2 || got.Shape()[0] != 1 || got.Shape()[1] != 3 {
t.Errorf("unsqueeze 0: shape = %v", got.Shape())
}
// Negative dim.
got, err = Unsqueeze(a, -1)
if err != nil {
t.Fatal(err)
}
if got.Shape()[0] != 3 || got.Shape()[1] != 1 {
t.Errorf("unsqueeze -1: shape = %v", got.Shape())
}
}
func TestTransposeAxes(t *testing.T) {
a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
got, err := TransposeAxes(a, 1, 0)
if err != nil {
t.Fatal(err)
}
if got.Shape()[0] != 3 || got.Shape()[1] != 2 {
t.Errorf("transpose: shape = %v", got.Shape())
}
// Check element: a[1,0] = 4 should land at out[0,1].
v, _ := FloatAt(got, 0, 1)
if v != 4 {
t.Errorf("transpose element: got %v, want 4", v)
}
// Invalid permutation.
if _, err := TransposeAxes(a, 0, 0); err == nil {
t.Error("transpose: expected error for duplicate dims")
}
}
func TestCopy(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2)
got := Copy(a)
if got.Shape()[0] != 2 || got.Shape()[1] != 2 {
t.Errorf("copy: shape = %v", got.Shape())
}
}
func TestPadConstant(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2)
// 2-D: pad = (left, right, top, bottom).
got, err := Pad(a, []int{1, 1, 1, 1}, "constant", 0)
if err != nil {
t.Fatal(err)
}
if got.Shape()[0] != 4 || got.Shape()[1] != 4 {
t.Errorf("pad constant: shape = %v", got.Shape())
}
// Corners should be zero.
c, _ := FloatAt(got, 0, 0)
if c != 0 {
t.Errorf("pad corner: got %v, want 0", c)
}
// Centre should be original a[0,0] = 1.
c, _ = FloatAt(got, 1, 1)
if c != 1 {
t.Errorf("pad centre: got %v, want 1", c)
}
// Wrong pad length.
if _, err := Pad(a, []int{1, 1, 1}, "constant", 0); err == nil {
t.Error("pad: expected error for odd-length pad")
}
// Unknown mode.
if _, err := Pad(a, []int{1, 1, 1, 1}, "weird", 0); err == nil {
t.Error("pad: expected error for unknown mode")
}
}
func TestPadReflect(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3}, 1, 3)
got, err := Pad(a, []int{2, 0, 0, 0}, "reflect", 0)
if err != nil {
t.Fatal(err)
}
// Reflect without repeating edge: at index 0 we mirror index 2 -> 3,
// at index 1 we mirror index 1 -> 2.
v, _ := FloatAt(got, 0, 0)
if v != 3 {
t.Errorf("pad reflect [0]: got %v, want 3", v)
}
v, _ = FloatAt(got, 0, 1)
if v != 2 {
t.Errorf("pad reflect [1]: got %v, want 2", v)
}
// Original index 0 = 1 should land at the position equal to the pre-pad.
v, _ = FloatAt(got, 0, 2)
if v != 1 {
t.Errorf("pad reflect [2]: got %v, want 1", v)
}
}
func TestPadReplicate(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2)
got, err := Pad(a, []int{2, 1, 0, 0}, "replicate", 0)
if err != nil {
t.Fatal(err)
}
// Shape (2, 5). At (0, 0) replicate original (0, 0) = 1.
v, _ := FloatAt(got, 0, 0)
if v != 1 {
t.Errorf("pad replicate [0,0]: got %v, want 1", v)
}
v, _ = FloatAt(got, 0, 1)
if v != 1 {
t.Errorf("pad replicate [0,1]: got %v, want 1", v)
}
// Post-pad replicates last column (index 1) for the last 1 column.
v, _ = FloatAt(got, 1, 4)
if v != 4 {
t.Errorf("pad replicate [1,4]: got %v, want 4", v)
}
}
func TestPadCircular(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2)
got, err := Pad(a, []int{2, 1, 0, 0}, "circular", 0)
if err != nil {
t.Fatal(err)
}
// Shape (2, 5). Pad left=2, right=1 on dim 1.
// Row 0 = [1, 2]. Wrapped with pre=2 -> [..., 1, 2, 1, 2, 1] then trim to 5:
// index 0 = wrap(0-2=-2) = 0 -> 1
// index 1 = wrap(0-1=-1) = 1 -> 2
// index 2 = 1
// index 3 = 2
// index 4 = wrap(0+2=2) mod 2 = 0 -> 1
v, _ := FloatAt(got, 0, 0)
if v != 1 {
t.Errorf("pad circular [0,0]: got %v, want 1", v)
}
v, _ = FloatAt(got, 0, 1)
if v != 2 {
t.Errorf("pad circular [0,1]: got %v, want 2", v)
}
v, _ = FloatAt(got, 0, 4)
if v != 1 {
t.Errorf("pad circular [0,4]: got %v, want 1", v)
}
}
func TestGather(t *testing.T) {
a, _ := FromFloats([]float64{10, 20, 30, 40, 50, 60}, 2, 3)
idx := mustFromInts(t, []int64{2, 0}, 2, 1)
got, err := Gather(a, 1, idx)
if err != nil {
t.Fatal(err)
}
// Expect [a[0,2], a[1,0]] = [30, 40].
v0, _ := FloatAt(got, 0, 0)
v1, _ := FloatAt(got, 1, 0)
if v0 != 30 || v1 != 40 {
t.Errorf("gather: got [%v, %v], want [30, 40]", v0, v1)
}
// Out-of-range index.
badIdx := mustFromInts(t, []int64{99, 0}, 2, 1)
if _, err := Gather(a, 1, badIdx); err == nil {
t.Error("gather: expected error for out-of-range index")
}
// Wrong dtype.
badDtype := mustFromFloats(t, []float64{1, 1}, 2, 1)
if _, err := Gather(a, 1, badDtype); err == nil {
t.Error("gather: expected error for non-int index")
}
}
func TestScatter(t *testing.T) {
a, _ := FromFloats([]float64{10, 20, 30, 40, 50, 60}, 2, 3)
idx := mustFromInts(t, []int64{2, 0}, 2, 1)
src := mustFromFloats(t, []float64{300, 400}, 2, 1)
got, err := Scatter(a, 1, idx, src)
if err != nil {
t.Fatal(err)
}
// Expect [10, 20, 300, 400, 50, 60].
for i, want := range []float64{10, 20, 300, 400, 50, 60} {
v, _ := FloatAt(got, i/3, i%3)
if v != want {
t.Errorf("scatter [%d]: got %v, want %v", i, v, want)
}
}
// Index/src shape mismatch.
badIdx := mustFromInts(t, []int64{0, 0, 0}, 3, 1)
if _, err := Scatter(a, 0, badIdx, idx); err == nil {
t.Error("scatter: expected error for shape mismatch")
}
}
func TestNonzero(t *testing.T) {
a, _ := FromFloats([]float64{0, 1, 0, 2, 3, 0}, 2, 3)
got, err := Nonzero(a)
if err != nil {
t.Fatal(err)
}
if len(got) != 2 {
t.Fatalf("nonzero dims: got %d, want 2", len(got))
}
// Flat layout [0, 1, 0, 2, 3, 0] -> row 0 [0,1,0], row 1 [2,3,0].
// Non-zero positions: flat 1 = (0,1), flat 3 = (1,0), flat 4 = (1,1).
want := [][2]int{{0, 1}, {1, 0}, {1, 1}}
for i, w := range want {
if got[0][i] != w[0] || got[1][i] != w[1] {
t.Errorf("nonzero [%d]: got (%d, %d), want %v", i, got[0][i], got[1][i], w)
}
}
// Complex rejected.
c, _ := FromComplexes([]complex128{1, 0}, 2)
if _, err := Nonzero(c); err == nil {
t.Error("nonzero: expected error for complex input")
}
}
func TestTake(t *testing.T) {
a := mustFromFloats(t, []float64{10, 20, 30, 40, 50}, 5)
idx := mustFromInts(t, []int64{4, 0, 2}, 3)
got, err := Take(a, idx)
if err != nil {
t.Fatal(err)
}
for i, want := range []float64{50, 10, 30} {
v, _ := FloatAt(got, i)
if v != want {
t.Errorf("take [%d]: got %v, want %v", i, v, want)
}
}
// Out-of-range.
badIdx := mustFromInts(t, []int64{99}, 1)
if _, err := Take(a, badIdx); err == nil {
t.Error("take: expected error for out-of-range index")
}
// Non-1-D indices.
badShape := mustFromInts(t, []int64{0, 0}, 2, 1)
if _, err := Take(a, badShape); err == nil {
t.Error("take: expected error for non-1-D indices")
}
}
func TestCumSum(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)
got, err := CumSum(a, 1)
if err != nil {
t.Fatal(err)
}
want, _ := FromFloats([]float64{1, 3, 6, 4, 9, 15}, 2, 3)
for i := range 6 {
v, _ := FloatAt(got, i/3, i%3)
w, _ := FloatAt(want, i/3, i%3)
if v != w {
t.Errorf("cumsum [%d]: got %v, want %v", i, v, w)
}
}
// Out-of-range dim.
if _, err := CumSum(a, 5); err == nil || !strings.Contains(err.Error(), "out of range") {
t.Errorf("cumsum: expected out-of-range error, got %v", err)
}
}
func TestCumProd(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)
got, err := CumProd(a, 0)
if err != nil {
t.Fatal(err)
}
want, _ := FromFloats([]float64{1, 2, 3, 4, 10, 18}, 2, 3)
for i := range 6 {
v, _ := FloatAt(got, i/3, i%3)
w, _ := FloatAt(want, i/3, i%3)
if v != w {
t.Errorf("cumprod [%d]: got %v, want %v", i, v, w)
}
}
}
func TestProd(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)
got, err := Prod(a, 1, false)
if err != nil {
t.Fatal(err)
}
// Prod across dim 1: row 0 = [1,2,3] -> 6, row 1 = [4,5,6] -> 120.
for i, w := range []float64{6, 120} {
v, _ := FloatAt(got, i)
if v != w {
t.Errorf("prod [%d]: got %v, want %v", i, v, w)
}
}
// With keepDim.
got, err = Prod(a, 1, true)
if err != nil {
t.Fatal(err)
}
if got.Shape()[0] != 2 || got.Shape()[1] != 1 {
t.Errorf("prod keepDim: shape = %v", got.Shape())
}
}
func TestNorm(t *testing.T) {
a := mustFromFloats(t, []float64{3, 4}, 2)
got, err := Norm(a, 2, 0, false)
if err != nil {
t.Fatal(err)
}
if math.Abs(got.RawFloats()[0]-5) > 1e-9 {
t.Errorf("L2 norm [3,4]: got %v, want 5", got.RawFloats()[0])
}
// L1.
got, _ = Norm(a, 1, 0, false)
if got.RawFloats()[0] != 7 {
t.Errorf("L1 norm: got %v, want 7", got.RawFloats()[0])
}
// L-inf.
got, _ = Norm(a, math.Inf(1), 0, false)
if got.RawFloats()[0] != 4 {
t.Errorf("L-inf norm: got %v, want 4", got.RawFloats()[0])
}
// Negative p.
if _, err := Norm(a, -1, 0, false); err == nil {
t.Error("norm: expected error for negative p")
}
}
func TestTrace(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3)
got, err := Trace(a)
if err != nil {
t.Fatal(err)
}
if got != 15 {
t.Errorf("trace: got %v, want 15", got)
}
// Non-square.
bad, _ := FromFloats([]float64{1, 2, 3, 4}, 2, 2)
// 2x2 is square; try non-square.
bad2, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
if _, err := Trace(bad2); err == nil {
t.Error("trace: expected error for non-square")
}
// Just to silence "unused" for bad.
_ = bad
}
func TestDiagonal(t *testing.T) {
a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3)
main, err := Diagonal(a, 0)
if err != nil {
t.Fatal(err)
}
for i, w := range []float64{1, 5, 9} {
v, _ := FloatAt(main, i)
if v != w {
t.Errorf("diag main [%d]: got %v, want %v", i, v, w)
}
}
// Super-diagonal offset 1.
sup, err := Diagonal(a, 1)
if err != nil {
t.Fatal(err)
}
for i, w := range []float64{2, 6} {
v, _ := FloatAt(sup, i)
if v != w {
t.Errorf("diag +1 [%d]: got %v, want %v", i, v, w)
}
}
// Sub-diagonal offset -1.
sub, err := Diagonal(a, -1)
if err != nil {
t.Fatal(err)
}
for i, w := range []float64{4, 8} {
v, _ := FloatAt(sub, i)
if v != w {
t.Errorf("diag -1 [%d]: got %v, want %v", i, v, w)
}
}
}
func TestKron(t *testing.T) {
a, _ := FromFloats([]float64{1, 2}, 1, 2)
b, _ := FromFloats([]float64{0, 5, 6, 7}, 2, 2)
got, err := Kron(a, b)
if err != nil {
t.Fatal(err)
}
if got.Shape()[0] != 2 || got.Shape()[1] != 4 {
t.Errorf("kron: shape = %v", got.Shape())
}
// Expect: a = [[1, 2]], b = [[0,5],[6,7]]
// Result = [[0,5,0,10],[6,7,12,14]].
for i, w := range []float64{0, 5, 0, 10, 6, 7, 12, 14} {
v, _ := FloatAt(got, i/4, i%4)
if v != w {
t.Errorf("kron [%d]: got %v, want %v", i, v, w)
}
}
// Int input: exercises setFromValue with int dtype.
aInt, _ := FromInts([]int64{1, 2}, 1, 2)
gotInt, err := Kron(aInt, b)
if err != nil {
t.Fatal(err)
}
v0, _ := IntAt(gotInt, 0, 0)
if v0 != 0 {
t.Errorf("kron int [0,0]: got %v, want 0", v0)
}
}
// TestPadReflectRejectsOversizedPads pins the reflect guard: a pad of
// the full axis length has no source to mirror, and the single fold
// used to emit a negative offset into the payload (panic).
func TestPadReflectRejectsOversizedPads(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3}, 3)
if _, err := Pad(a, []int{6, 0}, "reflect", 0); err == nil {
t.Error("Pad reflect pad=6 on length 3: expected error")
}
// A pad of n-1 is still mirrorable and must keep working.
got, err := Pad(a, []int{2, 0}, "reflect", 0)
if err != nil {
t.Fatalf("Pad reflect pad=2: %v", err)
}
want := mustFromFloats(t, []float64{3, 2, 1, 2, 3}, 5)
if !Equal(want, got) {
t.Errorf("Pad reflect pad=2: %s", got)
}
}
// TestScanFloat32CarryNarrows pins the float32 scan's carry rule: the
// running value narrows to float32 at every step, exactly as the value
// the walk stored fed the next one, so a term below the current
// precision is absorbed before the next term lands on it.
func TestScanFloat32CarryNarrows(t *testing.T) {
// 1 + 1e-9 rounds back to 1, so the -1 lands on a clean zero.
sums := mustFromFloat32s(t, []float32{1, 1e-9, -1}, 3)
cs, err := CumSum(sums, 0)
if err != nil {
t.Fatalf("CumSum float32: %v", err)
}
for i, want := range []float32{1, 1, 0} {
if got := cs.RawFloat32s()[i]; got != want {
t.Errorf("CumSum float32 [%d] = %v, want %v", i, got, want)
}
}
// A longer chain keeps the rule: the absorbed terms never
// accumulate into the running sum.
chain := mustFromFloat32s(t, []float32{1, 1e-9, 1e-9, 1e-9, -1}, 5)
cs2, err := CumSum(chain, 0)
if err != nil {
t.Fatalf("CumSum float32 chain: %v", err)
}
for i, want := range []float32{1, 1, 1, 1, 0} {
if got := cs2.RawFloat32s()[i]; got != want {
t.Errorf("CumSum float32 chain [%d] = %v, want %v", i, got, want)
}
}
// The product twin narrows the same way, so a product wider than
// float32 is rounded before the next factor multiplies it.
vals := []float32{1 + 1.0/4096, 1 + 1.0/4096, 1 + 1.0/4096, 1 + 1.0/4096}
prod := mustFromFloat32s(t, vals, len(vals))
cp, err := CumProd(prod, 0)
if err != nil {
t.Fatalf("CumProd float32: %v", err)
}
var want []float32
acc := 1.0
for i, v := range vals {
if i == 0 {
acc = float64(v)
} else {
acc = float64(float32(acc)) * float64(v)
}
want = append(want, float32(acc))
}
for i := range want {
if got := cp.RawFloat32s()[i]; got != want[i] {
t.Errorf("CumProd float32 [%d] = %v, want %v", i, got, want[i])
}
}
}