static_embeddings 0.1.1 → 0.1.2

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.
@@ -0,0 +1,378 @@
1
+ #include "se_internal.h"
2
+
3
+ #include <stdint.h>
4
+ #include <string.h>
5
+
6
+ #if defined(__ARM_NEON) || defined(__ARM_NEON__)
7
+ #include <arm_neon.h>
8
+ #if defined(__aarch64__)
9
+ #define SE_HAVE_NEON_FP16 1
10
+ #endif
11
+ #endif
12
+
13
+ #if defined(__x86_64__) || defined(__i386__)
14
+ #if defined(__GNUC__) || defined(__clang__)
15
+ #include <cpuid.h>
16
+ #include <immintrin.h>
17
+ #define SE_HAVE_X86_F16C_TARGET 1
18
+ #endif
19
+ #endif
20
+
21
+ static se_f16_backend_t se_f16_backend = SE_F16_BACKEND_LUT;
22
+
23
+ #if !defined(SE_HAVE_NEON_FP16)
24
+ #define SE_NEED_F16_LUT 1
25
+ static float se_f16_lut[65536];
26
+ #endif
27
+
28
+ size_t se_vector_format_element_bytes(se_vector_format_t format) {
29
+ return format == SE_VECTOR_FORMAT_F16 ? 2u : sizeof(float);
30
+ }
31
+
32
+ static void write_u16le(uint8_t *dst, uint16_t v) {
33
+ dst[0] = (uint8_t)(v & 0xffu);
34
+ dst[1] = (uint8_t)(v >> 8);
35
+ }
36
+
37
+ static uint16_t read_u16le(const uint8_t *src) {
38
+ return (uint16_t)src[0] | ((uint16_t)src[1] << 8);
39
+ }
40
+
41
+ uint16_t se_float_to_f16_bits(float value) {
42
+ uint32_t bits;
43
+ memcpy(&bits, &value, sizeof(bits));
44
+
45
+ uint32_t sign = (bits >> 16) & 0x8000u;
46
+ uint32_t exp = (bits >> 23) & 0xffu;
47
+ uint32_t mant = bits & 0x7fffffu;
48
+
49
+ if (exp == 0xffu) {
50
+ if (mant == 0)
51
+ return (uint16_t)(sign | 0x7c00u);
52
+ mant >>= 13;
53
+ return (uint16_t)(sign | 0x7c00u | mant | (mant == 0));
54
+ }
55
+
56
+ int32_t half_exp = (int32_t)exp - 127 + 15;
57
+ if (half_exp >= 31)
58
+ return (uint16_t)(sign | 0x7c00u);
59
+
60
+ if (half_exp <= 0) {
61
+ if (half_exp < -10)
62
+ return (uint16_t)sign;
63
+ mant |= 0x800000u;
64
+ uint32_t shift = (uint32_t)(14 - half_exp);
65
+ uint32_t rounded = (mant + (1u << (shift - 1))) >> shift;
66
+ return (uint16_t)(sign | rounded);
67
+ }
68
+
69
+ mant += 0x1000u;
70
+ if (mant & 0x800000u) {
71
+ mant = 0;
72
+ half_exp++;
73
+ if (half_exp >= 31)
74
+ return (uint16_t)(sign | 0x7c00u);
75
+ }
76
+
77
+ return (uint16_t)(sign | ((uint32_t)half_exp << 10) | (mant >> 13));
78
+ }
79
+
80
+ float se_f16_bits_to_float(uint16_t half) {
81
+ uint32_t sign = ((uint32_t)half & 0x8000u) << 16;
82
+ uint32_t exp = ((uint32_t)half >> 10) & 0x1fu;
83
+ uint32_t mant = (uint32_t)half & 0x03ffu;
84
+ uint32_t bits;
85
+
86
+ if (exp == 0) {
87
+ if (mant == 0) {
88
+ bits = sign;
89
+ } else {
90
+ exp = 1;
91
+ while ((mant & 0x0400u) == 0) {
92
+ mant <<= 1;
93
+ exp--;
94
+ }
95
+ mant &= 0x03ffu;
96
+ bits = sign | ((exp + (127 - 15)) << 23) | (mant << 13);
97
+ }
98
+ } else if (exp == 31) {
99
+ bits = sign | 0x7f800000u | (mant << 13);
100
+ } else {
101
+ bits = sign | ((exp + (127 - 15)) << 23) | (mant << 13);
102
+ }
103
+
104
+ float value;
105
+ memcpy(&value, &bits, sizeof(value));
106
+ return value;
107
+ }
108
+
109
+ void se_write_f16le(uint8_t *dst, float value) {
110
+ write_u16le(dst, se_float_to_f16_bits(value));
111
+ }
112
+
113
+ float se_read_f16le(const uint8_t *src) {
114
+ return se_f16_bits_to_float(read_u16le(src));
115
+ }
116
+
117
+ void se_encode_f16_from_floats(uint8_t *dst, const float *src, size_t count) {
118
+ for (size_t i = 0; i < count; i++)
119
+ write_u16le(dst + i * 2, se_float_to_f16_bits(src[i]));
120
+ }
121
+
122
+ void se_decode_f16_to_floats(float *dst, const uint8_t *src, size_t count) {
123
+ for (size_t i = 0; i < count; i++)
124
+ dst[i] = se_f16_bits_to_float(read_u16le(src + i * 2));
125
+ }
126
+
127
+ #if defined(SE_NEED_F16_LUT)
128
+ static void init_f16_lut(void) {
129
+ for (uint32_t i = 0; i <= 0xffffu; i++)
130
+ se_f16_lut[i] = se_f16_bits_to_float((uint16_t)i);
131
+ }
132
+ #endif
133
+
134
+ #if defined(SE_HAVE_X86_F16C_TARGET)
135
+ #ifndef bit_OSXSAVE
136
+ #define bit_OSXSAVE (1u << 27)
137
+ #endif
138
+ #ifndef bit_AVX
139
+ #define bit_AVX (1u << 28)
140
+ #endif
141
+ #ifndef bit_F16C
142
+ #define bit_F16C (1u << 29)
143
+ #endif
144
+
145
+ static int detect_x86_f16c(void) {
146
+ unsigned int eax = 0, ebx = 0, ecx = 0, edx = 0;
147
+ if (!__get_cpuid(1, &eax, &ebx, &ecx, &edx))
148
+ return 0;
149
+ if ((ecx & bit_OSXSAVE) == 0 || (ecx & bit_AVX) == 0 || (ecx & bit_F16C) == 0)
150
+ return 0;
151
+
152
+ uint32_t xcr0_lo = 0, xcr0_hi = 0;
153
+ #if defined(_MSC_VER)
154
+ return 0;
155
+ #else
156
+ __asm__ volatile("xgetbv" : "=a"(xcr0_lo), "=d"(xcr0_hi) : "c"(0));
157
+ (void)xcr0_hi;
158
+ return (xcr0_lo & 0x6u) == 0x6u;
159
+ #endif
160
+ }
161
+ #endif
162
+
163
+ #if defined(SE_NEED_F16_LUT)
164
+ static float dot_product_f16_lut(const float *q, const uint8_t *row, size_t dim) {
165
+ size_t j = 0;
166
+ float s0 = 0.0f;
167
+ float s1 = 0.0f;
168
+ float s2 = 0.0f;
169
+ float s3 = 0.0f;
170
+ for (; j + 3 < dim; j += 4) {
171
+ s0 += q[j] * se_f16_lut[read_u16le(row + j * 2)];
172
+ s1 += q[j + 1] * se_f16_lut[read_u16le(row + (j + 1) * 2)];
173
+ s2 += q[j + 2] * se_f16_lut[read_u16le(row + (j + 2) * 2)];
174
+ s3 += q[j + 3] * se_f16_lut[read_u16le(row + (j + 3) * 2)];
175
+ }
176
+ float dot = (s0 + s1) + (s2 + s3);
177
+ for (; j < dim; j++)
178
+ dot += q[j] * se_f16_lut[read_u16le(row + j * 2)];
179
+ return dot;
180
+ }
181
+
182
+ static float dot_and_row_sq_f16_lut(const float *q, const uint8_t *row, size_t dim,
183
+ float *row_sq_out) {
184
+ size_t j = 0;
185
+ float d0 = 0.0f;
186
+ float d1 = 0.0f;
187
+ float d2 = 0.0f;
188
+ float d3 = 0.0f;
189
+ float s0 = 0.0f;
190
+ float s1 = 0.0f;
191
+ float s2 = 0.0f;
192
+ float s3 = 0.0f;
193
+ for (; j + 3 < dim; j += 4) {
194
+ float r0 = se_f16_lut[read_u16le(row + j * 2)];
195
+ float r1 = se_f16_lut[read_u16le(row + (j + 1) * 2)];
196
+ float r2 = se_f16_lut[read_u16le(row + (j + 2) * 2)];
197
+ float r3 = se_f16_lut[read_u16le(row + (j + 3) * 2)];
198
+ d0 += q[j] * r0;
199
+ d1 += q[j + 1] * r1;
200
+ d2 += q[j + 2] * r2;
201
+ d3 += q[j + 3] * r3;
202
+ s0 += r0 * r0;
203
+ s1 += r1 * r1;
204
+ s2 += r2 * r2;
205
+ s3 += r3 * r3;
206
+ }
207
+ float dot = (d0 + d1) + (d2 + d3);
208
+ float row_sq = (s0 + s1) + (s2 + s3);
209
+ for (; j < dim; j++) {
210
+ float r = se_f16_lut[read_u16le(row + j * 2)];
211
+ dot += q[j] * r;
212
+ row_sq += r * r;
213
+ }
214
+ *row_sq_out = row_sq;
215
+ return dot;
216
+ }
217
+ #endif /* SE_NEED_F16_LUT */
218
+
219
+ #if defined(SE_HAVE_NEON_FP16)
220
+ static float dot_product_f16_neon(const float *q, const uint8_t *row, size_t dim) {
221
+ size_t j = 0;
222
+ float32x4_t a0 = vdupq_n_f32(0.0f);
223
+ float32x4_t a1 = vdupq_n_f32(0.0f);
224
+ for (; j + 7 < dim; j += 8) {
225
+ float16x4_t h0 =
226
+ vreinterpret_f16_u16(vld1_u16((const uint16_t *)(const void *)(row + j * 2)));
227
+ float16x4_t h1 =
228
+ vreinterpret_f16_u16(vld1_u16((const uint16_t *)(const void *)(row + (j + 4) * 2)));
229
+ float32x4_t r0 = vcvt_f32_f16(h0);
230
+ float32x4_t r1 = vcvt_f32_f16(h1);
231
+ a0 = vmlaq_f32(a0, vld1q_f32(q + j), r0);
232
+ a1 = vmlaq_f32(a1, vld1q_f32(q + j + 4), r1);
233
+ }
234
+ float32x4_t sumv = vaddq_f32(a0, a1);
235
+ float dot = vaddvq_f32(sumv);
236
+ for (; j < dim; j++)
237
+ dot += q[j] * se_f16_bits_to_float(read_u16le(row + j * 2));
238
+ return dot;
239
+ }
240
+
241
+ static float dot_and_row_sq_f16_neon(const float *q, const uint8_t *row, size_t dim,
242
+ float *row_sq_out) {
243
+ size_t j = 0;
244
+ float32x4_t d0 = vdupq_n_f32(0.0f);
245
+ float32x4_t d1 = vdupq_n_f32(0.0f);
246
+ float32x4_t s0 = vdupq_n_f32(0.0f);
247
+ float32x4_t s1 = vdupq_n_f32(0.0f);
248
+ for (; j + 7 < dim; j += 8) {
249
+ float16x4_t h0 =
250
+ vreinterpret_f16_u16(vld1_u16((const uint16_t *)(const void *)(row + j * 2)));
251
+ float16x4_t h1 =
252
+ vreinterpret_f16_u16(vld1_u16((const uint16_t *)(const void *)(row + (j + 4) * 2)));
253
+ float32x4_t r0 = vcvt_f32_f16(h0);
254
+ float32x4_t r1 = vcvt_f32_f16(h1);
255
+ d0 = vmlaq_f32(d0, vld1q_f32(q + j), r0);
256
+ d1 = vmlaq_f32(d1, vld1q_f32(q + j + 4), r1);
257
+ s0 = vmlaq_f32(s0, r0, r0);
258
+ s1 = vmlaq_f32(s1, r1, r1);
259
+ }
260
+ float dot = vaddvq_f32(vaddq_f32(d0, d1));
261
+ float row_sq = vaddvq_f32(vaddq_f32(s0, s1));
262
+ for (; j < dim; j++) {
263
+ float r = se_f16_bits_to_float(read_u16le(row + j * 2));
264
+ dot += q[j] * r;
265
+ row_sq += r * r;
266
+ }
267
+ *row_sq_out = row_sq;
268
+ return dot;
269
+ }
270
+ #endif
271
+
272
+ #if defined(SE_HAVE_X86_F16C_TARGET)
273
+ __attribute__((target("f16c,avx"))) static float hsum256_f16c(__m256 v) {
274
+ __m128 low = _mm256_castps256_ps128(v);
275
+ __m128 high = _mm256_extractf128_ps(v, 1);
276
+ __m128 sum = _mm_add_ps(low, high);
277
+ float tmp[4];
278
+ _mm_storeu_ps(tmp, sum);
279
+ return (tmp[0] + tmp[1]) + (tmp[2] + tmp[3]);
280
+ }
281
+
282
+ __attribute__((target("f16c,avx"))) static float
283
+ dot_product_f16_f16c(const float *q, const uint8_t *row, size_t dim) {
284
+ size_t j = 0;
285
+ __m256 a0 = _mm256_setzero_ps();
286
+ __m256 a1 = _mm256_setzero_ps();
287
+ for (; j + 15 < dim; j += 16) {
288
+ __m256 r0 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i *)(const void *)(row + j * 2)));
289
+ __m256 r1 =
290
+ _mm256_cvtph_ps(_mm_loadu_si128((const __m128i *)(const void *)(row + (j + 8) * 2)));
291
+ a0 = _mm256_add_ps(a0, _mm256_mul_ps(_mm256_loadu_ps(q + j), r0));
292
+ a1 = _mm256_add_ps(a1, _mm256_mul_ps(_mm256_loadu_ps(q + j + 8), r1));
293
+ }
294
+ float dot = hsum256_f16c(_mm256_add_ps(a0, a1));
295
+ for (; j + 7 < dim; j += 8) {
296
+ __m256 r = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i *)(const void *)(row + j * 2)));
297
+ dot += hsum256_f16c(_mm256_mul_ps(_mm256_loadu_ps(q + j), r));
298
+ }
299
+ for (; j < dim; j++)
300
+ dot += q[j] * se_f16_bits_to_float(read_u16le(row + j * 2));
301
+ return dot;
302
+ }
303
+
304
+ __attribute__((target("f16c,avx"))) static float
305
+ dot_and_row_sq_f16_f16c(const float *q, const uint8_t *row, size_t dim, float *row_sq_out) {
306
+ size_t j = 0;
307
+ __m256 d0 = _mm256_setzero_ps();
308
+ __m256 d1 = _mm256_setzero_ps();
309
+ __m256 s0 = _mm256_setzero_ps();
310
+ __m256 s1 = _mm256_setzero_ps();
311
+ for (; j + 15 < dim; j += 16) {
312
+ __m256 r0 = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i *)(const void *)(row + j * 2)));
313
+ __m256 r1 =
314
+ _mm256_cvtph_ps(_mm_loadu_si128((const __m128i *)(const void *)(row + (j + 8) * 2)));
315
+ d0 = _mm256_add_ps(d0, _mm256_mul_ps(_mm256_loadu_ps(q + j), r0));
316
+ d1 = _mm256_add_ps(d1, _mm256_mul_ps(_mm256_loadu_ps(q + j + 8), r1));
317
+ s0 = _mm256_add_ps(s0, _mm256_mul_ps(r0, r0));
318
+ s1 = _mm256_add_ps(s1, _mm256_mul_ps(r1, r1));
319
+ }
320
+ float dot = hsum256_f16c(_mm256_add_ps(d0, d1));
321
+ float row_sq = hsum256_f16c(_mm256_add_ps(s0, s1));
322
+ for (; j + 7 < dim; j += 8) {
323
+ __m256 r = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i *)(const void *)(row + j * 2)));
324
+ dot += hsum256_f16c(_mm256_mul_ps(_mm256_loadu_ps(q + j), r));
325
+ row_sq += hsum256_f16c(_mm256_mul_ps(r, r));
326
+ }
327
+ for (; j < dim; j++) {
328
+ float r = se_f16_bits_to_float(read_u16le(row + j * 2));
329
+ dot += q[j] * r;
330
+ row_sq += r * r;
331
+ }
332
+ *row_sq_out = row_sq;
333
+ return dot;
334
+ }
335
+ #endif
336
+
337
+ float se_dot_product_f16(const float *q, const uint8_t *row, size_t dim) {
338
+ #if defined(SE_HAVE_NEON_FP16)
339
+ return dot_product_f16_neon(q, row, dim);
340
+ #else
341
+ #if defined(SE_HAVE_X86_F16C_TARGET)
342
+ if (se_f16_backend == SE_F16_BACKEND_F16C)
343
+ return dot_product_f16_f16c(q, row, dim);
344
+ #endif
345
+ return dot_product_f16_lut(q, row, dim);
346
+ #endif
347
+ }
348
+
349
+ float se_dot_and_row_sq_f16(const float *q, const uint8_t *row, size_t dim, float *row_sq_out) {
350
+ #if defined(SE_HAVE_NEON_FP16)
351
+ return dot_and_row_sq_f16_neon(q, row, dim, row_sq_out);
352
+ #else
353
+ #if defined(SE_HAVE_X86_F16C_TARGET)
354
+ if (se_f16_backend == SE_F16_BACKEND_F16C)
355
+ return dot_and_row_sq_f16_f16c(q, row, dim, row_sq_out);
356
+ #endif
357
+ return dot_and_row_sq_f16_lut(q, row, dim, row_sq_out);
358
+ #endif
359
+ }
360
+
361
+ void se_select_f16_backend(void) {
362
+ #if defined(SE_HAVE_NEON_FP16)
363
+ se_f16_backend = SE_F16_BACKEND_NEON_FP16;
364
+ #else
365
+ #if defined(SE_HAVE_X86_F16C_TARGET)
366
+ if (detect_x86_f16c()) {
367
+ se_f16_backend = SE_F16_BACKEND_F16C;
368
+ return;
369
+ }
370
+ #endif
371
+ se_f16_backend = SE_F16_BACKEND_LUT;
372
+ init_f16_lut();
373
+ #endif
374
+ }
375
+
376
+ se_f16_backend_t se_current_f16_backend(void) {
377
+ return se_f16_backend;
378
+ }