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.
- checksums.yaml +4 -4
- data/CHANGELOG.md +121 -0
- data/README.md +167 -276
- data/docs/ARCHITECTURE.md +56 -6
- data/docs/MODEL_AUDIT.md +21 -8
- data/docs/PERFORMANCE.md +104 -36
- data/ext/static_embeddings/extconf.rb +4 -0
- data/ext/static_embeddings/se_alloc_stats.c +244 -0
- data/ext/static_embeddings/se_f16.c +378 -0
- data/ext/static_embeddings/se_format.c +235 -148
- data/ext/static_embeddings/se_internal.h +159 -0
- data/ext/static_embeddings/se_tokenizer.c +144 -41
- data/ext/static_embeddings/se_topk.c +236 -0
- data/ext/static_embeddings/static_embeddings.c +225 -744
- data/lib/static_embeddings/format.rb +22 -11
- data/lib/static_embeddings/version.rb +1 -1
- data/tools/benchmark.rb +13 -4
- metadata +4 -1
|
@@ -4,7 +4,6 @@
|
|
|
4
4
|
#include <stdlib.h>
|
|
5
5
|
#include <string.h>
|
|
6
6
|
|
|
7
|
-
#define SE_SIZE_MAX ((size_t)-1)
|
|
8
7
|
#define SE_CANCEL_CHECK_MASK 0x3ffu
|
|
9
8
|
|
|
10
9
|
void se_scratch_init(se_scratch_t *s) {
|
|
@@ -12,70 +11,58 @@ void se_scratch_init(se_scratch_t *s) {
|
|
|
12
11
|
}
|
|
13
12
|
|
|
14
13
|
void se_scratch_free(se_scratch_t *s) {
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
|
|
19
|
-
|
|
14
|
+
se_free(s->cps);
|
|
15
|
+
se_free(s->cps2);
|
|
16
|
+
se_free(s->bytes);
|
|
17
|
+
se_free(s->ids);
|
|
18
|
+
se_free(s->acc);
|
|
20
19
|
memset(s, 0, sizeof(*s));
|
|
21
20
|
}
|
|
22
21
|
|
|
23
|
-
static
|
|
22
|
+
static void *grow_buffer(void *buf, size_t *cap, size_t need, size_t elem_size,
|
|
23
|
+
size_t min_capacity) {
|
|
24
24
|
if (*cap >= need)
|
|
25
|
-
return
|
|
25
|
+
return buf;
|
|
26
26
|
|
|
27
|
-
size_t next =
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
|
|
27
|
+
size_t next = 0;
|
|
28
|
+
size_t bytes = 0;
|
|
29
|
+
if (!se_next_capacity(*cap, need, min_capacity, &next) ||
|
|
30
|
+
!se_array_bytes(next, elem_size, &bytes)) {
|
|
31
|
+
return NULL;
|
|
32
32
|
}
|
|
33
33
|
|
|
34
|
-
|
|
35
|
-
|
|
34
|
+
void *p = se_realloc(SE_ALLOC_SCRATCH, buf, bytes);
|
|
35
|
+
if (!p)
|
|
36
|
+
return NULL;
|
|
36
37
|
|
|
37
|
-
|
|
38
|
+
*cap = next;
|
|
39
|
+
return p;
|
|
40
|
+
}
|
|
41
|
+
|
|
42
|
+
static int grow_u32(uint32_t **buf, size_t *cap, size_t need) {
|
|
43
|
+
uint32_t *p = (uint32_t *)grow_buffer(*buf, cap, need, sizeof(uint32_t), 256);
|
|
38
44
|
if (!p)
|
|
39
45
|
return 0;
|
|
40
46
|
|
|
41
47
|
*buf = p;
|
|
42
|
-
*cap = next;
|
|
43
48
|
return 1;
|
|
44
49
|
}
|
|
45
50
|
|
|
46
51
|
static int grow_bytes(uint8_t **buf, size_t *cap, size_t need) {
|
|
47
|
-
|
|
48
|
-
return 1;
|
|
49
|
-
|
|
50
|
-
size_t next = *cap ? *cap : 256;
|
|
51
|
-
while (next < need) {
|
|
52
|
-
if (next > SE_SIZE_MAX / 2)
|
|
53
|
-
return 0;
|
|
54
|
-
next *= 2;
|
|
55
|
-
}
|
|
56
|
-
|
|
57
|
-
uint8_t *p = (uint8_t *)realloc(*buf, next);
|
|
52
|
+
uint8_t *p = (uint8_t *)grow_buffer(*buf, cap, need, sizeof(uint8_t), 256);
|
|
58
53
|
if (!p)
|
|
59
54
|
return 0;
|
|
60
55
|
|
|
61
56
|
*buf = p;
|
|
62
|
-
*cap = next;
|
|
63
57
|
return 1;
|
|
64
58
|
}
|
|
65
59
|
|
|
66
60
|
static int grow_float(float **buf, size_t *cap, size_t need) {
|
|
67
|
-
|
|
68
|
-
return 1;
|
|
69
|
-
|
|
70
|
-
if (need > SE_SIZE_MAX / sizeof(float))
|
|
71
|
-
return 0;
|
|
72
|
-
|
|
73
|
-
float *p = (float *)realloc(*buf, need * sizeof(float));
|
|
61
|
+
float *p = (float *)grow_buffer(*buf, cap, need, sizeof(float), need);
|
|
74
62
|
if (!p)
|
|
75
63
|
return 0;
|
|
76
64
|
|
|
77
65
|
*buf = p;
|
|
78
|
-
*cap = need;
|
|
79
66
|
return 1;
|
|
80
67
|
}
|
|
81
68
|
|
|
@@ -111,7 +98,7 @@ static int is_whitespace(const se_model_t *m, uint32_t cp) {
|
|
|
111
98
|
|
|
112
99
|
static int is_control(const se_model_t *m, uint32_t cp) {
|
|
113
100
|
if (cp < 0x80)
|
|
114
|
-
return cp < 0x20 && cp != '\t' && cp != '\n' && cp != '\r';
|
|
101
|
+
return (cp < 0x20 || cp == 0x7f) && cp != '\t' && cp != '\n' && cp != '\r';
|
|
115
102
|
return se_range_contains(m->norm.control, m->norm.control_count, cp);
|
|
116
103
|
}
|
|
117
104
|
|
|
@@ -223,6 +210,22 @@ size_t se_prefix_boundary_len(const se_model_t *model, const uint8_t *input, siz
|
|
|
223
210
|
return 0;
|
|
224
211
|
}
|
|
225
212
|
|
|
213
|
+
typedef enum {
|
|
214
|
+
SE_ASCII_DROP = 0,
|
|
215
|
+
SE_ASCII_SPACE = 1,
|
|
216
|
+
SE_ASCII_PUNCT = 2,
|
|
217
|
+
SE_ASCII_WORD = 3
|
|
218
|
+
} se_ascii_class_t;
|
|
219
|
+
|
|
220
|
+
static const uint8_t SE_ASCII_CLASS[128] = {
|
|
221
|
+
0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
|
|
222
|
+
1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 2, 2,
|
|
223
|
+
2, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 2,
|
|
224
|
+
2, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 0,
|
|
225
|
+
};
|
|
226
|
+
|
|
227
|
+
#define SE_ASCII_CANCEL_STRIDE (SE_CANCEL_CHECK_MASK + 1u)
|
|
228
|
+
|
|
226
229
|
typedef struct {
|
|
227
230
|
const se_model_t *model;
|
|
228
231
|
se_scratch_t *scratch;
|
|
@@ -282,9 +285,11 @@ static wordpiece_status_t wordpiece(const se_model_t *m, se_scratch_t *sc, const
|
|
|
282
285
|
return 1;
|
|
283
286
|
}
|
|
284
287
|
|
|
285
|
-
|
|
288
|
+
size_t bytes_need = 0;
|
|
289
|
+
if (!se_checked_mul_size(word_len, 4, &bytes_need) ||
|
|
290
|
+
!se_checked_add_size(bytes_need, 4, &bytes_need))
|
|
286
291
|
return 0;
|
|
287
|
-
if (!grow_bytes(&sc->bytes, &sc->bytes_cap,
|
|
292
|
+
if (!grow_bytes(&sc->bytes, &sc->bytes_cap, bytes_need))
|
|
288
293
|
return 0;
|
|
289
294
|
if (!grow_u32(&sc->cps2, &sc->cps2_cap, word_len + 1))
|
|
290
295
|
return 0;
|
|
@@ -452,6 +457,93 @@ static se_status_t emit_cleaned(token_state_t *st, uint32_t cp, int *stop) {
|
|
|
452
457
|
return emit_stripped(st, cp, stop);
|
|
453
458
|
}
|
|
454
459
|
|
|
460
|
+
static inline se_ascii_class_t ascii_class(uint8_t byte) {
|
|
461
|
+
return (se_ascii_class_t)SE_ASCII_CLASS[byte];
|
|
462
|
+
}
|
|
463
|
+
|
|
464
|
+
static size_t ascii_word_reserve_len(const uint8_t *input, size_t input_len, size_t i) {
|
|
465
|
+
size_t run = 1;
|
|
466
|
+
size_t limit = input_len - i;
|
|
467
|
+
|
|
468
|
+
if (limit > SE_ASCII_CANCEL_STRIDE)
|
|
469
|
+
limit = SE_ASCII_CANCEL_STRIDE;
|
|
470
|
+
|
|
471
|
+
while (run < limit && input[i + run] < 0x80 && ascii_class(input[i + run]) == SE_ASCII_WORD)
|
|
472
|
+
run++;
|
|
473
|
+
|
|
474
|
+
return run;
|
|
475
|
+
}
|
|
476
|
+
|
|
477
|
+
static se_status_t tokenize_ascii_run(token_state_t *st, const uint8_t *input, size_t input_len,
|
|
478
|
+
size_t *ip, int *stop) {
|
|
479
|
+
size_t i = *ip;
|
|
480
|
+
size_t next_cancel_check = i + SE_ASCII_CANCEL_STRIDE;
|
|
481
|
+
const int lc = (int)st->model->meta.do_lower_case;
|
|
482
|
+
se_scratch_t *sc = st->scratch;
|
|
483
|
+
|
|
484
|
+
while (i < input_len) {
|
|
485
|
+
uint8_t b = input[i];
|
|
486
|
+
if (b >= 0x80)
|
|
487
|
+
break;
|
|
488
|
+
|
|
489
|
+
if (i >= next_cancel_check) {
|
|
490
|
+
if (token_cancelled(st)) {
|
|
491
|
+
*ip = i;
|
|
492
|
+
se_error_set(st->err, SE_ERR_INTERNAL, "operation cancelled");
|
|
493
|
+
return SE_ERR_INTERNAL;
|
|
494
|
+
}
|
|
495
|
+
next_cancel_check = i + SE_ASCII_CANCEL_STRIDE;
|
|
496
|
+
}
|
|
497
|
+
|
|
498
|
+
se_ascii_class_t cls = ascii_class(b);
|
|
499
|
+
|
|
500
|
+
if (cls == SE_ASCII_WORD) {
|
|
501
|
+
if (st->segment_len >= sc->cps_cap) {
|
|
502
|
+
if (!grow_u32(&sc->cps, &sc->cps_cap,
|
|
503
|
+
st->segment_len + ascii_word_reserve_len(input, input_len, i))) {
|
|
504
|
+
*ip = i;
|
|
505
|
+
return oom(st, "building token segment");
|
|
506
|
+
}
|
|
507
|
+
}
|
|
508
|
+
sc->cps[st->segment_len++] =
|
|
509
|
+
(lc && b >= 'A' && b <= 'Z') ? (uint32_t)(b + 32u) : (uint32_t)b;
|
|
510
|
+
i++;
|
|
511
|
+
continue;
|
|
512
|
+
}
|
|
513
|
+
|
|
514
|
+
if (cls == SE_ASCII_DROP) {
|
|
515
|
+
i++;
|
|
516
|
+
continue;
|
|
517
|
+
}
|
|
518
|
+
|
|
519
|
+
se_status_t rc = flush_segment(st, stop);
|
|
520
|
+
if (rc != SE_OK) {
|
|
521
|
+
*ip = i;
|
|
522
|
+
return rc;
|
|
523
|
+
}
|
|
524
|
+
if (*stop) {
|
|
525
|
+
*ip = i;
|
|
526
|
+
return SE_OK;
|
|
527
|
+
}
|
|
528
|
+
|
|
529
|
+
if (cls == SE_ASCII_PUNCT) {
|
|
530
|
+
uint32_t cp = b;
|
|
531
|
+
rc = append_wordpiece(st, &cp, 1, stop);
|
|
532
|
+
if (rc != SE_OK) {
|
|
533
|
+
*ip = i;
|
|
534
|
+
return rc;
|
|
535
|
+
}
|
|
536
|
+
}
|
|
537
|
+
|
|
538
|
+
i++;
|
|
539
|
+
if (*stop)
|
|
540
|
+
break;
|
|
541
|
+
}
|
|
542
|
+
|
|
543
|
+
*ip = i;
|
|
544
|
+
return SE_OK;
|
|
545
|
+
}
|
|
546
|
+
|
|
455
547
|
se_status_t se_tokenize(const se_model_t *model, se_scratch_t *sc, const uint8_t *input,
|
|
456
548
|
size_t input_len, uint32_t max_tokens, se_token_stats_t *stats,
|
|
457
549
|
se_error_t *err, volatile sig_atomic_t *cancelled) {
|
|
@@ -474,13 +566,24 @@ se_status_t se_tokenize(const se_model_t *model, se_scratch_t *sc, const uint8_t
|
|
|
474
566
|
return SE_ERR_INTERNAL;
|
|
475
567
|
}
|
|
476
568
|
|
|
569
|
+
int stop = 0;
|
|
570
|
+
|
|
571
|
+
if (input[i] < 0x80) {
|
|
572
|
+
se_status_t frc = tokenize_ascii_run(&st, input, input_len, &i, &stop);
|
|
573
|
+
if (frc != SE_OK)
|
|
574
|
+
return frc;
|
|
575
|
+
if (stop)
|
|
576
|
+
break;
|
|
577
|
+
if (i >= input_len)
|
|
578
|
+
break;
|
|
579
|
+
}
|
|
580
|
+
|
|
477
581
|
uint32_t cp;
|
|
478
582
|
if (!decode_one(input, input_len, &i, &cp)) {
|
|
479
583
|
se_error_set(err, SE_ERR_INVALID_UTF8, "input is not valid UTF-8");
|
|
480
584
|
return SE_ERR_INVALID_UTF8;
|
|
481
585
|
}
|
|
482
586
|
|
|
483
|
-
int stop = 0;
|
|
484
587
|
se_status_t rc = emit_cleaned(&st, cp, &stop);
|
|
485
588
|
if (rc != SE_OK)
|
|
486
589
|
return rc;
|
|
@@ -0,0 +1,236 @@
|
|
|
1
|
+
#include "se_internal.h"
|
|
2
|
+
|
|
3
|
+
#include <math.h>
|
|
4
|
+
#include <stdint.h>
|
|
5
|
+
|
|
6
|
+
#if defined(__ARM_NEON) || defined(__ARM_NEON__)
|
|
7
|
+
#include <arm_neon.h>
|
|
8
|
+
#define SE_HAVE_NEON 1
|
|
9
|
+
#elif defined(__SSE__)
|
|
10
|
+
#include <xmmintrin.h>
|
|
11
|
+
#define SE_HAVE_SSE 1
|
|
12
|
+
#endif
|
|
13
|
+
|
|
14
|
+
float se_dot_product_f32(const float *q, const float *row, size_t dim) {
|
|
15
|
+
size_t j = 0;
|
|
16
|
+
#if defined(SE_HAVE_NEON)
|
|
17
|
+
float32x4_t a0 = vdupq_n_f32(0.0f);
|
|
18
|
+
float32x4_t a1 = vdupq_n_f32(0.0f);
|
|
19
|
+
float32x4_t a2 = vdupq_n_f32(0.0f);
|
|
20
|
+
float32x4_t a3 = vdupq_n_f32(0.0f);
|
|
21
|
+
for (; j + 15 < dim; j += 16) {
|
|
22
|
+
a0 = vmlaq_f32(a0, vld1q_f32(q + j), vld1q_f32(row + j));
|
|
23
|
+
a1 = vmlaq_f32(a1, vld1q_f32(q + j + 4), vld1q_f32(row + j + 4));
|
|
24
|
+
a2 = vmlaq_f32(a2, vld1q_f32(q + j + 8), vld1q_f32(row + j + 8));
|
|
25
|
+
a3 = vmlaq_f32(a3, vld1q_f32(q + j + 12), vld1q_f32(row + j + 12));
|
|
26
|
+
}
|
|
27
|
+
float32x4_t sumv = vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3));
|
|
28
|
+
#if defined(__aarch64__)
|
|
29
|
+
float dot = vaddvq_f32(sumv);
|
|
30
|
+
#else
|
|
31
|
+
float32x2_t pair = vadd_f32(vget_low_f32(sumv), vget_high_f32(sumv));
|
|
32
|
+
pair = vpadd_f32(pair, pair);
|
|
33
|
+
float dot = vget_lane_f32(pair, 0);
|
|
34
|
+
#endif
|
|
35
|
+
#elif defined(SE_HAVE_SSE)
|
|
36
|
+
__m128 a0 = _mm_setzero_ps();
|
|
37
|
+
__m128 a1 = _mm_setzero_ps();
|
|
38
|
+
__m128 a2 = _mm_setzero_ps();
|
|
39
|
+
__m128 a3 = _mm_setzero_ps();
|
|
40
|
+
for (; j + 15 < dim; j += 16) {
|
|
41
|
+
a0 = _mm_add_ps(a0, _mm_mul_ps(_mm_loadu_ps(q + j), _mm_loadu_ps(row + j)));
|
|
42
|
+
a1 = _mm_add_ps(a1, _mm_mul_ps(_mm_loadu_ps(q + j + 4), _mm_loadu_ps(row + j + 4)));
|
|
43
|
+
a2 = _mm_add_ps(a2, _mm_mul_ps(_mm_loadu_ps(q + j + 8), _mm_loadu_ps(row + j + 8)));
|
|
44
|
+
a3 = _mm_add_ps(a3, _mm_mul_ps(_mm_loadu_ps(q + j + 12), _mm_loadu_ps(row + j + 12)));
|
|
45
|
+
}
|
|
46
|
+
__m128 sumv = _mm_add_ps(_mm_add_ps(a0, a1), _mm_add_ps(a2, a3));
|
|
47
|
+
float tmp[4];
|
|
48
|
+
_mm_storeu_ps(tmp, sumv);
|
|
49
|
+
float dot = (tmp[0] + tmp[1]) + (tmp[2] + tmp[3]);
|
|
50
|
+
#else
|
|
51
|
+
float s0 = 0.0f;
|
|
52
|
+
float s1 = 0.0f;
|
|
53
|
+
float s2 = 0.0f;
|
|
54
|
+
float s3 = 0.0f;
|
|
55
|
+
float s4 = 0.0f;
|
|
56
|
+
float s5 = 0.0f;
|
|
57
|
+
float s6 = 0.0f;
|
|
58
|
+
float s7 = 0.0f;
|
|
59
|
+
for (; j + 7 < dim; j += 8) {
|
|
60
|
+
s0 += q[j] * row[j];
|
|
61
|
+
s1 += q[j + 1] * row[j + 1];
|
|
62
|
+
s2 += q[j + 2] * row[j + 2];
|
|
63
|
+
s3 += q[j + 3] * row[j + 3];
|
|
64
|
+
s4 += q[j + 4] * row[j + 4];
|
|
65
|
+
s5 += q[j + 5] * row[j + 5];
|
|
66
|
+
s6 += q[j + 6] * row[j + 6];
|
|
67
|
+
s7 += q[j + 7] * row[j + 7];
|
|
68
|
+
}
|
|
69
|
+
float dot = (s0 + s1) + (s2 + s3) + (s4 + s5) + (s6 + s7);
|
|
70
|
+
#endif
|
|
71
|
+
for (; j < dim; j++)
|
|
72
|
+
dot += q[j] * row[j];
|
|
73
|
+
return dot;
|
|
74
|
+
}
|
|
75
|
+
|
|
76
|
+
float se_dot_and_row_sq_f32(const float *q, const float *row, size_t dim,
|
|
77
|
+
float *row_sq_out) {
|
|
78
|
+
size_t j = 0;
|
|
79
|
+
#if defined(SE_HAVE_NEON)
|
|
80
|
+
float32x4_t d0 = vdupq_n_f32(0.0f);
|
|
81
|
+
float32x4_t d1 = vdupq_n_f32(0.0f);
|
|
82
|
+
float32x4_t d2 = vdupq_n_f32(0.0f);
|
|
83
|
+
float32x4_t d3 = vdupq_n_f32(0.0f);
|
|
84
|
+
float32x4_t s0 = vdupq_n_f32(0.0f);
|
|
85
|
+
float32x4_t s1 = vdupq_n_f32(0.0f);
|
|
86
|
+
float32x4_t s2 = vdupq_n_f32(0.0f);
|
|
87
|
+
float32x4_t s3 = vdupq_n_f32(0.0f);
|
|
88
|
+
for (; j + 15 < dim; j += 16) {
|
|
89
|
+
float32x4_t q0 = vld1q_f32(q + j);
|
|
90
|
+
float32x4_t r0 = vld1q_f32(row + j);
|
|
91
|
+
float32x4_t q1 = vld1q_f32(q + j + 4);
|
|
92
|
+
float32x4_t r1 = vld1q_f32(row + j + 4);
|
|
93
|
+
float32x4_t q2 = vld1q_f32(q + j + 8);
|
|
94
|
+
float32x4_t r2 = vld1q_f32(row + j + 8);
|
|
95
|
+
float32x4_t q3 = vld1q_f32(q + j + 12);
|
|
96
|
+
float32x4_t r3 = vld1q_f32(row + j + 12);
|
|
97
|
+
d0 = vmlaq_f32(d0, q0, r0);
|
|
98
|
+
d1 = vmlaq_f32(d1, q1, r1);
|
|
99
|
+
d2 = vmlaq_f32(d2, q2, r2);
|
|
100
|
+
d3 = vmlaq_f32(d3, q3, r3);
|
|
101
|
+
s0 = vmlaq_f32(s0, r0, r0);
|
|
102
|
+
s1 = vmlaq_f32(s1, r1, r1);
|
|
103
|
+
s2 = vmlaq_f32(s2, r2, r2);
|
|
104
|
+
s3 = vmlaq_f32(s3, r3, r3);
|
|
105
|
+
}
|
|
106
|
+
float32x4_t dotv = vaddq_f32(vaddq_f32(d0, d1), vaddq_f32(d2, d3));
|
|
107
|
+
float32x4_t sqv = vaddq_f32(vaddq_f32(s0, s1), vaddq_f32(s2, s3));
|
|
108
|
+
#if defined(__aarch64__)
|
|
109
|
+
float dot = vaddvq_f32(dotv);
|
|
110
|
+
float row_sq = vaddvq_f32(sqv);
|
|
111
|
+
#else
|
|
112
|
+
float32x2_t pair = vadd_f32(vget_low_f32(dotv), vget_high_f32(dotv));
|
|
113
|
+
pair = vpadd_f32(pair, pair);
|
|
114
|
+
float dot = vget_lane_f32(pair, 0);
|
|
115
|
+
pair = vadd_f32(vget_low_f32(sqv), vget_high_f32(sqv));
|
|
116
|
+
pair = vpadd_f32(pair, pair);
|
|
117
|
+
float row_sq = vget_lane_f32(pair, 0);
|
|
118
|
+
#endif
|
|
119
|
+
#elif defined(SE_HAVE_SSE)
|
|
120
|
+
__m128 d0 = _mm_setzero_ps();
|
|
121
|
+
__m128 d1 = _mm_setzero_ps();
|
|
122
|
+
__m128 d2 = _mm_setzero_ps();
|
|
123
|
+
__m128 d3 = _mm_setzero_ps();
|
|
124
|
+
__m128 s0 = _mm_setzero_ps();
|
|
125
|
+
__m128 s1 = _mm_setzero_ps();
|
|
126
|
+
__m128 s2 = _mm_setzero_ps();
|
|
127
|
+
__m128 s3 = _mm_setzero_ps();
|
|
128
|
+
for (; j + 15 < dim; j += 16) {
|
|
129
|
+
__m128 q0 = _mm_loadu_ps(q + j);
|
|
130
|
+
__m128 r0 = _mm_loadu_ps(row + j);
|
|
131
|
+
__m128 q1 = _mm_loadu_ps(q + j + 4);
|
|
132
|
+
__m128 r1 = _mm_loadu_ps(row + j + 4);
|
|
133
|
+
__m128 q2 = _mm_loadu_ps(q + j + 8);
|
|
134
|
+
__m128 r2 = _mm_loadu_ps(row + j + 8);
|
|
135
|
+
__m128 q3 = _mm_loadu_ps(q + j + 12);
|
|
136
|
+
__m128 r3 = _mm_loadu_ps(row + j + 12);
|
|
137
|
+
d0 = _mm_add_ps(d0, _mm_mul_ps(q0, r0));
|
|
138
|
+
d1 = _mm_add_ps(d1, _mm_mul_ps(q1, r1));
|
|
139
|
+
d2 = _mm_add_ps(d2, _mm_mul_ps(q2, r2));
|
|
140
|
+
d3 = _mm_add_ps(d3, _mm_mul_ps(q3, r3));
|
|
141
|
+
s0 = _mm_add_ps(s0, _mm_mul_ps(r0, r0));
|
|
142
|
+
s1 = _mm_add_ps(s1, _mm_mul_ps(r1, r1));
|
|
143
|
+
s2 = _mm_add_ps(s2, _mm_mul_ps(r2, r2));
|
|
144
|
+
s3 = _mm_add_ps(s3, _mm_mul_ps(r3, r3));
|
|
145
|
+
}
|
|
146
|
+
__m128 dotv = _mm_add_ps(_mm_add_ps(d0, d1), _mm_add_ps(d2, d3));
|
|
147
|
+
__m128 sqv = _mm_add_ps(_mm_add_ps(s0, s1), _mm_add_ps(s2, s3));
|
|
148
|
+
float tmp[4];
|
|
149
|
+
_mm_storeu_ps(tmp, dotv);
|
|
150
|
+
float dot = (tmp[0] + tmp[1]) + (tmp[2] + tmp[3]);
|
|
151
|
+
_mm_storeu_ps(tmp, sqv);
|
|
152
|
+
float row_sq = (tmp[0] + tmp[1]) + (tmp[2] + tmp[3]);
|
|
153
|
+
#else
|
|
154
|
+
float d0 = 0.0f;
|
|
155
|
+
float d1 = 0.0f;
|
|
156
|
+
float d2 = 0.0f;
|
|
157
|
+
float d3 = 0.0f;
|
|
158
|
+
float s0 = 0.0f;
|
|
159
|
+
float s1 = 0.0f;
|
|
160
|
+
float s2 = 0.0f;
|
|
161
|
+
float s3 = 0.0f;
|
|
162
|
+
for (; j + 3 < dim; j += 4) {
|
|
163
|
+
float r0 = row[j];
|
|
164
|
+
float r1 = row[j + 1];
|
|
165
|
+
float r2 = row[j + 2];
|
|
166
|
+
float r3 = row[j + 3];
|
|
167
|
+
d0 += q[j] * r0;
|
|
168
|
+
d1 += q[j + 1] * r1;
|
|
169
|
+
d2 += q[j + 2] * r2;
|
|
170
|
+
d3 += q[j + 3] * r3;
|
|
171
|
+
s0 += r0 * r0;
|
|
172
|
+
s1 += r1 * r1;
|
|
173
|
+
s2 += r2 * r2;
|
|
174
|
+
s3 += r3 * r3;
|
|
175
|
+
}
|
|
176
|
+
float dot = (d0 + d1) + (d2 + d3);
|
|
177
|
+
float row_sq = (s0 + s1) + (s2 + s3);
|
|
178
|
+
#endif
|
|
179
|
+
for (; j < dim; j++) {
|
|
180
|
+
float r = row[j];
|
|
181
|
+
dot += q[j] * r;
|
|
182
|
+
row_sq += r * r;
|
|
183
|
+
}
|
|
184
|
+
*row_sq_out = row_sq;
|
|
185
|
+
return dot;
|
|
186
|
+
}
|
|
187
|
+
|
|
188
|
+
static float cosine_score(float dot, float row_sq, float inv_query_norm) {
|
|
189
|
+
if (row_sq > 0.0f)
|
|
190
|
+
return dot * inv_query_norm / sqrtf(row_sq);
|
|
191
|
+
if (row_sq == 0.0f)
|
|
192
|
+
return 0.0f;
|
|
193
|
+
return NAN;
|
|
194
|
+
}
|
|
195
|
+
|
|
196
|
+
void *se_topk_execute(void *arg) {
|
|
197
|
+
se_topk_job_t *job = (se_topk_job_t *)arg;
|
|
198
|
+
for (size_t r = 0; r < job->rows; r++) {
|
|
199
|
+
if ((r & 1023u) == 0 && job->cancelled)
|
|
200
|
+
return NULL;
|
|
201
|
+
float score;
|
|
202
|
+
|
|
203
|
+
if (job->format == SE_VECTOR_FORMAT_F16) {
|
|
204
|
+
const uint8_t *row = (const uint8_t *)job->m + r * job->dim * 2u;
|
|
205
|
+
if (job->cosine) {
|
|
206
|
+
float row_sq = 0.0f;
|
|
207
|
+
score = se_dot_and_row_sq_f16(job->q, row, job->dim, &row_sq);
|
|
208
|
+
score = cosine_score(score, row_sq, job->inv_query_norm);
|
|
209
|
+
} else {
|
|
210
|
+
score = se_dot_product_f16(job->q, row, job->dim);
|
|
211
|
+
}
|
|
212
|
+
} else {
|
|
213
|
+
const float *row = (const float *)job->m + r * job->dim;
|
|
214
|
+
if (job->cosine) {
|
|
215
|
+
float row_sq = 0.0f;
|
|
216
|
+
score = se_dot_and_row_sq_f32(job->q, row, job->dim, &row_sq);
|
|
217
|
+
score = cosine_score(score, row_sq, job->inv_query_norm);
|
|
218
|
+
} else {
|
|
219
|
+
score = se_dot_product_f32(job->q, row, job->dim);
|
|
220
|
+
}
|
|
221
|
+
}
|
|
222
|
+
|
|
223
|
+
if (!(score > job->best_score[job->k - 1]))
|
|
224
|
+
continue;
|
|
225
|
+
|
|
226
|
+
long pos = job->k - 1;
|
|
227
|
+
while (pos > 0 && job->best_score[pos - 1] < score) {
|
|
228
|
+
job->best_score[pos] = job->best_score[pos - 1];
|
|
229
|
+
job->best_idx[pos] = job->best_idx[pos - 1];
|
|
230
|
+
pos--;
|
|
231
|
+
}
|
|
232
|
+
job->best_score[pos] = score;
|
|
233
|
+
job->best_idx[pos] = r;
|
|
234
|
+
}
|
|
235
|
+
return NULL;
|
|
236
|
+
}
|