From 65d687794f7cab5ff57ab2a129f0c83924a1538e Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 25 Feb 2026 13:13:41 +0100 Subject: Add tensor package --- tensor/dense_test.go | 123 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 123 insertions(+) create mode 100644 tensor/dense_test.go (limited to 'tensor/dense_test.go') diff --git a/tensor/dense_test.go b/tensor/dense_test.go new file mode 100644 index 0000000..5bc4e6f --- /dev/null +++ b/tensor/dense_test.go @@ -0,0 +1,123 @@ +package tensor + +import ( + "slices" + "testing" +) + +func TestNewDense_Strides(t *testing.T) { + d := NewDense[float32]([]int{2, 3, 4}) + + expected := []int{12, 4, 1} + + if !slices.Equal(d.strides, expected) { + t.Fatalf("expected %v but got %v", expected, d.strides) + } +} + +func TestDense_IsContiguous(t *testing.T) { + d := NewDense[float32]([]int{2, 3}) + + if !d.IsContiguous() { + t.Fatalf("expected contiguous true") + } + + permuted := d.Permute([]int{1, 0}) + + if permuted.IsContiguous() { + t.Fatalf("expected contiguous false") + } +} + +func TestDense_IsContiguousStrict(t *testing.T) { + d := NewDense[int]([]int{1, 3}) + + v := Dense[int]{ + base: d.base, + offset: 0, + shape: []int{1, 3}, + strides: []int{999, 1}, + } + + if v.IsContiguous() { + t.Fatalf("expected contiguous false") + } +} + +func TestDense_Contiguous(t *testing.T) { + d := NewDense[float32]([]int{2, 3}) + + permuted := d.Permute([]int{1, 0}) + + contiguous := permuted.Contiguous() + + if !contiguous.IsContiguous() { + t.Fatalf("expected contiguous true") + } +} + +func TestDense_Permute(t *testing.T) { + d := NewDense[float32]([]int{2, 3}) + + for i := range d.Size() { + d.base[i] = float32(i) + } + + permuted := d.Permute([]int{1, 0}) + + expected := []float32{ + 0, 3, + 1, 4, + 2, 5, + } + + buffer := make([]int, permuted.Rank()) + + linear := 0 + + for _, v := range permuted.All(buffer) { + a := v + b := expected[linear] + + if a != b { + t.Fatalf("expected %.2f but got %.2f at linear %d", b, a, linear) + } + + linear++ + } +} + +func TestDense_All(t *testing.T) { + d := NewDense[float32]([]int{2, 3}) + + for i := range d.Size() { + d.base[i] = float32(i) + } + + d = d.Permute([]int{1, 0}) + + out := []float32{ + 0, 3, + 1, 4, + 2, 5, + } + + buffer := make([]int, d.Rank()) + + step := 0 + + for idxs := range d.All(buffer) { + a := d.At(idxs) + b := out[step] + + if a != b { + t.Fatalf("expected %.2f but got %.2f", a, b) + } + + step++ + } + + if step != d.Size() { + t.Fatalf("premature termination") + } +} -- cgit v1.3.1