summaryrefslogtreecommitdiff
path: root/tensor/dense_test.go
blob: 5bc4e6fa6d1f8a964c75be26f2788b8a52e844ac (plain)
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")
	}
}