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.
Files changed (119) hide show
  1. checksums.yaml +4 -4
  2. data/ext/Cargo.toml +12 -6
  3. data/ext/db-index/Cargo.toml +29 -0
  4. data/ext/db-index/src/hnsw/error.rs +66 -0
  5. data/ext/db-index/src/hnsw/graph_store.rs +272 -0
  6. data/ext/db-index/src/hnsw/id.rs +36 -0
  7. data/ext/db-index/src/hnsw/implementation.rs +1517 -0
  8. data/ext/db-index/src/hnsw/kernels.rs +1228 -0
  9. data/ext/db-index/src/hnsw/metric.rs +244 -0
  10. data/ext/db-index/src/hnsw/scalar.rs +72 -0
  11. data/ext/db-index/src/hnsw/simple_store.rs +140 -0
  12. data/ext/db-index/src/hnsw/store.rs +472 -0
  13. data/ext/db-index/src/hnsw/vector.rs +105 -0
  14. data/ext/db-index/src/hnsw/vectors.rs +568 -0
  15. data/ext/db-index/src/hnsw.rs +42 -0
  16. data/ext/db-index/src/lib.rs +3 -0
  17. data/ext/gqlitedb/Cargo.toml +19 -8
  18. data/ext/gqlitedb/benches/common/pokec.rs +62 -2
  19. data/ext/gqlitedb/benches/pokec_divan.rs +66 -2
  20. data/ext/gqlitedb/benches/pokec_iai.rs +60 -3
  21. data/ext/gqlitedb/release.toml +2 -2
  22. data/ext/gqlitedb/src/aggregators/arithmetic.rs +3 -1
  23. data/ext/gqlitedb/src/aggregators/containers.rs +1 -1
  24. data/ext/gqlitedb/src/aggregators/stats.rs +21 -12
  25. data/ext/gqlitedb/src/capi.rs +20 -24
  26. data/ext/gqlitedb/src/compiler/expression_analyser.rs +28 -20
  27. data/ext/gqlitedb/src/compiler/variables_manager.rs +104 -38
  28. data/ext/gqlitedb/src/compiler.rs +506 -227
  29. data/ext/gqlitedb/src/connection.rs +151 -11
  30. data/ext/gqlitedb/src/consts.rs +1 -2
  31. data/ext/gqlitedb/src/error.rs +123 -61
  32. data/ext/gqlitedb/src/functions/common.rs +64 -0
  33. data/ext/gqlitedb/src/functions/containers.rs +2 -46
  34. data/ext/gqlitedb/src/functions/edge.rs +1 -1
  35. data/ext/gqlitedb/src/functions/math.rs +87 -15
  36. data/ext/gqlitedb/src/functions/node.rs +1 -1
  37. data/ext/gqlitedb/src/functions/path.rs +48 -11
  38. data/ext/gqlitedb/src/functions/scalar.rs +47 -4
  39. data/ext/gqlitedb/src/functions/string.rs +123 -1
  40. data/ext/gqlitedb/src/functions/value.rs +1 -1
  41. data/ext/gqlitedb/src/functions.rs +165 -28
  42. data/ext/gqlitedb/src/graph.rs +1 -9
  43. data/ext/gqlitedb/src/interpreter/evaluators.rs +968 -130
  44. data/ext/gqlitedb/src/interpreter/instructions.rs +39 -2
  45. data/ext/gqlitedb/src/lib.rs +5 -4
  46. data/ext/gqlitedb/src/parser/gql.pest +1 -1
  47. data/ext/gqlitedb/src/planner.rs +329 -0
  48. data/ext/gqlitedb/src/prelude.rs +3 -3
  49. data/ext/gqlitedb/src/store/pgrx.rs +1 -1
  50. data/ext/gqlitedb/src/store/postgres.rs +742 -16
  51. data/ext/gqlitedb/src/store/redb/hnsw_store.rs +702 -0
  52. data/ext/gqlitedb/src/store/redb/index.rs +274 -0
  53. data/ext/gqlitedb/src/store/redb.rs +1268 -113
  54. data/ext/gqlitedb/src/store/sqlbase/sqlmetadata.rs +28 -0
  55. data/ext/gqlitedb/src/store/sqlbase/sqlstore.rs +103 -0
  56. data/ext/gqlitedb/src/store/sqlbase.rs +146 -16
  57. data/ext/gqlitedb/src/store/sqlite.rs +576 -14
  58. data/ext/gqlitedb/src/store/vector_extract.rs +56 -0
  59. data/ext/gqlitedb/src/store.rs +140 -29
  60. data/ext/gqlitedb/src/tests/compiler.rs +207 -10
  61. data/ext/gqlitedb/src/tests/connection/postgres.rs +38 -3
  62. data/ext/gqlitedb/src/tests/connection/redb.rs +23 -0
  63. data/ext/gqlitedb/src/tests/connection/sqlite.rs +31 -0
  64. data/ext/gqlitedb/src/tests/connection.rs +61 -1
  65. data/ext/gqlitedb/src/tests/evaluators.rs +162 -7
  66. data/ext/gqlitedb/src/tests/parser.rs +54 -23
  67. data/ext/gqlitedb/src/tests/planner.rs +511 -0
  68. data/ext/gqlitedb/src/tests/store/postgres.rs +7 -0
  69. data/ext/gqlitedb/src/tests/store/redb.rs +8 -0
  70. data/ext/gqlitedb/src/tests/store/sqlite.rs +8 -0
  71. data/ext/gqlitedb/src/tests/store/vector_index/postgres.rs +182 -0
  72. data/ext/gqlitedb/src/tests/store/vector_index/redb.rs +386 -0
  73. data/ext/gqlitedb/src/tests/store/vector_index/sqlite.rs +166 -0
  74. data/ext/gqlitedb/src/tests/store/vector_index.rs +2313 -0
  75. data/ext/gqlitedb/src/tests/store.rs +78 -14
  76. data/ext/gqlitedb/src/tests/templates/ast.rs +92 -16
  77. data/ext/gqlitedb/src/tests/templates/programs.rs +14 -7
  78. data/ext/gqlitedb/src/tests.rs +15 -9
  79. data/ext/gqlitedb/src/utils.rs +6 -1
  80. data/ext/gqlitedb/src/value/compare.rs +61 -3
  81. data/ext/gqlitedb/src/value.rs +136 -7
  82. data/ext/gqlitedb/templates/sql/postgres/metadata_delete.sql +1 -0
  83. data/ext/gqlitedb/templates/sql/postgres/node_select.sql +1 -1
  84. data/ext/gqlitedb/templates/sql/sqlite/metadata_delete.sql +1 -0
  85. data/ext/gqliterb/src/lib.rs +53 -8
  86. data/ext/gqlparser/Cargo.toml +25 -0
  87. data/ext/gqlparser/README.MD +9 -0
  88. data/ext/gqlparser/benches/pokec_divan.rs +34 -0
  89. data/ext/gqlparser/src/common.rs +69 -0
  90. data/ext/gqlparser/src/gqls/ast.rs +69 -0
  91. data/ext/gqlparser/src/gqls/constraint.rs +79 -0
  92. data/ext/gqlparser/src/gqls/error.rs +22 -0
  93. data/ext/gqlparser/src/gqls/parser.rs +813 -0
  94. data/ext/gqlparser/src/gqls/prelude.rs +8 -0
  95. data/ext/gqlparser/src/gqls/properties.rs +115 -0
  96. data/ext/gqlparser/src/gqls/resolve.rs +207 -0
  97. data/ext/gqlparser/src/gqls.rs +265 -0
  98. data/ext/gqlparser/src/lib.rs +7 -0
  99. data/ext/gqlparser/src/oc/ast.rs +680 -0
  100. data/ext/gqlparser/src/oc/error.rs +172 -0
  101. data/ext/gqlparser/src/oc/lexer.rs +429 -0
  102. data/ext/gqlparser/src/oc/parser/tests.rs +2284 -0
  103. data/ext/gqlparser/src/oc/parser.rs +2005 -0
  104. data/ext/gqlparser/src/oc.rs +33 -0
  105. data/ext/gqlparser/src/prelude.rs +3 -0
  106. data/ext/graphcore/Cargo.toml +1 -0
  107. data/ext/graphcore/src/error.rs +26 -0
  108. data/ext/graphcore/src/graph.rs +177 -51
  109. data/ext/graphcore/src/lib.rs +4 -2
  110. data/ext/graphcore/src/open_cypher.rs +12 -0
  111. data/ext/graphcore/src/table.rs +50 -2
  112. data/ext/graphcore/src/timestamp.rs +127 -104
  113. data/ext/graphcore/src/value/tensor.rs +739 -0
  114. data/ext/graphcore/src/value/value_map.rs +1 -1
  115. data/ext/graphcore/src/value.rs +343 -19
  116. metadata +93 -28
  117. data/ext/gqlitedb/src/parser/ast.rs +0 -604
  118. data/ext/gqlitedb/src/parser/parser_impl.rs +0 -1213
  119. 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
+ }