-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_sparse.flow
More file actions
129 lines (120 loc) · 3.2 KB
/
Copy pathtest_sparse.flow
File metadata and controls
129 lines (120 loc) · 3.2 KB
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
124
125
126
127
128
129
import "lib/scikit/scikit.flow"
function main() -> i32 {
println("Test: sparse matrix operations")
# Create a dense matrix with some zeros
let M: Matrix = matrix_new(4, 5)
matrix_set(M, 0, 0, 1.0)
matrix_set(M, 0, 2, 3.0)
matrix_set(M, 1, 1, 5.0)
matrix_set(M, 1, 4, 7.0)
matrix_set(M, 2, 0, 2.0)
matrix_set(M, 3, 3, 9.0)
# Row 0: [1, 0, 3, 0, 0]
# Row 1: [0, 5, 0, 0, 7]
# Row 2: [2, 0, 0, 0, 0]
# Row 3: [0, 0, 0, 9, 0]
# Convert to sparse
let S: SparseMatrix = sparse_from_dense(M)
if S.n_rows != 4 || S.n_cols != 5 {
println(" FAIL: wrong dimensions")
sparse_free(S)
matrix_free(M)
return 1
}
if S.nnz != 6 {
println(" FAIL: wrong nnz")
sparse_free(S)
matrix_free(M)
return 1
}
# Check density
let dens: f32 = sparse_density(S)
# 6 / 20 = 0.3
if dens < 0.29 || dens > 0.31 {
println(" FAIL: wrong density")
sparse_free(S)
matrix_free(M)
return 1
}
# Check sparse_at
if sparse_at(S, 0, 0) != 1.0 {
println(" FAIL: sparse_at(0,0) wrong")
sparse_free(S)
matrix_free(M)
return 1
}
if sparse_at(S, 0, 1) != 0.0 {
println(" FAIL: sparse_at(0,1) wrong")
sparse_free(S)
matrix_free(M)
return 1
}
if sparse_at(S, 3, 3) != 9.0 {
println(" FAIL: sparse_at(3,3) wrong")
sparse_free(S)
matrix_free(M)
return 1
}
# Convert back to dense and verify
let M2: Matrix = sparse_to_dense(S)
let mut i: i32 = 0
while i < 4 {
let mut j: i32 = 0
while j < 5 {
if matrix_at(M, i, j) != matrix_at(M2, i, j) {
println(" FAIL: roundtrip mismatch")
sparse_free(S)
matrix_free(M)
matrix_free(M2)
return 1
}
j = j + 1
}
i = i + 1
}
# Test sparse_dot_row
let v_arr: ptr<f32> = array_new_f32(5)
v_arr[0] = 1.0
v_arr[1] = 2.0
v_arr[2] = 3.0
v_arr[3] = 4.0
v_arr[4] = 5.0
# Row 0: 1*1 + 3*3 = 10
let dot0: f32 = sparse_dot_row(S, 0, v_arr)
if dot0 < 9.9 || dot0 > 10.1 {
println(" FAIL: sparse_dot_row row 0 wrong")
sparse_free(S)
matrix_free(M)
matrix_free(M2)
array_free_f32(v_arr)
return 1
}
# Row 1: 5*2 + 7*5 = 45
let dot1: f32 = sparse_dot_row(S, 1, v_arr)
if dot1 < 44.9 || dot1 > 45.1 {
println(" FAIL: sparse_dot_row row 1 wrong")
sparse_free(S)
matrix_free(M)
matrix_free(M2)
array_free_f32(v_arr)
return 1
}
# Test nnz_per_row
let counts: ptr<i32> = sparse_nnz_per_row(S)
if counts[0] != 2 || counts[1] != 2 || counts[2] != 1 || counts[3] != 1 {
println(" FAIL: nnz_per_row wrong")
free(counts as ptr<void>)
sparse_free(S)
matrix_free(M)
matrix_free(M2)
array_free_f32(v_arr)
return 1
}
println(" OK: all sparse matrix tests passed")
free(counts as ptr<void>)
array_free_f32(v_arr)
sparse_free(S)
matrix_free(M)
matrix_free(M2)
return 0
}