1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
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")
}
}
|