gqlite 1.5.0 → 1.7.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +4 -4
- data/ext/Cargo.toml +12 -6
- data/ext/db-index/Cargo.toml +29 -0
- data/ext/db-index/src/hnsw/error.rs +66 -0
- data/ext/db-index/src/hnsw/graph_store.rs +272 -0
- data/ext/db-index/src/hnsw/id.rs +36 -0
- data/ext/db-index/src/hnsw/implementation.rs +1517 -0
- data/ext/db-index/src/hnsw/kernels.rs +1228 -0
- data/ext/db-index/src/hnsw/metric.rs +244 -0
- data/ext/db-index/src/hnsw/scalar.rs +72 -0
- data/ext/db-index/src/hnsw/simple_store.rs +140 -0
- data/ext/db-index/src/hnsw/store.rs +472 -0
- data/ext/db-index/src/hnsw/vector.rs +105 -0
- data/ext/db-index/src/hnsw/vectors.rs +568 -0
- data/ext/db-index/src/hnsw.rs +42 -0
- data/ext/db-index/src/lib.rs +3 -0
- data/ext/gqlitedb/Cargo.toml +19 -8
- data/ext/gqlitedb/benches/common/pokec.rs +62 -2
- data/ext/gqlitedb/benches/pokec_divan.rs +66 -2
- data/ext/gqlitedb/benches/pokec_iai.rs +60 -3
- data/ext/gqlitedb/release.toml +2 -2
- data/ext/gqlitedb/src/aggregators/arithmetic.rs +3 -1
- data/ext/gqlitedb/src/aggregators/containers.rs +1 -1
- data/ext/gqlitedb/src/aggregators/stats.rs +21 -12
- data/ext/gqlitedb/src/capi.rs +20 -24
- data/ext/gqlitedb/src/compiler/expression_analyser.rs +28 -20
- data/ext/gqlitedb/src/compiler/variables_manager.rs +104 -38
- data/ext/gqlitedb/src/compiler.rs +506 -227
- data/ext/gqlitedb/src/connection.rs +151 -11
- data/ext/gqlitedb/src/consts.rs +1 -2
- data/ext/gqlitedb/src/error.rs +123 -61
- data/ext/gqlitedb/src/functions/common.rs +64 -0
- data/ext/gqlitedb/src/functions/containers.rs +2 -46
- data/ext/gqlitedb/src/functions/edge.rs +1 -1
- data/ext/gqlitedb/src/functions/math.rs +87 -15
- data/ext/gqlitedb/src/functions/node.rs +1 -1
- data/ext/gqlitedb/src/functions/path.rs +48 -11
- data/ext/gqlitedb/src/functions/scalar.rs +47 -4
- data/ext/gqlitedb/src/functions/string.rs +123 -1
- data/ext/gqlitedb/src/functions/value.rs +1 -1
- data/ext/gqlitedb/src/functions.rs +165 -28
- data/ext/gqlitedb/src/graph.rs +1 -9
- data/ext/gqlitedb/src/interpreter/evaluators.rs +968 -130
- data/ext/gqlitedb/src/interpreter/instructions.rs +39 -2
- data/ext/gqlitedb/src/lib.rs +5 -4
- data/ext/gqlitedb/src/parser/gql.pest +1 -1
- data/ext/gqlitedb/src/planner.rs +329 -0
- data/ext/gqlitedb/src/prelude.rs +3 -3
- data/ext/gqlitedb/src/store/pgrx.rs +1 -1
- data/ext/gqlitedb/src/store/postgres.rs +742 -16
- data/ext/gqlitedb/src/store/redb/hnsw_store.rs +702 -0
- data/ext/gqlitedb/src/store/redb/index.rs +274 -0
- data/ext/gqlitedb/src/store/redb.rs +1268 -113
- data/ext/gqlitedb/src/store/sqlbase/sqlmetadata.rs +28 -0
- data/ext/gqlitedb/src/store/sqlbase/sqlstore.rs +103 -0
- data/ext/gqlitedb/src/store/sqlbase.rs +146 -16
- data/ext/gqlitedb/src/store/sqlite.rs +576 -14
- data/ext/gqlitedb/src/store/vector_extract.rs +56 -0
- data/ext/gqlitedb/src/store.rs +140 -29
- data/ext/gqlitedb/src/tests/compiler.rs +207 -10
- data/ext/gqlitedb/src/tests/connection/postgres.rs +38 -3
- data/ext/gqlitedb/src/tests/connection/redb.rs +23 -0
- data/ext/gqlitedb/src/tests/connection/sqlite.rs +31 -0
- data/ext/gqlitedb/src/tests/connection.rs +61 -1
- data/ext/gqlitedb/src/tests/evaluators.rs +162 -7
- data/ext/gqlitedb/src/tests/parser.rs +54 -23
- data/ext/gqlitedb/src/tests/planner.rs +511 -0
- data/ext/gqlitedb/src/tests/store/postgres.rs +7 -0
- data/ext/gqlitedb/src/tests/store/redb.rs +8 -0
- data/ext/gqlitedb/src/tests/store/sqlite.rs +8 -0
- data/ext/gqlitedb/src/tests/store/vector_index/postgres.rs +182 -0
- data/ext/gqlitedb/src/tests/store/vector_index/redb.rs +386 -0
- data/ext/gqlitedb/src/tests/store/vector_index/sqlite.rs +166 -0
- data/ext/gqlitedb/src/tests/store/vector_index.rs +2313 -0
- data/ext/gqlitedb/src/tests/store.rs +78 -14
- data/ext/gqlitedb/src/tests/templates/ast.rs +92 -16
- data/ext/gqlitedb/src/tests/templates/programs.rs +14 -7
- data/ext/gqlitedb/src/tests.rs +15 -9
- data/ext/gqlitedb/src/utils.rs +6 -1
- data/ext/gqlitedb/src/value/compare.rs +61 -3
- data/ext/gqlitedb/src/value.rs +136 -7
- data/ext/gqlitedb/templates/sql/postgres/metadata_delete.sql +1 -0
- data/ext/gqlitedb/templates/sql/postgres/node_select.sql +1 -1
- data/ext/gqlitedb/templates/sql/sqlite/metadata_delete.sql +1 -0
- data/ext/gqliterb/src/lib.rs +53 -8
- data/ext/gqlparser/Cargo.toml +25 -0
- data/ext/gqlparser/README.MD +9 -0
- data/ext/gqlparser/benches/pokec_divan.rs +34 -0
- data/ext/gqlparser/src/common.rs +69 -0
- data/ext/gqlparser/src/gqls/ast.rs +69 -0
- data/ext/gqlparser/src/gqls/constraint.rs +79 -0
- data/ext/gqlparser/src/gqls/error.rs +22 -0
- data/ext/gqlparser/src/gqls/parser.rs +813 -0
- data/ext/gqlparser/src/gqls/prelude.rs +8 -0
- data/ext/gqlparser/src/gqls/properties.rs +115 -0
- data/ext/gqlparser/src/gqls/resolve.rs +207 -0
- data/ext/gqlparser/src/gqls.rs +265 -0
- data/ext/gqlparser/src/lib.rs +7 -0
- data/ext/gqlparser/src/oc/ast.rs +680 -0
- data/ext/gqlparser/src/oc/error.rs +172 -0
- data/ext/gqlparser/src/oc/lexer.rs +429 -0
- data/ext/gqlparser/src/oc/parser/tests.rs +2284 -0
- data/ext/gqlparser/src/oc/parser.rs +2005 -0
- data/ext/gqlparser/src/oc.rs +33 -0
- data/ext/gqlparser/src/prelude.rs +3 -0
- data/ext/graphcore/Cargo.toml +1 -0
- data/ext/graphcore/src/error.rs +26 -0
- data/ext/graphcore/src/graph.rs +177 -51
- data/ext/graphcore/src/lib.rs +4 -2
- data/ext/graphcore/src/open_cypher.rs +12 -0
- data/ext/graphcore/src/table.rs +50 -2
- data/ext/graphcore/src/timestamp.rs +127 -104
- data/ext/graphcore/src/value/tensor.rs +739 -0
- data/ext/graphcore/src/value/value_map.rs +1 -1
- data/ext/graphcore/src/value.rs +343 -19
- metadata +93 -28
- data/ext/gqlitedb/src/parser/ast.rs +0 -604
- data/ext/gqlitedb/src/parser/parser_impl.rs +0 -1213
- data/ext/gqlitedb/src/parser.rs +0 -4
|
@@ -0,0 +1,1228 @@
|
|
|
1
|
+
use half::bf16;
|
|
2
|
+
use half::f16;
|
|
3
|
+
|
|
4
|
+
pub trait Kernel: Copy + Send + Sync + 'static
|
|
5
|
+
{
|
|
6
|
+
fn dot(a: &[Self], b: &[Self]) -> f32;
|
|
7
|
+
fn l2_sq(a: &[Self], b: &[Self]) -> f32;
|
|
8
|
+
fn dot_and_norms(a: &[Self], b: &[Self]) -> (f32, f32, f32);
|
|
9
|
+
}
|
|
10
|
+
|
|
11
|
+
#[derive(Clone, Copy, Debug)]
|
|
12
|
+
struct I8Sums
|
|
13
|
+
{
|
|
14
|
+
dot: i32,
|
|
15
|
+
sum_a: i32,
|
|
16
|
+
sum_b: i32,
|
|
17
|
+
sum_sq_a: i32,
|
|
18
|
+
sum_sq_b: i32,
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
pub fn cosine_distance<T: Kernel>(a: &[T], b: &[T]) -> f32
|
|
22
|
+
{
|
|
23
|
+
let (dot, norm_sq_a, norm_sq_b) = T::dot_and_norms(a, b);
|
|
24
|
+
if norm_sq_a == 0.0 || norm_sq_b == 0.0
|
|
25
|
+
{
|
|
26
|
+
return 1.0;
|
|
27
|
+
}
|
|
28
|
+
let denom = (norm_sq_a * norm_sq_b).sqrt();
|
|
29
|
+
if denom == 0.0
|
|
30
|
+
{
|
|
31
|
+
1.0
|
|
32
|
+
}
|
|
33
|
+
else
|
|
34
|
+
{
|
|
35
|
+
1.0 - dot / denom
|
|
36
|
+
}
|
|
37
|
+
}
|
|
38
|
+
|
|
39
|
+
pub fn inner_product_distance<T: Kernel>(a: &[T], b: &[T]) -> f32
|
|
40
|
+
{
|
|
41
|
+
1.0 - T::dot(a, b)
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
fn i8_sums_scalar(a: &[i8], b: &[i8]) -> I8Sums
|
|
45
|
+
{
|
|
46
|
+
debug_assert_eq!(a.len(), b.len());
|
|
47
|
+
|
|
48
|
+
let mut dot = 0i32;
|
|
49
|
+
let mut sum_a = 0i32;
|
|
50
|
+
let mut sum_b = 0i32;
|
|
51
|
+
let mut sum_sq_a = 0i32;
|
|
52
|
+
let mut sum_sq_b = 0i32;
|
|
53
|
+
|
|
54
|
+
for (&a, &b) in a.iter().zip(b)
|
|
55
|
+
{
|
|
56
|
+
let a = a as i32;
|
|
57
|
+
let b = b as i32;
|
|
58
|
+
dot += a * b;
|
|
59
|
+
sum_a += a;
|
|
60
|
+
sum_b += b;
|
|
61
|
+
sum_sq_a += a * a;
|
|
62
|
+
sum_sq_b += b * b;
|
|
63
|
+
}
|
|
64
|
+
|
|
65
|
+
I8Sums {
|
|
66
|
+
dot,
|
|
67
|
+
sum_a,
|
|
68
|
+
sum_b,
|
|
69
|
+
sum_sq_a,
|
|
70
|
+
sum_sq_b,
|
|
71
|
+
}
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
fn i8_sums(a: &[i8], b: &[i8]) -> I8Sums
|
|
75
|
+
{
|
|
76
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
77
|
+
{
|
|
78
|
+
if std::is_x86_feature_detected!("avx2")
|
|
79
|
+
{
|
|
80
|
+
// SAFETY: guarded by runtime feature check.
|
|
81
|
+
unsafe { return x86::i8_sums_avx2(a, b) };
|
|
82
|
+
}
|
|
83
|
+
}
|
|
84
|
+
#[cfg(target_arch = "aarch64")]
|
|
85
|
+
{
|
|
86
|
+
// AArch64 always has Neon.
|
|
87
|
+
// SAFETY: intrinsic usage.
|
|
88
|
+
unsafe { return aarch64::i8_sums_neon(a, b) };
|
|
89
|
+
}
|
|
90
|
+
#[allow(unreachable_code)]
|
|
91
|
+
i8_sums_scalar(a, b)
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
pub fn dot_and_norms_qi8(
|
|
95
|
+
a: &[i8],
|
|
96
|
+
a_scale: f32,
|
|
97
|
+
a_zero_point: i8,
|
|
98
|
+
b: &[i8],
|
|
99
|
+
b_scale: f32,
|
|
100
|
+
b_zero_point: i8,
|
|
101
|
+
) -> (f32, f32, f32)
|
|
102
|
+
{
|
|
103
|
+
assert_eq!(a.len(), b.len());
|
|
104
|
+
|
|
105
|
+
let sums = i8_sums(a, b);
|
|
106
|
+
let d = a.len() as i32;
|
|
107
|
+
let za = a_zero_point as i32;
|
|
108
|
+
let zb = b_zero_point as i32;
|
|
109
|
+
|
|
110
|
+
let ab = sums.dot - zb * sums.sum_a - za * sums.sum_b + d * za * zb;
|
|
111
|
+
let a2 = sums.sum_sq_a - 2 * za * sums.sum_a + d * za * za;
|
|
112
|
+
let b2 = sums.sum_sq_b - 2 * zb * sums.sum_b + d * zb * zb;
|
|
113
|
+
|
|
114
|
+
let dot = (a_scale * b_scale) * (ab as f32);
|
|
115
|
+
let norm_sq_a = (a_scale * a_scale) * (a2 as f32);
|
|
116
|
+
let norm_sq_b = (b_scale * b_scale) * (b2 as f32);
|
|
117
|
+
(dot, norm_sq_a, norm_sq_b)
|
|
118
|
+
}
|
|
119
|
+
|
|
120
|
+
pub fn l2_sq_qi8(
|
|
121
|
+
a: &[i8],
|
|
122
|
+
a_scale: f32,
|
|
123
|
+
a_zero_point: i8,
|
|
124
|
+
b: &[i8],
|
|
125
|
+
b_scale: f32,
|
|
126
|
+
b_zero_point: i8,
|
|
127
|
+
) -> f32
|
|
128
|
+
{
|
|
129
|
+
let (dot, norm_sq_a, norm_sq_b) =
|
|
130
|
+
dot_and_norms_qi8(a, a_scale, a_zero_point, b, b_scale, b_zero_point);
|
|
131
|
+
norm_sq_a + norm_sq_b - 2.0 * dot
|
|
132
|
+
}
|
|
133
|
+
|
|
134
|
+
pub fn inner_product_distance_qi8(
|
|
135
|
+
a: &[i8],
|
|
136
|
+
a_scale: f32,
|
|
137
|
+
a_zero_point: i8,
|
|
138
|
+
b: &[i8],
|
|
139
|
+
b_scale: f32,
|
|
140
|
+
b_zero_point: i8,
|
|
141
|
+
) -> f32
|
|
142
|
+
{
|
|
143
|
+
let (dot, _, _) = dot_and_norms_qi8(a, a_scale, a_zero_point, b, b_scale, b_zero_point);
|
|
144
|
+
1.0 - dot
|
|
145
|
+
}
|
|
146
|
+
|
|
147
|
+
pub fn cosine_distance_qi8(
|
|
148
|
+
a: &[i8],
|
|
149
|
+
a_scale: f32,
|
|
150
|
+
a_zero_point: i8,
|
|
151
|
+
b: &[i8],
|
|
152
|
+
b_scale: f32,
|
|
153
|
+
b_zero_point: i8,
|
|
154
|
+
) -> f32
|
|
155
|
+
{
|
|
156
|
+
let (dot, norm_sq_a, norm_sq_b) =
|
|
157
|
+
dot_and_norms_qi8(a, a_scale, a_zero_point, b, b_scale, b_zero_point);
|
|
158
|
+
if norm_sq_a == 0.0 || norm_sq_b == 0.0
|
|
159
|
+
{
|
|
160
|
+
return 1.0;
|
|
161
|
+
}
|
|
162
|
+
let denom = (norm_sq_a * norm_sq_b).sqrt();
|
|
163
|
+
if denom == 0.0
|
|
164
|
+
{
|
|
165
|
+
1.0
|
|
166
|
+
}
|
|
167
|
+
else
|
|
168
|
+
{
|
|
169
|
+
1.0 - dot / denom
|
|
170
|
+
}
|
|
171
|
+
}
|
|
172
|
+
|
|
173
|
+
impl Kernel for f32
|
|
174
|
+
{
|
|
175
|
+
fn dot(a: &[Self], b: &[Self]) -> f32
|
|
176
|
+
{
|
|
177
|
+
assert_eq!(a.len(), b.len());
|
|
178
|
+
dot_f32(a, b)
|
|
179
|
+
}
|
|
180
|
+
|
|
181
|
+
fn l2_sq(a: &[Self], b: &[Self]) -> f32
|
|
182
|
+
{
|
|
183
|
+
assert_eq!(a.len(), b.len());
|
|
184
|
+
l2_sq_f32(a, b)
|
|
185
|
+
}
|
|
186
|
+
|
|
187
|
+
fn dot_and_norms(a: &[Self], b: &[Self]) -> (f32, f32, f32)
|
|
188
|
+
{
|
|
189
|
+
assert_eq!(a.len(), b.len());
|
|
190
|
+
dot_and_norms_f32(a, b)
|
|
191
|
+
}
|
|
192
|
+
}
|
|
193
|
+
|
|
194
|
+
impl Kernel for bf16
|
|
195
|
+
{
|
|
196
|
+
fn dot(a: &[Self], b: &[Self]) -> f32
|
|
197
|
+
{
|
|
198
|
+
assert_eq!(a.len(), b.len());
|
|
199
|
+
dot_bf16(a, b)
|
|
200
|
+
}
|
|
201
|
+
|
|
202
|
+
fn l2_sq(a: &[Self], b: &[Self]) -> f32
|
|
203
|
+
{
|
|
204
|
+
assert_eq!(a.len(), b.len());
|
|
205
|
+
l2_sq_bf16(a, b)
|
|
206
|
+
}
|
|
207
|
+
|
|
208
|
+
fn dot_and_norms(a: &[Self], b: &[Self]) -> (f32, f32, f32)
|
|
209
|
+
{
|
|
210
|
+
assert_eq!(a.len(), b.len());
|
|
211
|
+
dot_and_norms_bf16(a, b)
|
|
212
|
+
}
|
|
213
|
+
}
|
|
214
|
+
|
|
215
|
+
impl Kernel for f16
|
|
216
|
+
{
|
|
217
|
+
fn dot(a: &[Self], b: &[Self]) -> f32
|
|
218
|
+
{
|
|
219
|
+
assert_eq!(a.len(), b.len());
|
|
220
|
+
dot_f16(a, b)
|
|
221
|
+
}
|
|
222
|
+
|
|
223
|
+
fn l2_sq(a: &[Self], b: &[Self]) -> f32
|
|
224
|
+
{
|
|
225
|
+
assert_eq!(a.len(), b.len());
|
|
226
|
+
l2_sq_f16(a, b)
|
|
227
|
+
}
|
|
228
|
+
|
|
229
|
+
fn dot_and_norms(a: &[Self], b: &[Self]) -> (f32, f32, f32)
|
|
230
|
+
{
|
|
231
|
+
assert_eq!(a.len(), b.len());
|
|
232
|
+
dot_and_norms_f16(a, b)
|
|
233
|
+
}
|
|
234
|
+
}
|
|
235
|
+
|
|
236
|
+
fn dot_f32(a: &[f32], b: &[f32]) -> f32
|
|
237
|
+
{
|
|
238
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
239
|
+
{
|
|
240
|
+
if std::is_x86_feature_detected!("avx512f")
|
|
241
|
+
{
|
|
242
|
+
// SAFETY: guarded by runtime feature check.
|
|
243
|
+
unsafe { return x86::dot_f32_avx512(a, b) };
|
|
244
|
+
}
|
|
245
|
+
if std::is_x86_feature_detected!("avx")
|
|
246
|
+
{
|
|
247
|
+
unsafe { return x86::dot_f32_avx(a, b) };
|
|
248
|
+
}
|
|
249
|
+
if std::is_x86_feature_detected!("sse")
|
|
250
|
+
{
|
|
251
|
+
unsafe { return x86::dot_f32_sse(a, b) };
|
|
252
|
+
}
|
|
253
|
+
}
|
|
254
|
+
dot_f32_scalar(a, b)
|
|
255
|
+
}
|
|
256
|
+
|
|
257
|
+
fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32
|
|
258
|
+
{
|
|
259
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
260
|
+
{
|
|
261
|
+
if std::is_x86_feature_detected!("avx512f")
|
|
262
|
+
{
|
|
263
|
+
unsafe { return x86::l2_sq_f32_avx512(a, b) };
|
|
264
|
+
}
|
|
265
|
+
if std::is_x86_feature_detected!("avx")
|
|
266
|
+
{
|
|
267
|
+
unsafe { return x86::l2_sq_f32_avx(a, b) };
|
|
268
|
+
}
|
|
269
|
+
if std::is_x86_feature_detected!("sse")
|
|
270
|
+
{
|
|
271
|
+
unsafe { return x86::l2_sq_f32_sse(a, b) };
|
|
272
|
+
}
|
|
273
|
+
}
|
|
274
|
+
l2_sq_f32_scalar(a, b)
|
|
275
|
+
}
|
|
276
|
+
|
|
277
|
+
fn dot_and_norms_f32(a: &[f32], b: &[f32]) -> (f32, f32, f32)
|
|
278
|
+
{
|
|
279
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
280
|
+
{
|
|
281
|
+
if std::is_x86_feature_detected!("avx512f")
|
|
282
|
+
{
|
|
283
|
+
unsafe { return x86::dot_and_norms_f32_avx512(a, b) };
|
|
284
|
+
}
|
|
285
|
+
if std::is_x86_feature_detected!("avx")
|
|
286
|
+
{
|
|
287
|
+
unsafe { return x86::dot_and_norms_f32_avx(a, b) };
|
|
288
|
+
}
|
|
289
|
+
if std::is_x86_feature_detected!("sse")
|
|
290
|
+
{
|
|
291
|
+
unsafe { return x86::dot_and_norms_f32_sse(a, b) };
|
|
292
|
+
}
|
|
293
|
+
}
|
|
294
|
+
dot_and_norms_f32_scalar(a, b)
|
|
295
|
+
}
|
|
296
|
+
|
|
297
|
+
fn dot_bf16(a: &[bf16], b: &[bf16]) -> f32
|
|
298
|
+
{
|
|
299
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
300
|
+
{
|
|
301
|
+
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("avx512bf16")
|
|
302
|
+
{
|
|
303
|
+
unsafe { return x86::dot_bf16_avx512bf16(a, b) };
|
|
304
|
+
}
|
|
305
|
+
}
|
|
306
|
+
dot_bf16_scalar(a, b)
|
|
307
|
+
}
|
|
308
|
+
|
|
309
|
+
fn l2_sq_bf16(a: &[bf16], b: &[bf16]) -> f32
|
|
310
|
+
{
|
|
311
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
312
|
+
{
|
|
313
|
+
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("avx512bf16")
|
|
314
|
+
{
|
|
315
|
+
unsafe { return x86::l2_sq_bf16_avx512bf16(a, b) };
|
|
316
|
+
}
|
|
317
|
+
}
|
|
318
|
+
l2_sq_bf16_scalar(a, b)
|
|
319
|
+
}
|
|
320
|
+
|
|
321
|
+
fn dot_and_norms_bf16(a: &[bf16], b: &[bf16]) -> (f32, f32, f32)
|
|
322
|
+
{
|
|
323
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
324
|
+
{
|
|
325
|
+
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("avx512bf16")
|
|
326
|
+
{
|
|
327
|
+
unsafe { return x86::dot_and_norms_bf16_avx512bf16(a, b) };
|
|
328
|
+
}
|
|
329
|
+
}
|
|
330
|
+
dot_and_norms_bf16_scalar(a, b)
|
|
331
|
+
}
|
|
332
|
+
|
|
333
|
+
fn dot_f16(a: &[f16], b: &[f16]) -> f32
|
|
334
|
+
{
|
|
335
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
336
|
+
{
|
|
337
|
+
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("f16c")
|
|
338
|
+
{
|
|
339
|
+
unsafe { return x86::dot_f16_avx512f_f16c(a, b) };
|
|
340
|
+
}
|
|
341
|
+
if std::is_x86_feature_detected!("avx") && std::is_x86_feature_detected!("f16c")
|
|
342
|
+
{
|
|
343
|
+
unsafe { return x86::dot_f16_avx_f16c(a, b) };
|
|
344
|
+
}
|
|
345
|
+
}
|
|
346
|
+
dot_f16_scalar(a, b)
|
|
347
|
+
}
|
|
348
|
+
|
|
349
|
+
fn l2_sq_f16(a: &[f16], b: &[f16]) -> f32
|
|
350
|
+
{
|
|
351
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
352
|
+
{
|
|
353
|
+
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("f16c")
|
|
354
|
+
{
|
|
355
|
+
unsafe { return x86::l2_sq_f16_avx512f_f16c(a, b) };
|
|
356
|
+
}
|
|
357
|
+
if std::is_x86_feature_detected!("avx") && std::is_x86_feature_detected!("f16c")
|
|
358
|
+
{
|
|
359
|
+
unsafe { return x86::l2_sq_f16_avx_f16c(a, b) };
|
|
360
|
+
}
|
|
361
|
+
}
|
|
362
|
+
l2_sq_f16_scalar(a, b)
|
|
363
|
+
}
|
|
364
|
+
|
|
365
|
+
fn dot_and_norms_f16(a: &[f16], b: &[f16]) -> (f32, f32, f32)
|
|
366
|
+
{
|
|
367
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
368
|
+
{
|
|
369
|
+
if std::is_x86_feature_detected!("avx512f") && std::is_x86_feature_detected!("f16c")
|
|
370
|
+
{
|
|
371
|
+
unsafe { return x86::dot_and_norms_f16_avx512f_f16c(a, b) };
|
|
372
|
+
}
|
|
373
|
+
if std::is_x86_feature_detected!("avx") && std::is_x86_feature_detected!("f16c")
|
|
374
|
+
{
|
|
375
|
+
unsafe { return x86::dot_and_norms_f16_avx_f16c(a, b) };
|
|
376
|
+
}
|
|
377
|
+
}
|
|
378
|
+
dot_and_norms_f16_scalar(a, b)
|
|
379
|
+
}
|
|
380
|
+
|
|
381
|
+
fn dot_f32_scalar(a: &[f32], b: &[f32]) -> f32
|
|
382
|
+
{
|
|
383
|
+
a.iter().zip(b).map(|(a, b)| a * b).sum()
|
|
384
|
+
}
|
|
385
|
+
|
|
386
|
+
fn l2_sq_f32_scalar(a: &[f32], b: &[f32]) -> f32
|
|
387
|
+
{
|
|
388
|
+
a.iter()
|
|
389
|
+
.zip(b)
|
|
390
|
+
.map(|(a, b)| {
|
|
391
|
+
let d = a - b;
|
|
392
|
+
d * d
|
|
393
|
+
})
|
|
394
|
+
.sum()
|
|
395
|
+
}
|
|
396
|
+
|
|
397
|
+
fn dot_and_norms_f32_scalar(a: &[f32], b: &[f32]) -> (f32, f32, f32)
|
|
398
|
+
{
|
|
399
|
+
let mut dot = 0.0f32;
|
|
400
|
+
let mut norm_a = 0.0f32;
|
|
401
|
+
let mut norm_b = 0.0f32;
|
|
402
|
+
for (a, b) in a.iter().zip(b)
|
|
403
|
+
{
|
|
404
|
+
dot += a * b;
|
|
405
|
+
norm_a += a * a;
|
|
406
|
+
norm_b += b * b;
|
|
407
|
+
}
|
|
408
|
+
(dot, norm_a, norm_b)
|
|
409
|
+
}
|
|
410
|
+
|
|
411
|
+
fn dot_bf16_scalar(a: &[bf16], b: &[bf16]) -> f32
|
|
412
|
+
{
|
|
413
|
+
let mut dot = 0.0f32;
|
|
414
|
+
for (a, b) in a.iter().zip(b)
|
|
415
|
+
{
|
|
416
|
+
dot += a.to_f32() * b.to_f32();
|
|
417
|
+
}
|
|
418
|
+
dot
|
|
419
|
+
}
|
|
420
|
+
|
|
421
|
+
fn l2_sq_bf16_scalar(a: &[bf16], b: &[bf16]) -> f32
|
|
422
|
+
{
|
|
423
|
+
let mut sum = 0.0f32;
|
|
424
|
+
for (a, b) in a.iter().zip(b)
|
|
425
|
+
{
|
|
426
|
+
let d = a.to_f32() - b.to_f32();
|
|
427
|
+
sum += d * d;
|
|
428
|
+
}
|
|
429
|
+
sum
|
|
430
|
+
}
|
|
431
|
+
|
|
432
|
+
fn dot_and_norms_bf16_scalar(a: &[bf16], b: &[bf16]) -> (f32, f32, f32)
|
|
433
|
+
{
|
|
434
|
+
let mut dot = 0.0f32;
|
|
435
|
+
let mut norm_a = 0.0f32;
|
|
436
|
+
let mut norm_b = 0.0f32;
|
|
437
|
+
for (a, b) in a.iter().zip(b)
|
|
438
|
+
{
|
|
439
|
+
let a = a.to_f32();
|
|
440
|
+
let b = b.to_f32();
|
|
441
|
+
dot += a * b;
|
|
442
|
+
norm_a += a * a;
|
|
443
|
+
norm_b += b * b;
|
|
444
|
+
}
|
|
445
|
+
(dot, norm_a, norm_b)
|
|
446
|
+
}
|
|
447
|
+
|
|
448
|
+
fn dot_f16_scalar(a: &[f16], b: &[f16]) -> f32
|
|
449
|
+
{
|
|
450
|
+
let mut dot = 0.0f32;
|
|
451
|
+
for (a, b) in a.iter().zip(b)
|
|
452
|
+
{
|
|
453
|
+
dot += a.to_f32() * b.to_f32();
|
|
454
|
+
}
|
|
455
|
+
dot
|
|
456
|
+
}
|
|
457
|
+
|
|
458
|
+
fn l2_sq_f16_scalar(a: &[f16], b: &[f16]) -> f32
|
|
459
|
+
{
|
|
460
|
+
let mut sum = 0.0f32;
|
|
461
|
+
for (a, b) in a.iter().zip(b)
|
|
462
|
+
{
|
|
463
|
+
let d = a.to_f32() - b.to_f32();
|
|
464
|
+
sum += d * d;
|
|
465
|
+
}
|
|
466
|
+
sum
|
|
467
|
+
}
|
|
468
|
+
|
|
469
|
+
fn dot_and_norms_f16_scalar(a: &[f16], b: &[f16]) -> (f32, f32, f32)
|
|
470
|
+
{
|
|
471
|
+
let mut dot = 0.0f32;
|
|
472
|
+
let mut norm_a = 0.0f32;
|
|
473
|
+
let mut norm_b = 0.0f32;
|
|
474
|
+
for (a, b) in a.iter().zip(b)
|
|
475
|
+
{
|
|
476
|
+
let a = a.to_f32();
|
|
477
|
+
let b = b.to_f32();
|
|
478
|
+
dot += a * b;
|
|
479
|
+
norm_a += a * a;
|
|
480
|
+
norm_b += b * b;
|
|
481
|
+
}
|
|
482
|
+
(dot, norm_a, norm_b)
|
|
483
|
+
}
|
|
484
|
+
|
|
485
|
+
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
|
|
486
|
+
mod x86
|
|
487
|
+
{
|
|
488
|
+
use super::I8Sums;
|
|
489
|
+
use half::bf16;
|
|
490
|
+
use half::f16;
|
|
491
|
+
#[cfg(target_arch = "x86")]
|
|
492
|
+
use std::arch::x86::*;
|
|
493
|
+
#[cfg(target_arch = "x86_64")]
|
|
494
|
+
use std::arch::x86_64::*;
|
|
495
|
+
|
|
496
|
+
#[inline]
|
|
497
|
+
pub unsafe fn dot_f32_sse(a: &[f32], b: &[f32]) -> f32
|
|
498
|
+
{
|
|
499
|
+
let mut sum = _mm_setzero_ps();
|
|
500
|
+
let mut i = 0usize;
|
|
501
|
+
while i + 4 <= a.len()
|
|
502
|
+
{
|
|
503
|
+
let va = _mm_loadu_ps(a.as_ptr().add(i));
|
|
504
|
+
let vb = _mm_loadu_ps(b.as_ptr().add(i));
|
|
505
|
+
sum = _mm_add_ps(sum, _mm_mul_ps(va, vb));
|
|
506
|
+
i += 4;
|
|
507
|
+
}
|
|
508
|
+
let mut tmp = [0.0f32; 4];
|
|
509
|
+
_mm_storeu_ps(tmp.as_mut_ptr(), sum);
|
|
510
|
+
let mut dot = tmp.iter().sum::<f32>();
|
|
511
|
+
while i < a.len()
|
|
512
|
+
{
|
|
513
|
+
dot += *a.get_unchecked(i) * *b.get_unchecked(i);
|
|
514
|
+
i += 1;
|
|
515
|
+
}
|
|
516
|
+
dot
|
|
517
|
+
}
|
|
518
|
+
|
|
519
|
+
#[inline]
|
|
520
|
+
pub unsafe fn dot_f32_avx(a: &[f32], b: &[f32]) -> f32
|
|
521
|
+
{
|
|
522
|
+
let mut sum = _mm256_setzero_ps();
|
|
523
|
+
let mut i = 0usize;
|
|
524
|
+
while i + 8 <= a.len()
|
|
525
|
+
{
|
|
526
|
+
let va = _mm256_loadu_ps(a.as_ptr().add(i));
|
|
527
|
+
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
|
|
528
|
+
sum = _mm256_add_ps(sum, _mm256_mul_ps(va, vb));
|
|
529
|
+
i += 8;
|
|
530
|
+
}
|
|
531
|
+
let mut tmp = [0.0f32; 8];
|
|
532
|
+
_mm256_storeu_ps(tmp.as_mut_ptr(), sum);
|
|
533
|
+
let mut dot = tmp.iter().sum::<f32>();
|
|
534
|
+
while i < a.len()
|
|
535
|
+
{
|
|
536
|
+
dot += *a.get_unchecked(i) * *b.get_unchecked(i);
|
|
537
|
+
i += 1;
|
|
538
|
+
}
|
|
539
|
+
dot
|
|
540
|
+
}
|
|
541
|
+
|
|
542
|
+
#[inline]
|
|
543
|
+
pub unsafe fn dot_f32_avx512(a: &[f32], b: &[f32]) -> f32
|
|
544
|
+
{
|
|
545
|
+
let mut sum = _mm512_setzero_ps();
|
|
546
|
+
let mut i = 0usize;
|
|
547
|
+
while i + 16 <= a.len()
|
|
548
|
+
{
|
|
549
|
+
let va = _mm512_loadu_ps(a.as_ptr().add(i));
|
|
550
|
+
let vb = _mm512_loadu_ps(b.as_ptr().add(i));
|
|
551
|
+
sum = _mm512_fmadd_ps(va, vb, sum);
|
|
552
|
+
i += 16;
|
|
553
|
+
}
|
|
554
|
+
let mut dot = _mm512_reduce_add_ps(sum);
|
|
555
|
+
while i < a.len()
|
|
556
|
+
{
|
|
557
|
+
dot += *a.get_unchecked(i) * *b.get_unchecked(i);
|
|
558
|
+
i += 1;
|
|
559
|
+
}
|
|
560
|
+
dot
|
|
561
|
+
}
|
|
562
|
+
|
|
563
|
+
#[inline]
|
|
564
|
+
pub unsafe fn l2_sq_f32_sse(a: &[f32], b: &[f32]) -> f32
|
|
565
|
+
{
|
|
566
|
+
let mut sum = _mm_setzero_ps();
|
|
567
|
+
let mut i = 0usize;
|
|
568
|
+
while i + 4 <= a.len()
|
|
569
|
+
{
|
|
570
|
+
let va = _mm_loadu_ps(a.as_ptr().add(i));
|
|
571
|
+
let vb = _mm_loadu_ps(b.as_ptr().add(i));
|
|
572
|
+
let diff = _mm_sub_ps(va, vb);
|
|
573
|
+
sum = _mm_add_ps(sum, _mm_mul_ps(diff, diff));
|
|
574
|
+
i += 4;
|
|
575
|
+
}
|
|
576
|
+
let mut tmp = [0.0f32; 4];
|
|
577
|
+
_mm_storeu_ps(tmp.as_mut_ptr(), sum);
|
|
578
|
+
let mut l2 = tmp.iter().sum::<f32>();
|
|
579
|
+
while i < a.len()
|
|
580
|
+
{
|
|
581
|
+
let d = *a.get_unchecked(i) - *b.get_unchecked(i);
|
|
582
|
+
l2 += d * d;
|
|
583
|
+
i += 1;
|
|
584
|
+
}
|
|
585
|
+
l2
|
|
586
|
+
}
|
|
587
|
+
|
|
588
|
+
#[inline]
|
|
589
|
+
pub unsafe fn l2_sq_f32_avx(a: &[f32], b: &[f32]) -> f32
|
|
590
|
+
{
|
|
591
|
+
let mut sum = _mm256_setzero_ps();
|
|
592
|
+
let mut i = 0usize;
|
|
593
|
+
while i + 8 <= a.len()
|
|
594
|
+
{
|
|
595
|
+
let va = _mm256_loadu_ps(a.as_ptr().add(i));
|
|
596
|
+
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
|
|
597
|
+
let diff = _mm256_sub_ps(va, vb);
|
|
598
|
+
sum = _mm256_add_ps(sum, _mm256_mul_ps(diff, diff));
|
|
599
|
+
i += 8;
|
|
600
|
+
}
|
|
601
|
+
let mut tmp = [0.0f32; 8];
|
|
602
|
+
_mm256_storeu_ps(tmp.as_mut_ptr(), sum);
|
|
603
|
+
let mut l2 = tmp.iter().sum::<f32>();
|
|
604
|
+
while i < a.len()
|
|
605
|
+
{
|
|
606
|
+
let d = *a.get_unchecked(i) - *b.get_unchecked(i);
|
|
607
|
+
l2 += d * d;
|
|
608
|
+
i += 1;
|
|
609
|
+
}
|
|
610
|
+
l2
|
|
611
|
+
}
|
|
612
|
+
|
|
613
|
+
#[inline]
|
|
614
|
+
pub unsafe fn l2_sq_f32_avx512(a: &[f32], b: &[f32]) -> f32
|
|
615
|
+
{
|
|
616
|
+
let mut sum = _mm512_setzero_ps();
|
|
617
|
+
let mut i = 0usize;
|
|
618
|
+
while i + 16 <= a.len()
|
|
619
|
+
{
|
|
620
|
+
let va = _mm512_loadu_ps(a.as_ptr().add(i));
|
|
621
|
+
let vb = _mm512_loadu_ps(b.as_ptr().add(i));
|
|
622
|
+
let diff = _mm512_sub_ps(va, vb);
|
|
623
|
+
sum = _mm512_fmadd_ps(diff, diff, sum);
|
|
624
|
+
i += 16;
|
|
625
|
+
}
|
|
626
|
+
let mut l2 = _mm512_reduce_add_ps(sum);
|
|
627
|
+
while i < a.len()
|
|
628
|
+
{
|
|
629
|
+
let d = *a.get_unchecked(i) - *b.get_unchecked(i);
|
|
630
|
+
l2 += d * d;
|
|
631
|
+
i += 1;
|
|
632
|
+
}
|
|
633
|
+
l2
|
|
634
|
+
}
|
|
635
|
+
|
|
636
|
+
#[inline]
|
|
637
|
+
pub unsafe fn i8_sums_avx2(a: &[i8], b: &[i8]) -> I8Sums
|
|
638
|
+
{
|
|
639
|
+
debug_assert_eq!(a.len(), b.len());
|
|
640
|
+
|
|
641
|
+
let ones_i16 = _mm256_set1_epi16(1);
|
|
642
|
+
let mut dot_acc = _mm256_setzero_si256();
|
|
643
|
+
let mut sum_a_acc = _mm256_setzero_si256();
|
|
644
|
+
let mut sum_b_acc = _mm256_setzero_si256();
|
|
645
|
+
let mut sum_sq_a_acc = _mm256_setzero_si256();
|
|
646
|
+
let mut sum_sq_b_acc = _mm256_setzero_si256();
|
|
647
|
+
|
|
648
|
+
let mut i = 0usize;
|
|
649
|
+
while i + 32 <= a.len()
|
|
650
|
+
{
|
|
651
|
+
let a_i8 = _mm256_loadu_si256(a.as_ptr().add(i) as *const __m256i);
|
|
652
|
+
let b_i8 = _mm256_loadu_si256(b.as_ptr().add(i) as *const __m256i);
|
|
653
|
+
|
|
654
|
+
let a_lo_i16 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(a_i8));
|
|
655
|
+
let a_hi_i16 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(a_i8, 1));
|
|
656
|
+
let b_lo_i16 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(b_i8));
|
|
657
|
+
let b_hi_i16 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(b_i8, 1));
|
|
658
|
+
|
|
659
|
+
let prod_lo = _mm256_mullo_epi16(a_lo_i16, b_lo_i16);
|
|
660
|
+
let prod_hi = _mm256_mullo_epi16(a_hi_i16, b_hi_i16);
|
|
661
|
+
dot_acc = _mm256_add_epi32(dot_acc, _mm256_madd_epi16(prod_lo, ones_i16));
|
|
662
|
+
dot_acc = _mm256_add_epi32(dot_acc, _mm256_madd_epi16(prod_hi, ones_i16));
|
|
663
|
+
|
|
664
|
+
sum_a_acc = _mm256_add_epi32(sum_a_acc, _mm256_madd_epi16(a_lo_i16, ones_i16));
|
|
665
|
+
sum_a_acc = _mm256_add_epi32(sum_a_acc, _mm256_madd_epi16(a_hi_i16, ones_i16));
|
|
666
|
+
sum_b_acc = _mm256_add_epi32(sum_b_acc, _mm256_madd_epi16(b_lo_i16, ones_i16));
|
|
667
|
+
sum_b_acc = _mm256_add_epi32(sum_b_acc, _mm256_madd_epi16(b_hi_i16, ones_i16));
|
|
668
|
+
|
|
669
|
+
let a2_lo = _mm256_mullo_epi16(a_lo_i16, a_lo_i16);
|
|
670
|
+
let a2_hi = _mm256_mullo_epi16(a_hi_i16, a_hi_i16);
|
|
671
|
+
let b2_lo = _mm256_mullo_epi16(b_lo_i16, b_lo_i16);
|
|
672
|
+
let b2_hi = _mm256_mullo_epi16(b_hi_i16, b_hi_i16);
|
|
673
|
+
sum_sq_a_acc = _mm256_add_epi32(sum_sq_a_acc, _mm256_madd_epi16(a2_lo, ones_i16));
|
|
674
|
+
sum_sq_a_acc = _mm256_add_epi32(sum_sq_a_acc, _mm256_madd_epi16(a2_hi, ones_i16));
|
|
675
|
+
sum_sq_b_acc = _mm256_add_epi32(sum_sq_b_acc, _mm256_madd_epi16(b2_lo, ones_i16));
|
|
676
|
+
sum_sq_b_acc = _mm256_add_epi32(sum_sq_b_acc, _mm256_madd_epi16(b2_hi, ones_i16));
|
|
677
|
+
|
|
678
|
+
i += 32;
|
|
679
|
+
}
|
|
680
|
+
|
|
681
|
+
let mut dot_buf = [0i32; 8];
|
|
682
|
+
let mut sum_a_buf = [0i32; 8];
|
|
683
|
+
let mut sum_b_buf = [0i32; 8];
|
|
684
|
+
let mut sum_sq_a_buf = [0i32; 8];
|
|
685
|
+
let mut sum_sq_b_buf = [0i32; 8];
|
|
686
|
+
_mm256_storeu_si256(dot_buf.as_mut_ptr() as *mut __m256i, dot_acc);
|
|
687
|
+
_mm256_storeu_si256(sum_a_buf.as_mut_ptr() as *mut __m256i, sum_a_acc);
|
|
688
|
+
_mm256_storeu_si256(sum_b_buf.as_mut_ptr() as *mut __m256i, sum_b_acc);
|
|
689
|
+
_mm256_storeu_si256(sum_sq_a_buf.as_mut_ptr() as *mut __m256i, sum_sq_a_acc);
|
|
690
|
+
_mm256_storeu_si256(sum_sq_b_buf.as_mut_ptr() as *mut __m256i, sum_sq_b_acc);
|
|
691
|
+
|
|
692
|
+
let mut dot = dot_buf.iter().sum::<i32>();
|
|
693
|
+
let mut sum_a = sum_a_buf.iter().sum::<i32>();
|
|
694
|
+
let mut sum_b = sum_b_buf.iter().sum::<i32>();
|
|
695
|
+
let mut sum_sq_a = sum_sq_a_buf.iter().sum::<i32>();
|
|
696
|
+
let mut sum_sq_b = sum_sq_b_buf.iter().sum::<i32>();
|
|
697
|
+
|
|
698
|
+
while i < a.len()
|
|
699
|
+
{
|
|
700
|
+
let av = *a.get_unchecked(i) as i32;
|
|
701
|
+
let bv = *b.get_unchecked(i) as i32;
|
|
702
|
+
dot += av * bv;
|
|
703
|
+
sum_a += av;
|
|
704
|
+
sum_b += bv;
|
|
705
|
+
sum_sq_a += av * av;
|
|
706
|
+
sum_sq_b += bv * bv;
|
|
707
|
+
i += 1;
|
|
708
|
+
}
|
|
709
|
+
|
|
710
|
+
I8Sums {
|
|
711
|
+
dot,
|
|
712
|
+
sum_a,
|
|
713
|
+
sum_b,
|
|
714
|
+
sum_sq_a,
|
|
715
|
+
sum_sq_b,
|
|
716
|
+
}
|
|
717
|
+
}
|
|
718
|
+
|
|
719
|
+
#[inline]
|
|
720
|
+
pub unsafe fn dot_and_norms_f32_sse(a: &[f32], b: &[f32]) -> (f32, f32, f32)
|
|
721
|
+
{
|
|
722
|
+
let mut dot_acc = _mm_setzero_ps();
|
|
723
|
+
let mut na_acc = _mm_setzero_ps();
|
|
724
|
+
let mut nb_acc = _mm_setzero_ps();
|
|
725
|
+
let mut i = 0usize;
|
|
726
|
+
while i + 4 <= a.len()
|
|
727
|
+
{
|
|
728
|
+
let va = _mm_loadu_ps(a.as_ptr().add(i));
|
|
729
|
+
let vb = _mm_loadu_ps(b.as_ptr().add(i));
|
|
730
|
+
dot_acc = _mm_add_ps(dot_acc, _mm_mul_ps(va, vb));
|
|
731
|
+
na_acc = _mm_add_ps(na_acc, _mm_mul_ps(va, va));
|
|
732
|
+
nb_acc = _mm_add_ps(nb_acc, _mm_mul_ps(vb, vb));
|
|
733
|
+
i += 4;
|
|
734
|
+
}
|
|
735
|
+
let mut dot_tmp = [0.0f32; 4];
|
|
736
|
+
let mut na_tmp = [0.0f32; 4];
|
|
737
|
+
let mut nb_tmp = [0.0f32; 4];
|
|
738
|
+
_mm_storeu_ps(dot_tmp.as_mut_ptr(), dot_acc);
|
|
739
|
+
_mm_storeu_ps(na_tmp.as_mut_ptr(), na_acc);
|
|
740
|
+
_mm_storeu_ps(nb_tmp.as_mut_ptr(), nb_acc);
|
|
741
|
+
let mut dot = dot_tmp.iter().sum::<f32>();
|
|
742
|
+
let mut na = na_tmp.iter().sum::<f32>();
|
|
743
|
+
let mut nb = nb_tmp.iter().sum::<f32>();
|
|
744
|
+
while i < a.len()
|
|
745
|
+
{
|
|
746
|
+
let va = *a.get_unchecked(i);
|
|
747
|
+
let vb = *b.get_unchecked(i);
|
|
748
|
+
dot += va * vb;
|
|
749
|
+
na += va * va;
|
|
750
|
+
nb += vb * vb;
|
|
751
|
+
i += 1;
|
|
752
|
+
}
|
|
753
|
+
(dot, na, nb)
|
|
754
|
+
}
|
|
755
|
+
|
|
756
|
+
#[inline]
|
|
757
|
+
pub unsafe fn dot_and_norms_f32_avx(a: &[f32], b: &[f32]) -> (f32, f32, f32)
|
|
758
|
+
{
|
|
759
|
+
let mut dot_acc = _mm256_setzero_ps();
|
|
760
|
+
let mut na_acc = _mm256_setzero_ps();
|
|
761
|
+
let mut nb_acc = _mm256_setzero_ps();
|
|
762
|
+
let mut i = 0usize;
|
|
763
|
+
while i + 8 <= a.len()
|
|
764
|
+
{
|
|
765
|
+
let va = _mm256_loadu_ps(a.as_ptr().add(i));
|
|
766
|
+
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
|
|
767
|
+
dot_acc = _mm256_add_ps(dot_acc, _mm256_mul_ps(va, vb));
|
|
768
|
+
na_acc = _mm256_add_ps(na_acc, _mm256_mul_ps(va, va));
|
|
769
|
+
nb_acc = _mm256_add_ps(nb_acc, _mm256_mul_ps(vb, vb));
|
|
770
|
+
i += 8;
|
|
771
|
+
}
|
|
772
|
+
let mut dot_tmp = [0.0f32; 8];
|
|
773
|
+
let mut na_tmp = [0.0f32; 8];
|
|
774
|
+
let mut nb_tmp = [0.0f32; 8];
|
|
775
|
+
_mm256_storeu_ps(dot_tmp.as_mut_ptr(), dot_acc);
|
|
776
|
+
_mm256_storeu_ps(na_tmp.as_mut_ptr(), na_acc);
|
|
777
|
+
_mm256_storeu_ps(nb_tmp.as_mut_ptr(), nb_acc);
|
|
778
|
+
let mut dot = dot_tmp.iter().sum::<f32>();
|
|
779
|
+
let mut na = na_tmp.iter().sum::<f32>();
|
|
780
|
+
let mut nb = nb_tmp.iter().sum::<f32>();
|
|
781
|
+
while i < a.len()
|
|
782
|
+
{
|
|
783
|
+
let va = *a.get_unchecked(i);
|
|
784
|
+
let vb = *b.get_unchecked(i);
|
|
785
|
+
dot += va * vb;
|
|
786
|
+
na += va * va;
|
|
787
|
+
nb += vb * vb;
|
|
788
|
+
i += 1;
|
|
789
|
+
}
|
|
790
|
+
(dot, na, nb)
|
|
791
|
+
}
|
|
792
|
+
|
|
793
|
+
#[inline]
|
|
794
|
+
pub unsafe fn dot_and_norms_f32_avx512(a: &[f32], b: &[f32]) -> (f32, f32, f32)
|
|
795
|
+
{
|
|
796
|
+
let mut dot_acc = _mm512_setzero_ps();
|
|
797
|
+
let mut na_acc = _mm512_setzero_ps();
|
|
798
|
+
let mut nb_acc = _mm512_setzero_ps();
|
|
799
|
+
let mut i = 0usize;
|
|
800
|
+
while i + 16 <= a.len()
|
|
801
|
+
{
|
|
802
|
+
let va = _mm512_loadu_ps(a.as_ptr().add(i));
|
|
803
|
+
let vb = _mm512_loadu_ps(b.as_ptr().add(i));
|
|
804
|
+
dot_acc = _mm512_fmadd_ps(va, vb, dot_acc);
|
|
805
|
+
na_acc = _mm512_fmadd_ps(va, va, na_acc);
|
|
806
|
+
nb_acc = _mm512_fmadd_ps(vb, vb, nb_acc);
|
|
807
|
+
i += 16;
|
|
808
|
+
}
|
|
809
|
+
let mut dot = _mm512_reduce_add_ps(dot_acc);
|
|
810
|
+
let mut na = _mm512_reduce_add_ps(na_acc);
|
|
811
|
+
let mut nb = _mm512_reduce_add_ps(nb_acc);
|
|
812
|
+
while i < a.len()
|
|
813
|
+
{
|
|
814
|
+
let va = *a.get_unchecked(i);
|
|
815
|
+
let vb = *b.get_unchecked(i);
|
|
816
|
+
dot += va * vb;
|
|
817
|
+
na += va * va;
|
|
818
|
+
nb += vb * vb;
|
|
819
|
+
i += 1;
|
|
820
|
+
}
|
|
821
|
+
(dot, na, nb)
|
|
822
|
+
}
|
|
823
|
+
|
|
824
|
+
#[inline]
|
|
825
|
+
pub unsafe fn dot_bf16_avx512bf16(a: &[bf16], b: &[bf16]) -> f32
|
|
826
|
+
{
|
|
827
|
+
let len = a.len();
|
|
828
|
+
let mut acc = _mm512_setzero_ps();
|
|
829
|
+
let mut i = 0usize;
|
|
830
|
+
while i + 32 <= len
|
|
831
|
+
{
|
|
832
|
+
let a_ptr = a.as_ptr().add(i) as *const i8;
|
|
833
|
+
let b_ptr = b.as_ptr().add(i) as *const i8;
|
|
834
|
+
let a_i = _mm512_loadu_si512(a_ptr as *const _);
|
|
835
|
+
let b_i = _mm512_loadu_si512(b_ptr as *const _);
|
|
836
|
+
let a_bh: __m512bh = std::mem::transmute(a_i);
|
|
837
|
+
let b_bh: __m512bh = std::mem::transmute(b_i);
|
|
838
|
+
acc = _mm512_dpbf16_ps(acc, a_bh, b_bh);
|
|
839
|
+
i += 32;
|
|
840
|
+
}
|
|
841
|
+
let mut dot = _mm512_reduce_add_ps(acc);
|
|
842
|
+
while i < len
|
|
843
|
+
{
|
|
844
|
+
dot += (*a.get_unchecked(i)).to_f32() * (*b.get_unchecked(i)).to_f32();
|
|
845
|
+
i += 1;
|
|
846
|
+
}
|
|
847
|
+
dot
|
|
848
|
+
}
|
|
849
|
+
|
|
850
|
+
#[inline]
|
|
851
|
+
pub unsafe fn dot_and_norms_bf16_avx512bf16(a: &[bf16], b: &[bf16]) -> (f32, f32, f32)
|
|
852
|
+
{
|
|
853
|
+
let len = a.len();
|
|
854
|
+
let mut dot_acc = _mm512_setzero_ps();
|
|
855
|
+
let mut na_acc = _mm512_setzero_ps();
|
|
856
|
+
let mut nb_acc = _mm512_setzero_ps();
|
|
857
|
+
let mut i = 0usize;
|
|
858
|
+
while i + 32 <= len
|
|
859
|
+
{
|
|
860
|
+
let a_ptr = a.as_ptr().add(i) as *const i8;
|
|
861
|
+
let b_ptr = b.as_ptr().add(i) as *const i8;
|
|
862
|
+
let a_i = _mm512_loadu_si512(a_ptr as *const _);
|
|
863
|
+
let b_i = _mm512_loadu_si512(b_ptr as *const _);
|
|
864
|
+
let a_bh: __m512bh = std::mem::transmute(a_i);
|
|
865
|
+
let b_bh: __m512bh = std::mem::transmute(b_i);
|
|
866
|
+
dot_acc = _mm512_dpbf16_ps(dot_acc, a_bh, b_bh);
|
|
867
|
+
na_acc = _mm512_dpbf16_ps(na_acc, a_bh, a_bh);
|
|
868
|
+
nb_acc = _mm512_dpbf16_ps(nb_acc, b_bh, b_bh);
|
|
869
|
+
i += 32;
|
|
870
|
+
}
|
|
871
|
+
let mut dot = _mm512_reduce_add_ps(dot_acc);
|
|
872
|
+
let mut na = _mm512_reduce_add_ps(na_acc);
|
|
873
|
+
let mut nb = _mm512_reduce_add_ps(nb_acc);
|
|
874
|
+
while i < len
|
|
875
|
+
{
|
|
876
|
+
let va = (*a.get_unchecked(i)).to_f32();
|
|
877
|
+
let vb = (*b.get_unchecked(i)).to_f32();
|
|
878
|
+
dot += va * vb;
|
|
879
|
+
na += va * va;
|
|
880
|
+
nb += vb * vb;
|
|
881
|
+
i += 1;
|
|
882
|
+
}
|
|
883
|
+
(dot, na, nb)
|
|
884
|
+
}
|
|
885
|
+
|
|
886
|
+
#[inline]
|
|
887
|
+
pub unsafe fn l2_sq_bf16_avx512bf16(a: &[bf16], b: &[bf16]) -> f32
|
|
888
|
+
{
|
|
889
|
+
use std::arch::x86_64::{__m256bh, __m256i};
|
|
890
|
+
|
|
891
|
+
let len = a.len();
|
|
892
|
+
let mut acc = _mm512_setzero_ps();
|
|
893
|
+
let mut i = 0usize;
|
|
894
|
+
while i + 32 <= len
|
|
895
|
+
{
|
|
896
|
+
let a_ptr = a.as_ptr().add(i) as *const i8;
|
|
897
|
+
let b_ptr = b.as_ptr().add(i) as *const i8;
|
|
898
|
+
let a_i = _mm512_loadu_si512(a_ptr as *const _);
|
|
899
|
+
let b_i = _mm512_loadu_si512(b_ptr as *const _);
|
|
900
|
+
|
|
901
|
+
let a_low = _mm512_castsi512_si256(a_i);
|
|
902
|
+
let b_low = _mm512_castsi512_si256(b_i);
|
|
903
|
+
let a_low_ps = _mm512_cvtpbh_ps(std::mem::transmute::<__m256i, __m256bh>(a_low));
|
|
904
|
+
let b_low_ps = _mm512_cvtpbh_ps(std::mem::transmute::<__m256i, __m256bh>(b_low));
|
|
905
|
+
|
|
906
|
+
let a_high = _mm512_extracti64x4_epi64(a_i, 1);
|
|
907
|
+
let b_high = _mm512_extracti64x4_epi64(b_i, 1);
|
|
908
|
+
let a_high_ps = _mm512_cvtpbh_ps(std::mem::transmute::<__m256i, __m256bh>(a_high));
|
|
909
|
+
let b_high_ps = _mm512_cvtpbh_ps(std::mem::transmute::<__m256i, __m256bh>(b_high));
|
|
910
|
+
|
|
911
|
+
let diff_low = _mm512_sub_ps(a_low_ps, b_low_ps);
|
|
912
|
+
let diff_high = _mm512_sub_ps(a_high_ps, b_high_ps);
|
|
913
|
+
acc = _mm512_fmadd_ps(diff_low, diff_low, acc);
|
|
914
|
+
acc = _mm512_fmadd_ps(diff_high, diff_high, acc);
|
|
915
|
+
i += 32;
|
|
916
|
+
}
|
|
917
|
+
let mut sum = _mm512_reduce_add_ps(acc);
|
|
918
|
+
while i < len
|
|
919
|
+
{
|
|
920
|
+
let d = (*a.get_unchecked(i)).to_f32() - (*b.get_unchecked(i)).to_f32();
|
|
921
|
+
sum += d * d;
|
|
922
|
+
i += 1;
|
|
923
|
+
}
|
|
924
|
+
sum
|
|
925
|
+
}
|
|
926
|
+
|
|
927
|
+
#[inline]
|
|
928
|
+
pub unsafe fn dot_f16_avx512f_f16c(a: &[f16], b: &[f16]) -> f32
|
|
929
|
+
{
|
|
930
|
+
let len = a.len();
|
|
931
|
+
let mut acc = _mm512_setzero_ps();
|
|
932
|
+
let mut i = 0usize;
|
|
933
|
+
while i + 16 <= len
|
|
934
|
+
{
|
|
935
|
+
let a_i = _mm256_loadu_si256(a.as_ptr().add(i) as *const _);
|
|
936
|
+
let b_i = _mm256_loadu_si256(b.as_ptr().add(i) as *const _);
|
|
937
|
+
let a_ps = _mm512_cvtph_ps(a_i);
|
|
938
|
+
let b_ps = _mm512_cvtph_ps(b_i);
|
|
939
|
+
acc = _mm512_fmadd_ps(a_ps, b_ps, acc);
|
|
940
|
+
i += 16;
|
|
941
|
+
}
|
|
942
|
+
let mut dot = _mm512_reduce_add_ps(acc);
|
|
943
|
+
while i < len
|
|
944
|
+
{
|
|
945
|
+
dot += (*a.get_unchecked(i)).to_f32() * (*b.get_unchecked(i)).to_f32();
|
|
946
|
+
i += 1;
|
|
947
|
+
}
|
|
948
|
+
dot
|
|
949
|
+
}
|
|
950
|
+
|
|
951
|
+
#[inline]
|
|
952
|
+
pub unsafe fn l2_sq_f16_avx512f_f16c(a: &[f16], b: &[f16]) -> f32
|
|
953
|
+
{
|
|
954
|
+
let len = a.len();
|
|
955
|
+
let mut acc = _mm512_setzero_ps();
|
|
956
|
+
let mut i = 0usize;
|
|
957
|
+
while i + 16 <= len
|
|
958
|
+
{
|
|
959
|
+
let a_i = _mm256_loadu_si256(a.as_ptr().add(i) as *const _);
|
|
960
|
+
let b_i = _mm256_loadu_si256(b.as_ptr().add(i) as *const _);
|
|
961
|
+
let a_ps = _mm512_cvtph_ps(a_i);
|
|
962
|
+
let b_ps = _mm512_cvtph_ps(b_i);
|
|
963
|
+
let diff = _mm512_sub_ps(a_ps, b_ps);
|
|
964
|
+
acc = _mm512_fmadd_ps(diff, diff, acc);
|
|
965
|
+
i += 16;
|
|
966
|
+
}
|
|
967
|
+
let mut sum = _mm512_reduce_add_ps(acc);
|
|
968
|
+
while i < len
|
|
969
|
+
{
|
|
970
|
+
let d = (*a.get_unchecked(i)).to_f32() - (*b.get_unchecked(i)).to_f32();
|
|
971
|
+
sum += d * d;
|
|
972
|
+
i += 1;
|
|
973
|
+
}
|
|
974
|
+
sum
|
|
975
|
+
}
|
|
976
|
+
|
|
977
|
+
#[inline]
|
|
978
|
+
pub unsafe fn dot_and_norms_f16_avx512f_f16c(a: &[f16], b: &[f16]) -> (f32, f32, f32)
|
|
979
|
+
{
|
|
980
|
+
let len = a.len();
|
|
981
|
+
let mut dot_acc = _mm512_setzero_ps();
|
|
982
|
+
let mut na_acc = _mm512_setzero_ps();
|
|
983
|
+
let mut nb_acc = _mm512_setzero_ps();
|
|
984
|
+
let mut i = 0usize;
|
|
985
|
+
while i + 16 <= len
|
|
986
|
+
{
|
|
987
|
+
let a_i = _mm256_loadu_si256(a.as_ptr().add(i) as *const _);
|
|
988
|
+
let b_i = _mm256_loadu_si256(b.as_ptr().add(i) as *const _);
|
|
989
|
+
let a_ps = _mm512_cvtph_ps(a_i);
|
|
990
|
+
let b_ps = _mm512_cvtph_ps(b_i);
|
|
991
|
+
dot_acc = _mm512_fmadd_ps(a_ps, b_ps, dot_acc);
|
|
992
|
+
na_acc = _mm512_fmadd_ps(a_ps, a_ps, na_acc);
|
|
993
|
+
nb_acc = _mm512_fmadd_ps(b_ps, b_ps, nb_acc);
|
|
994
|
+
i += 16;
|
|
995
|
+
}
|
|
996
|
+
let mut dot = _mm512_reduce_add_ps(dot_acc);
|
|
997
|
+
let mut na = _mm512_reduce_add_ps(na_acc);
|
|
998
|
+
let mut nb = _mm512_reduce_add_ps(nb_acc);
|
|
999
|
+
while i < len
|
|
1000
|
+
{
|
|
1001
|
+
let va = (*a.get_unchecked(i)).to_f32();
|
|
1002
|
+
let vb = (*b.get_unchecked(i)).to_f32();
|
|
1003
|
+
dot += va * vb;
|
|
1004
|
+
na += va * va;
|
|
1005
|
+
nb += vb * vb;
|
|
1006
|
+
i += 1;
|
|
1007
|
+
}
|
|
1008
|
+
(dot, na, nb)
|
|
1009
|
+
}
|
|
1010
|
+
|
|
1011
|
+
#[inline]
|
|
1012
|
+
pub unsafe fn dot_f16_avx_f16c(a: &[f16], b: &[f16]) -> f32
|
|
1013
|
+
{
|
|
1014
|
+
let len = a.len();
|
|
1015
|
+
let mut acc = _mm256_setzero_ps();
|
|
1016
|
+
let mut i = 0usize;
|
|
1017
|
+
while i + 8 <= len
|
|
1018
|
+
{
|
|
1019
|
+
let a_i = _mm_loadu_si128(a.as_ptr().add(i) as *const _);
|
|
1020
|
+
let b_i = _mm_loadu_si128(b.as_ptr().add(i) as *const _);
|
|
1021
|
+
let a_ps = _mm256_cvtph_ps(a_i);
|
|
1022
|
+
let b_ps = _mm256_cvtph_ps(b_i);
|
|
1023
|
+
acc = _mm256_add_ps(acc, _mm256_mul_ps(a_ps, b_ps));
|
|
1024
|
+
i += 8;
|
|
1025
|
+
}
|
|
1026
|
+
let mut tmp = [0.0f32; 8];
|
|
1027
|
+
_mm256_storeu_ps(tmp.as_mut_ptr(), acc);
|
|
1028
|
+
let mut dot = tmp.iter().sum::<f32>();
|
|
1029
|
+
while i < len
|
|
1030
|
+
{
|
|
1031
|
+
dot += (*a.get_unchecked(i)).to_f32() * (*b.get_unchecked(i)).to_f32();
|
|
1032
|
+
i += 1;
|
|
1033
|
+
}
|
|
1034
|
+
dot
|
|
1035
|
+
}
|
|
1036
|
+
|
|
1037
|
+
#[inline]
|
|
1038
|
+
pub unsafe fn l2_sq_f16_avx_f16c(a: &[f16], b: &[f16]) -> f32
|
|
1039
|
+
{
|
|
1040
|
+
let len = a.len();
|
|
1041
|
+
let mut acc = _mm256_setzero_ps();
|
|
1042
|
+
let mut i = 0usize;
|
|
1043
|
+
while i + 8 <= len
|
|
1044
|
+
{
|
|
1045
|
+
let a_i = _mm_loadu_si128(a.as_ptr().add(i) as *const _);
|
|
1046
|
+
let b_i = _mm_loadu_si128(b.as_ptr().add(i) as *const _);
|
|
1047
|
+
let a_ps = _mm256_cvtph_ps(a_i);
|
|
1048
|
+
let b_ps = _mm256_cvtph_ps(b_i);
|
|
1049
|
+
let diff = _mm256_sub_ps(a_ps, b_ps);
|
|
1050
|
+
acc = _mm256_add_ps(acc, _mm256_mul_ps(diff, diff));
|
|
1051
|
+
i += 8;
|
|
1052
|
+
}
|
|
1053
|
+
let mut tmp = [0.0f32; 8];
|
|
1054
|
+
_mm256_storeu_ps(tmp.as_mut_ptr(), acc);
|
|
1055
|
+
let mut sum = tmp.iter().sum::<f32>();
|
|
1056
|
+
while i < len
|
|
1057
|
+
{
|
|
1058
|
+
let d = (*a.get_unchecked(i)).to_f32() - (*b.get_unchecked(i)).to_f32();
|
|
1059
|
+
sum += d * d;
|
|
1060
|
+
i += 1;
|
|
1061
|
+
}
|
|
1062
|
+
sum
|
|
1063
|
+
}
|
|
1064
|
+
|
|
1065
|
+
#[inline]
|
|
1066
|
+
pub unsafe fn dot_and_norms_f16_avx_f16c(a: &[f16], b: &[f16]) -> (f32, f32, f32)
|
|
1067
|
+
{
|
|
1068
|
+
let len = a.len();
|
|
1069
|
+
let mut dot_acc = _mm256_setzero_ps();
|
|
1070
|
+
let mut na_acc = _mm256_setzero_ps();
|
|
1071
|
+
let mut nb_acc = _mm256_setzero_ps();
|
|
1072
|
+
let mut i = 0usize;
|
|
1073
|
+
while i + 8 <= len
|
|
1074
|
+
{
|
|
1075
|
+
let a_i = _mm_loadu_si128(a.as_ptr().add(i) as *const _);
|
|
1076
|
+
let b_i = _mm_loadu_si128(b.as_ptr().add(i) as *const _);
|
|
1077
|
+
let a_ps = _mm256_cvtph_ps(a_i);
|
|
1078
|
+
let b_ps = _mm256_cvtph_ps(b_i);
|
|
1079
|
+
dot_acc = _mm256_add_ps(dot_acc, _mm256_mul_ps(a_ps, b_ps));
|
|
1080
|
+
na_acc = _mm256_add_ps(na_acc, _mm256_mul_ps(a_ps, a_ps));
|
|
1081
|
+
nb_acc = _mm256_add_ps(nb_acc, _mm256_mul_ps(b_ps, b_ps));
|
|
1082
|
+
i += 8;
|
|
1083
|
+
}
|
|
1084
|
+
let mut dot_tmp = [0.0f32; 8];
|
|
1085
|
+
let mut na_tmp = [0.0f32; 8];
|
|
1086
|
+
let mut nb_tmp = [0.0f32; 8];
|
|
1087
|
+
_mm256_storeu_ps(dot_tmp.as_mut_ptr(), dot_acc);
|
|
1088
|
+
_mm256_storeu_ps(na_tmp.as_mut_ptr(), na_acc);
|
|
1089
|
+
_mm256_storeu_ps(nb_tmp.as_mut_ptr(), nb_acc);
|
|
1090
|
+
let mut dot = dot_tmp.iter().sum::<f32>();
|
|
1091
|
+
let mut na = na_tmp.iter().sum::<f32>();
|
|
1092
|
+
let mut nb = nb_tmp.iter().sum::<f32>();
|
|
1093
|
+
while i < len
|
|
1094
|
+
{
|
|
1095
|
+
let va = (*a.get_unchecked(i)).to_f32();
|
|
1096
|
+
let vb = (*b.get_unchecked(i)).to_f32();
|
|
1097
|
+
dot += va * vb;
|
|
1098
|
+
na += va * va;
|
|
1099
|
+
nb += vb * vb;
|
|
1100
|
+
i += 1;
|
|
1101
|
+
}
|
|
1102
|
+
(dot, na, nb)
|
|
1103
|
+
}
|
|
1104
|
+
}
|
|
1105
|
+
|
|
1106
|
+
#[cfg(target_arch = "aarch64")]
|
|
1107
|
+
mod aarch64
|
|
1108
|
+
{
|
|
1109
|
+
use super::I8Sums;
|
|
1110
|
+
use std::arch::aarch64::*;
|
|
1111
|
+
|
|
1112
|
+
#[inline]
|
|
1113
|
+
pub unsafe fn i8_sums_neon(a: &[i8], b: &[i8]) -> I8Sums
|
|
1114
|
+
{
|
|
1115
|
+
debug_assert_eq!(a.len(), b.len());
|
|
1116
|
+
|
|
1117
|
+
let mut dot_acc = vdupq_n_s32(0);
|
|
1118
|
+
let mut sum_a_acc = vdupq_n_s32(0);
|
|
1119
|
+
let mut sum_b_acc = vdupq_n_s32(0);
|
|
1120
|
+
let mut sum_sq_a_acc = vdupq_n_s32(0);
|
|
1121
|
+
let mut sum_sq_b_acc = vdupq_n_s32(0);
|
|
1122
|
+
|
|
1123
|
+
let mut i = 0usize;
|
|
1124
|
+
while i + 16 <= a.len()
|
|
1125
|
+
{
|
|
1126
|
+
let a_i8 = vld1q_s8(a.as_ptr().add(i));
|
|
1127
|
+
let b_i8 = vld1q_s8(b.as_ptr().add(i));
|
|
1128
|
+
|
|
1129
|
+
// dot
|
|
1130
|
+
let prod_lo = vmull_s8(vget_low_s8(a_i8), vget_low_s8(b_i8));
|
|
1131
|
+
let prod_hi = vmull_s8(vget_high_s8(a_i8), vget_high_s8(b_i8));
|
|
1132
|
+
dot_acc = vaddq_s32(dot_acc, vpaddlq_s16(prod_lo));
|
|
1133
|
+
dot_acc = vaddq_s32(dot_acc, vpaddlq_s16(prod_hi));
|
|
1134
|
+
|
|
1135
|
+
// sum
|
|
1136
|
+
let a_lo_i16 = vmovl_s8(vget_low_s8(a_i8));
|
|
1137
|
+
let a_hi_i16 = vmovl_s8(vget_high_s8(a_i8));
|
|
1138
|
+
let b_lo_i16 = vmovl_s8(vget_low_s8(b_i8));
|
|
1139
|
+
let b_hi_i16 = vmovl_s8(vget_high_s8(b_i8));
|
|
1140
|
+
sum_a_acc = vaddq_s32(sum_a_acc, vpaddlq_s16(a_lo_i16));
|
|
1141
|
+
sum_a_acc = vaddq_s32(sum_a_acc, vpaddlq_s16(a_hi_i16));
|
|
1142
|
+
sum_b_acc = vaddq_s32(sum_b_acc, vpaddlq_s16(b_lo_i16));
|
|
1143
|
+
sum_b_acc = vaddq_s32(sum_b_acc, vpaddlq_s16(b_hi_i16));
|
|
1144
|
+
|
|
1145
|
+
// sum_sq
|
|
1146
|
+
let a2_lo = vmull_s8(vget_low_s8(a_i8), vget_low_s8(a_i8));
|
|
1147
|
+
let a2_hi = vmull_s8(vget_high_s8(a_i8), vget_high_s8(a_i8));
|
|
1148
|
+
let b2_lo = vmull_s8(vget_low_s8(b_i8), vget_low_s8(b_i8));
|
|
1149
|
+
let b2_hi = vmull_s8(vget_high_s8(b_i8), vget_high_s8(b_i8));
|
|
1150
|
+
sum_sq_a_acc = vaddq_s32(sum_sq_a_acc, vpaddlq_s16(a2_lo));
|
|
1151
|
+
sum_sq_a_acc = vaddq_s32(sum_sq_a_acc, vpaddlq_s16(a2_hi));
|
|
1152
|
+
sum_sq_b_acc = vaddq_s32(sum_sq_b_acc, vpaddlq_s16(b2_lo));
|
|
1153
|
+
sum_sq_b_acc = vaddq_s32(sum_sq_b_acc, vpaddlq_s16(b2_hi));
|
|
1154
|
+
|
|
1155
|
+
i += 16;
|
|
1156
|
+
}
|
|
1157
|
+
|
|
1158
|
+
let mut dot = vaddvq_s32(dot_acc);
|
|
1159
|
+
let mut sum_a = vaddvq_s32(sum_a_acc);
|
|
1160
|
+
let mut sum_b = vaddvq_s32(sum_b_acc);
|
|
1161
|
+
let mut sum_sq_a = vaddvq_s32(sum_sq_a_acc);
|
|
1162
|
+
let mut sum_sq_b = vaddvq_s32(sum_sq_b_acc);
|
|
1163
|
+
|
|
1164
|
+
while i < a.len()
|
|
1165
|
+
{
|
|
1166
|
+
let av = *a.get_unchecked(i) as i32;
|
|
1167
|
+
let bv = *b.get_unchecked(i) as i32;
|
|
1168
|
+
dot += av * bv;
|
|
1169
|
+
sum_a += av;
|
|
1170
|
+
sum_b += bv;
|
|
1171
|
+
sum_sq_a += av * av;
|
|
1172
|
+
sum_sq_b += bv * bv;
|
|
1173
|
+
i += 1;
|
|
1174
|
+
}
|
|
1175
|
+
|
|
1176
|
+
I8Sums {
|
|
1177
|
+
dot,
|
|
1178
|
+
sum_a,
|
|
1179
|
+
sum_b,
|
|
1180
|
+
sum_sq_a,
|
|
1181
|
+
sum_sq_b,
|
|
1182
|
+
}
|
|
1183
|
+
}
|
|
1184
|
+
}
|
|
1185
|
+
|
|
1186
|
+
#[cfg(test)]
|
|
1187
|
+
mod qi8_tests
|
|
1188
|
+
{
|
|
1189
|
+
use super::*;
|
|
1190
|
+
use rand_09::prelude::Rng;
|
|
1191
|
+
use rand_09::rngs::StdRng;
|
|
1192
|
+
use rand_09::SeedableRng;
|
|
1193
|
+
|
|
1194
|
+
fn dequant(v: &[i8], scale: f32, zero: i8) -> Vec<f32>
|
|
1195
|
+
{
|
|
1196
|
+
v.iter()
|
|
1197
|
+
.map(|&x| (x as i32 - zero as i32) as f32 * scale)
|
|
1198
|
+
.collect()
|
|
1199
|
+
}
|
|
1200
|
+
|
|
1201
|
+
#[test]
|
|
1202
|
+
fn qi8_l2_sq_matches_f32_reference()
|
|
1203
|
+
{
|
|
1204
|
+
let mut rng = StdRng::seed_from_u64(123);
|
|
1205
|
+
for dim in [
|
|
1206
|
+
1usize, 2, 3, 7, 8, 15, 16, 31, 32, 63, 64, 127, 128, 255, 256, 1024,
|
|
1207
|
+
]
|
|
1208
|
+
{
|
|
1209
|
+
for _ in 0..200
|
|
1210
|
+
{
|
|
1211
|
+
let a = (0..dim).map(|_| rng.random::<i8>()).collect::<Vec<_>>();
|
|
1212
|
+
let b = (0..dim).map(|_| rng.random::<i8>()).collect::<Vec<_>>();
|
|
1213
|
+
let a_scale = rng.random_range(0.0001..10.0);
|
|
1214
|
+
let b_scale = rng.random_range(0.0001..10.0);
|
|
1215
|
+
let a_zero = rng.random::<i8>();
|
|
1216
|
+
let b_zero = rng.random::<i8>();
|
|
1217
|
+
|
|
1218
|
+
let af = dequant(&a, a_scale, a_zero);
|
|
1219
|
+
let bf = dequant(&b, b_scale, b_zero);
|
|
1220
|
+
let ref_l2 = l2_sq_f32_scalar(&af, &bf);
|
|
1221
|
+
let got = l2_sq_qi8(&a, a_scale, a_zero, &b, b_scale, b_zero);
|
|
1222
|
+
|
|
1223
|
+
let denom = ref_l2.abs().max(1.0);
|
|
1224
|
+
assert!(((got - ref_l2) / denom).abs() < 1e-3);
|
|
1225
|
+
}
|
|
1226
|
+
}
|
|
1227
|
+
}
|
|
1228
|
+
}
|