diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-02-25 13:13:41 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-02-25 13:13:41 +0100 |
| commit | 65d687794f7cab5ff57ab2a129f0c83924a1538e (patch) | |
| tree | 4e976e665837a19dff4d2763cb8b39b23e82f35e /tensor/dense_test.go | |
| parent | f35de835c3b0068f5a2e992c4a0af2acf59f7bb9 (diff) | |
Add tensor package
Diffstat (limited to 'tensor/dense_test.go')
| -rw-r--r-- | tensor/dense_test.go | 123 |
1 files changed, 123 insertions, 0 deletions
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") + } +} |
