stride-align 0.6.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 (190) hide show
  1. checksums.yaml +7 -0
  2. data/LICENSE +201 -0
  3. data/NOTICE +255 -0
  4. data/README.md +107 -0
  5. data/data/bmpm_data/gen_approx_any.txt +131 -0
  6. data/data/bmpm_data/gen_approx_arabic.txt +26 -0
  7. data/data/bmpm_data/gen_approx_common.txt +233 -0
  8. data/data/bmpm_data/gen_approx_cyrillic.txt +18 -0
  9. data/data/bmpm_data/gen_approx_czech.txt +18 -0
  10. data/data/bmpm_data/gen_approx_dutch.txt +18 -0
  11. data/data/bmpm_data/gen_approx_english.txt +47 -0
  12. data/data/bmpm_data/gen_approx_french.txt +25 -0
  13. data/data/bmpm_data/gen_approx_german.txt +73 -0
  14. data/data/bmpm_data/gen_approx_greek.txt +18 -0
  15. data/data/bmpm_data/gen_approx_greeklatin.txt +20 -0
  16. data/data/bmpm_data/gen_approx_hebrew.txt +18 -0
  17. data/data/bmpm_data/gen_approx_hungarian.txt +18 -0
  18. data/data/bmpm_data/gen_approx_italian.txt +18 -0
  19. data/data/bmpm_data/gen_approx_polish.txt +84 -0
  20. data/data/bmpm_data/gen_approx_portuguese.txt +18 -0
  21. data/data/bmpm_data/gen_approx_romanian.txt +18 -0
  22. data/data/bmpm_data/gen_approx_russian.txt +48 -0
  23. data/data/bmpm_data/gen_approx_spanish.txt +21 -0
  24. data/data/bmpm_data/gen_approx_turkish.txt +18 -0
  25. data/data/bmpm_data/gen_exact_any.txt +40 -0
  26. data/data/bmpm_data/gen_exact_approx_common.txt +79 -0
  27. data/data/bmpm_data/gen_exact_arabic.txt +18 -0
  28. data/data/bmpm_data/gen_exact_common.txt +32 -0
  29. data/data/bmpm_data/gen_exact_cyrillic.txt +18 -0
  30. data/data/bmpm_data/gen_exact_czech.txt +18 -0
  31. data/data/bmpm_data/gen_exact_dutch.txt +18 -0
  32. data/data/bmpm_data/gen_exact_english.txt +18 -0
  33. data/data/bmpm_data/gen_exact_french.txt +18 -0
  34. data/data/bmpm_data/gen_exact_german.txt +18 -0
  35. data/data/bmpm_data/gen_exact_greek.txt +18 -0
  36. data/data/bmpm_data/gen_exact_greeklatin.txt +18 -0
  37. data/data/bmpm_data/gen_exact_hebrew.txt +18 -0
  38. data/data/bmpm_data/gen_exact_hungarian.txt +18 -0
  39. data/data/bmpm_data/gen_exact_italian.txt +18 -0
  40. data/data/bmpm_data/gen_exact_polish.txt +23 -0
  41. data/data/bmpm_data/gen_exact_portuguese.txt +18 -0
  42. data/data/bmpm_data/gen_exact_romanian.txt +18 -0
  43. data/data/bmpm_data/gen_exact_russian.txt +19 -0
  44. data/data/bmpm_data/gen_exact_spanish.txt +19 -0
  45. data/data/bmpm_data/gen_exact_turkish.txt +18 -0
  46. data/data/bmpm_data/gen_hebrew_common.txt +113 -0
  47. data/data/bmpm_data/gen_lang.txt +295 -0
  48. data/data/bmpm_data/gen_languages.txt +36 -0
  49. data/data/bmpm_data/gen_rules_any.txt +367 -0
  50. data/data/bmpm_data/gen_rules_arabic.txt +76 -0
  51. data/data/bmpm_data/gen_rules_cyrillic.txt +99 -0
  52. data/data/bmpm_data/gen_rules_czech.txt +67 -0
  53. data/data/bmpm_data/gen_rules_dutch.txt +78 -0
  54. data/data/bmpm_data/gen_rules_english.txt +113 -0
  55. data/data/bmpm_data/gen_rules_french.txt +114 -0
  56. data/data/bmpm_data/gen_rules_german.txt +129 -0
  57. data/data/bmpm_data/gen_rules_greek.txt +97 -0
  58. data/data/bmpm_data/gen_rules_greeklatin.txt +118 -0
  59. data/data/bmpm_data/gen_rules_hebrew.txt +62 -0
  60. data/data/bmpm_data/gen_rules_hungarian.txt +83 -0
  61. data/data/bmpm_data/gen_rules_italian.txt +77 -0
  62. data/data/bmpm_data/gen_rules_polish.txt +185 -0
  63. data/data/bmpm_data/gen_rules_portuguese.txt +105 -0
  64. data/data/bmpm_data/gen_rules_romanian.txt +64 -0
  65. data/data/bmpm_data/gen_rules_russian.txt +142 -0
  66. data/data/bmpm_data/gen_rules_spanish.txt +85 -0
  67. data/data/bmpm_data/gen_rules_turkish.txt +50 -0
  68. data/data/keyboard_data/qwerty.npy +0 -0
  69. data/data/matrix_data/BLOSUM100 +31 -0
  70. data/data/matrix_data/BLOSUM30 +31 -0
  71. data/data/matrix_data/BLOSUM35 +31 -0
  72. data/data/matrix_data/BLOSUM40 +31 -0
  73. data/data/matrix_data/BLOSUM45 +25 -0
  74. data/data/matrix_data/BLOSUM50 +25 -0
  75. data/data/matrix_data/BLOSUM55 +31 -0
  76. data/data/matrix_data/BLOSUM60 +31 -0
  77. data/data/matrix_data/BLOSUM62 +25 -0
  78. data/data/matrix_data/BLOSUM65 +31 -0
  79. data/data/matrix_data/BLOSUM70 +31 -0
  80. data/data/matrix_data/BLOSUM75 +31 -0
  81. data/data/matrix_data/BLOSUM80 +25 -0
  82. data/data/matrix_data/BLOSUM85 +31 -0
  83. data/data/matrix_data/BLOSUM90 +25 -0
  84. data/data/matrix_data/NUC.4.4 +25 -0
  85. data/data/matrix_data/PAM10 +34 -0
  86. data/data/matrix_data/PAM100 +34 -0
  87. data/data/matrix_data/PAM110 +34 -0
  88. data/data/matrix_data/PAM120 +34 -0
  89. data/data/matrix_data/PAM130 +34 -0
  90. data/data/matrix_data/PAM140 +34 -0
  91. data/data/matrix_data/PAM150 +34 -0
  92. data/data/matrix_data/PAM160 +34 -0
  93. data/data/matrix_data/PAM170 +34 -0
  94. data/data/matrix_data/PAM180 +34 -0
  95. data/data/matrix_data/PAM190 +34 -0
  96. data/data/matrix_data/PAM20 +34 -0
  97. data/data/matrix_data/PAM200 +34 -0
  98. data/data/matrix_data/PAM210 +34 -0
  99. data/data/matrix_data/PAM220 +34 -0
  100. data/data/matrix_data/PAM230 +34 -0
  101. data/data/matrix_data/PAM240 +34 -0
  102. data/data/matrix_data/PAM250 +25 -0
  103. data/data/matrix_data/PAM260 +34 -0
  104. data/data/matrix_data/PAM270 +34 -0
  105. data/data/matrix_data/PAM280 +34 -0
  106. data/data/matrix_data/PAM290 +34 -0
  107. data/data/matrix_data/PAM30 +25 -0
  108. data/data/matrix_data/PAM300 +34 -0
  109. data/data/matrix_data/PAM310 +34 -0
  110. data/data/matrix_data/PAM320 +34 -0
  111. data/data/matrix_data/PAM330 +34 -0
  112. data/data/matrix_data/PAM340 +34 -0
  113. data/data/matrix_data/PAM350 +34 -0
  114. data/data/matrix_data/PAM360 +34 -0
  115. data/data/matrix_data/PAM370 +34 -0
  116. data/data/matrix_data/PAM380 +34 -0
  117. data/data/matrix_data/PAM390 +34 -0
  118. data/data/matrix_data/PAM40 +34 -0
  119. data/data/matrix_data/PAM400 +34 -0
  120. data/data/matrix_data/PAM410 +34 -0
  121. data/data/matrix_data/PAM420 +34 -0
  122. data/data/matrix_data/PAM430 +34 -0
  123. data/data/matrix_data/PAM440 +34 -0
  124. data/data/matrix_data/PAM450 +34 -0
  125. data/data/matrix_data/PAM460 +34 -0
  126. data/data/matrix_data/PAM470 +34 -0
  127. data/data/matrix_data/PAM480 +34 -0
  128. data/data/matrix_data/PAM490 +34 -0
  129. data/data/matrix_data/PAM50 +34 -0
  130. data/data/matrix_data/PAM500 +34 -0
  131. data/data/matrix_data/PAM60 +34 -0
  132. data/data/matrix_data/PAM70 +25 -0
  133. data/data/matrix_data/PAM80 +34 -0
  134. data/data/matrix_data/PAM90 +34 -0
  135. data/ext/stride_align/backend_avx2.cpp +2 -0
  136. data/ext/stride_align/backend_avx512bwvl.cpp +2 -0
  137. data/ext/stride_align/backend_generic.cpp +3 -0
  138. data/ext/stride_align/backend_impl.hpp +983 -0
  139. data/ext/stride_align/backend_lasx.cpp +2 -0
  140. data/ext/stride_align/backend_lsx.cpp +2 -0
  141. data/ext/stride_align/backend_neon.cpp +2 -0
  142. data/ext/stride_align/backend_rvv.cpp +2 -0
  143. data/ext/stride_align/backend_sse41.cpp +2 -0
  144. data/ext/stride_align/backend_sve.cpp +2 -0
  145. data/ext/stride_align/backend_sve2.cpp +2 -0
  146. data/ext/stride_align/backend_vsx.cpp +2 -0
  147. data/ext/stride_align/beider_morse_impl.cpp +5 -0
  148. data/ext/stride_align/cpu_detect.cpp +172 -0
  149. data/ext/stride_align/cpu_detect.hpp +6 -0
  150. data/ext/stride_align/extconf.rb +114 -0
  151. data/ext/stride_align/target_profile.hpp +82 -0
  152. data/ext/stride_align/vendor/beider_morse_impl.cpp +1467 -0
  153. data/ext/stride_align/vendor/stride_align/alignment.hpp +199 -0
  154. data/ext/stride_align/vendor/stride_align/batch.hpp +812 -0
  155. data/ext/stride_align/vendor/stride_align/beider_morse.hpp +121 -0
  156. data/ext/stride_align/vendor/stride_align/caverphone.hpp +222 -0
  157. data/ext/stride_align/vendor/stride_align/cologne_phonetic.hpp +202 -0
  158. data/ext/stride_align/vendor/stride_align/core.hpp +731 -0
  159. data/ext/stride_align/vendor/stride_align/daitch_mokotoff.hpp +631 -0
  160. data/ext/stride_align/vendor/stride_align/double_metaphone.hpp +796 -0
  161. data/ext/stride_align/vendor/stride_align/dtw.hpp +300 -0
  162. data/ext/stride_align/vendor/stride_align/encoded.hpp +235 -0
  163. data/ext/stride_align/vendor/stride_align/hamming.hpp +55 -0
  164. data/ext/stride_align/vendor/stride_align/indel.hpp +1200 -0
  165. data/ext/stride_align/vendor/stride_align/jaro.hpp +517 -0
  166. data/ext/stride_align/vendor/stride_align/lcs.hpp +159 -0
  167. data/ext/stride_align/vendor/stride_align/levenshtein.hpp +1247 -0
  168. data/ext/stride_align/vendor/stride_align/levenshtein_prepared.hpp +193 -0
  169. data/ext/stride_align/vendor/stride_align/match_rating.hpp +168 -0
  170. data/ext/stride_align/vendor/stride_align/metaphone.hpp +291 -0
  171. data/ext/stride_align/vendor/stride_align/ngram.hpp +176 -0
  172. data/ext/stride_align/vendor/stride_align/nysiis.hpp +199 -0
  173. data/ext/stride_align/vendor/stride_align/pairwise_alignment.hpp +465 -0
  174. data/ext/stride_align/vendor/stride_align/partial_ratio.hpp +486 -0
  175. data/ext/stride_align/vendor/stride_align/ratcliff_obershelp.hpp +101 -0
  176. data/ext/stride_align/vendor/stride_align/soundex.hpp +108 -0
  177. data/ext/stride_align/vendor/stride_align/token_ratios.hpp +445 -0
  178. data/ext/stride_align/vendor/stride_align/types.hpp +16 -0
  179. data/ext/stride_align/vendor/stride_align/utf8.hpp +512 -0
  180. data/ext/stride_align/vendor/stride_align/wratio.hpp +363 -0
  181. data/lib/stride_align/algorithms.rb +296 -0
  182. data/lib/stride_align/alignment_path.rb +217 -0
  183. data/lib/stride_align/backend.rb +47 -0
  184. data/lib/stride_align/batch.rb +705 -0
  185. data/lib/stride_align/core.rb +180 -0
  186. data/lib/stride_align/keyboard.rb +200 -0
  187. data/lib/stride_align/matrices.rb +403 -0
  188. data/lib/stride_align/version.rb +5 -0
  189. data/lib/stride_align.rb +87 -0
  190. metadata +231 -0
@@ -0,0 +1,512 @@
1
+ #pragma once
2
+
3
+ #include <algorithm>
4
+ #include <array>
5
+ #include <bit>
6
+ #include <cstddef>
7
+ #include <cstdint>
8
+ #include <cstring>
9
+ #include <limits>
10
+ #include <span>
11
+ #include <stdexcept>
12
+ #include <string>
13
+ #include <string_view>
14
+ #include <type_traits>
15
+ #include <utility>
16
+ #include <variant>
17
+ #include <vector>
18
+
19
+ #if defined(__AVX512BW__)
20
+ #include <immintrin.h>
21
+ #elif defined(__AVX2__)
22
+ #include <immintrin.h>
23
+ #elif defined(__SSE2__)
24
+ #include <emmintrin.h>
25
+ #endif
26
+
27
+ #if defined(__ARM_NEON) || defined(__ARM_NEON__)
28
+ #include <arm_neon.h>
29
+ #endif
30
+
31
+ #if defined(__loongarch_asx)
32
+ #include <lsxintrin.h>
33
+ #include <lasxintrin.h>
34
+ #elif defined(__loongarch_sx)
35
+ #include <lsxintrin.h>
36
+ #endif
37
+
38
+ namespace stride_align::utf8 {
39
+
40
+ enum class TokenWidth : std::uint8_t {
41
+ u8 = 1,
42
+ u16 = 2,
43
+ u32 = 4,
44
+ };
45
+
46
+ enum class PreparationMode : std::uint8_t {
47
+ pair,
48
+ streaming,
49
+ };
50
+
51
+ enum class NonAsciiPolicy : std::uint8_t {
52
+ // Decode directly to Unicode code points. Pair preparation uses UCS-2 when
53
+ // both strings fit the BMP, otherwise UCS-4.
54
+ fixed_width,
55
+ // For sufficiently long pairs, remap the combined alphabet to dense tokens
56
+ // using the narrowest supported width. Equality is preserved exactly.
57
+ pack_long,
58
+ };
59
+
60
+ struct PreparationOptions {
61
+ PreparationMode mode = PreparationMode::pair;
62
+ NonAsciiPolicy non_ascii = NonAsciiPolicy::pack_long;
63
+ std::size_t pack_threshold = 64;
64
+ };
65
+
66
+ using TokenBuffer = std::variant<
67
+ std::span<const std::uint8_t>,
68
+ std::vector<std::uint8_t>,
69
+ std::span<const std::uint16_t>,
70
+ std::vector<std::uint16_t>,
71
+ std::span<const std::uint32_t>,
72
+ std::vector<std::uint32_t>>;
73
+
74
+ struct PreparedPair {
75
+ TokenWidth width = TokenWidth::u8;
76
+ bool borrowed_ascii = true;
77
+ bool packed = false;
78
+ TokenBuffer query;
79
+ TokenBuffer target;
80
+
81
+ // Borrowed fast paths store non-owning spans. This is ASCII for the UTF-8
82
+ // pair adapter; native single-byte host adapters may also borrow high-bit
83
+ // characters. Streaming host adapters can borrow stable UCS-2/UCS-4
84
+ // buffers after promoting their inputs once. Every source view must outlive
85
+ // every consumer of this pair.
86
+
87
+ std::size_t query_size() const noexcept {
88
+ return std::visit([](const auto& value) { return value.size(); }, query);
89
+ }
90
+
91
+ std::size_t target_size() const noexcept {
92
+ return std::visit([](const auto& value) { return value.size(); }, target);
93
+ }
94
+ };
95
+
96
+ class InvalidUtf8 : public std::invalid_argument {
97
+ public:
98
+ InvalidUtf8(std::size_t offset, const char* reason)
99
+ : std::invalid_argument(
100
+ "invalid UTF-8 at byte " + std::to_string(offset) + ": " + reason),
101
+ offset_(offset) {}
102
+
103
+ std::size_t offset() const noexcept { return offset_; }
104
+
105
+ private:
106
+ std::size_t offset_;
107
+ };
108
+
109
+ inline bool is_ascii(std::span<const std::uint8_t> input) noexcept {
110
+ const auto* data = input.data();
111
+ std::size_t index = 0;
112
+
113
+ #if defined(__AVX512BW__)
114
+ for (; index + 64U <= input.size(); index += 64U) {
115
+ const __m512i value = _mm512_loadu_si512(
116
+ reinterpret_cast<const void*>(data + index));
117
+ if (_mm512_movepi8_mask(value) != 0U) {
118
+ return false;
119
+ }
120
+ }
121
+ #elif defined(__AVX2__)
122
+ for (; index + 32U <= input.size(); index += 32U) {
123
+ const __m256i value = _mm256_loadu_si256(
124
+ reinterpret_cast<const __m256i*>(data + index));
125
+ if (_mm256_movemask_epi8(value) != 0) {
126
+ return false;
127
+ }
128
+ }
129
+ #elif defined(__SSE2__)
130
+ for (; index + 16U <= input.size(); index += 16U) {
131
+ const __m128i value = _mm_loadu_si128(
132
+ reinterpret_cast<const __m128i*>(data + index));
133
+ if (_mm_movemask_epi8(value) != 0) {
134
+ return false;
135
+ }
136
+ }
137
+ #elif defined(__ARM_NEON) || defined(__ARM_NEON__)
138
+ for (; index + 16U <= input.size(); index += 16U) {
139
+ const uint8x16_t value = vld1q_u8(data + index);
140
+ const uint8x16_t high_bits = vandq_u8(value, vdupq_n_u8(0x80U));
141
+ #if defined(__aarch64__)
142
+ if (vmaxvq_u8(high_bits) != 0U) {
143
+ return false;
144
+ }
145
+ #else
146
+ const uint64x2_t words = vreinterpretq_u64_u8(high_bits);
147
+ if ((vgetq_lane_u64(words, 0) | vgetq_lane_u64(words, 1)) != 0U) {
148
+ return false;
149
+ }
150
+ #endif
151
+ }
152
+ #elif defined(__loongarch_asx)
153
+ for (; index + 32U <= input.size(); index += 32U) {
154
+ const __m256i value = __lasx_xvld(
155
+ const_cast<void*>(static_cast<const void*>(data + index)), 0);
156
+ const __m256i high_bits = __lasx_xvmskltz_b(value);
157
+ if ((__lasx_xvpickve2gr_du(high_bits, 0) |
158
+ __lasx_xvpickve2gr_du(high_bits, 1) |
159
+ __lasx_xvpickve2gr_du(high_bits, 2) |
160
+ __lasx_xvpickve2gr_du(high_bits, 3)) != 0UL) {
161
+ return false;
162
+ }
163
+ }
164
+ #elif defined(__loongarch_sx)
165
+ for (; index + 16U <= input.size(); index += 16U) {
166
+ const __m128i value = __lsx_vld(
167
+ const_cast<void*>(static_cast<const void*>(data + index)), 0);
168
+ const __m128i high_bits = __lsx_vmskltz_b(value);
169
+ if ((__lsx_vpickve2gr_du(high_bits, 0) |
170
+ __lsx_vpickve2gr_du(high_bits, 1)) != 0UL) {
171
+ return false;
172
+ }
173
+ }
174
+ #endif
175
+
176
+ constexpr std::size_t kWordBytes = sizeof(std::size_t);
177
+ constexpr std::size_t kHighBits = std::numeric_limits<std::size_t>::max() / 0xffU * 0x80U;
178
+ for (; index + kWordBytes <= input.size(); index += kWordBytes) {
179
+ std::size_t word = 0;
180
+ std::memcpy(&word, data + index, kWordBytes);
181
+ if ((word & kHighBits) != 0U) {
182
+ return false;
183
+ }
184
+ }
185
+ for (; index < input.size(); ++index) {
186
+ if ((data[index] & 0x80U) != 0U) {
187
+ return false;
188
+ }
189
+ }
190
+ return true;
191
+ }
192
+
193
+ inline bool is_ascii(std::string_view input) noexcept {
194
+ return is_ascii(std::span<const std::uint8_t>(
195
+ reinterpret_cast<const std::uint8_t*>(input.data()), input.size()));
196
+ }
197
+
198
+ namespace detail {
199
+
200
+ inline bool continuation(std::uint8_t value) noexcept {
201
+ return (value & 0xc0U) == 0x80U;
202
+ }
203
+
204
+ inline std::vector<std::uint32_t> decode(std::string_view input) {
205
+ const auto* data = reinterpret_cast<const std::uint8_t*>(input.data());
206
+ std::vector<std::uint32_t> output;
207
+ output.reserve(input.size());
208
+
209
+ std::size_t index = 0;
210
+ while (index < input.size()) {
211
+ const std::uint8_t first = data[index];
212
+ if (first < 0x80U) {
213
+ output.push_back(first);
214
+ ++index;
215
+ continue;
216
+ }
217
+
218
+ std::uint32_t codepoint = 0;
219
+ std::size_t length = 0;
220
+ if (first >= 0xc2U && first <= 0xdfU) {
221
+ codepoint = first & 0x1fU;
222
+ length = 2U;
223
+ } else if (first >= 0xe0U && first <= 0xefU) {
224
+ codepoint = first & 0x0fU;
225
+ length = 3U;
226
+ } else if (first >= 0xf0U && first <= 0xf4U) {
227
+ codepoint = first & 0x07U;
228
+ length = 4U;
229
+ } else {
230
+ throw InvalidUtf8(index, "invalid leading byte");
231
+ }
232
+
233
+ if (index + length > input.size()) {
234
+ throw InvalidUtf8(index, "truncated sequence");
235
+ }
236
+ for (std::size_t part = 1; part < length; ++part) {
237
+ if (!continuation(data[index + part])) {
238
+ throw InvalidUtf8(index + part, "invalid continuation byte");
239
+ }
240
+ codepoint = (codepoint << 6U) | (data[index + part] & 0x3fU);
241
+ }
242
+
243
+ if ((length == 3U && codepoint < 0x800U) ||
244
+ (length == 4U && codepoint < 0x10000U)) {
245
+ throw InvalidUtf8(index, "overlong sequence");
246
+ }
247
+ if (codepoint >= 0xd800U && codepoint <= 0xdfffU) {
248
+ throw InvalidUtf8(index, "surrogate code point");
249
+ }
250
+ if (codepoint > 0x10ffffU) {
251
+ throw InvalidUtf8(index, "code point above U+10FFFF");
252
+ }
253
+
254
+ output.push_back(codepoint);
255
+ index += length;
256
+ }
257
+ return output;
258
+ }
259
+
260
+ inline std::size_t table_capacity(std::size_t expected) {
261
+ constexpr std::size_t kHighestPowerOfTwo =
262
+ std::size_t{1} << (std::numeric_limits<std::size_t>::digits - 1U);
263
+ if (expected > (kHighestPowerOfTwo - 1U) / 2U) {
264
+ throw std::length_error("Swiss table capacity overflow");
265
+ }
266
+ return std::bit_ceil(std::max<std::size_t>(expected * 2U + 1U, 16U));
267
+ }
268
+
269
+ inline std::uint32_t control_mask(
270
+ const std::uint8_t* control,
271
+ std::uint8_t needle) noexcept {
272
+ #if defined(__SSE2__)
273
+ const __m128i values = _mm_loadu_si128(
274
+ reinterpret_cast<const __m128i*>(control));
275
+ const __m128i matches = _mm_cmpeq_epi8(
276
+ values, _mm_set1_epi8(static_cast<char>(needle)));
277
+ return static_cast<std::uint32_t>(_mm_movemask_epi8(matches));
278
+ #elif defined(__ARM_NEON) || defined(__ARM_NEON__)
279
+ const uint8x16_t matches = vceqq_u8(vld1q_u8(control), vdupq_n_u8(needle));
280
+ #if defined(__aarch64__)
281
+ static constexpr std::array<std::uint8_t, 16> kBits{
282
+ 1U, 2U, 4U, 8U, 16U, 32U, 64U, 128U,
283
+ 1U, 2U, 4U, 8U, 16U, 32U, 64U, 128U};
284
+ const uint8x16_t weighted = vandq_u8(matches, vld1q_u8(kBits.data()));
285
+ const std::uint32_t low = vaddv_u8(vget_low_u8(weighted));
286
+ const std::uint32_t high = vaddv_u8(vget_high_u8(weighted));
287
+ return low | (high << 8U);
288
+ #else
289
+ std::array<std::uint8_t, 16> lanes{};
290
+ vst1q_u8(lanes.data(), matches);
291
+ std::uint32_t mask = 0;
292
+ for (std::size_t lane = 0; lane < lanes.size(); ++lane) {
293
+ mask |= static_cast<std::uint32_t>(lanes[lane] != 0U) << lane;
294
+ }
295
+ return mask;
296
+ #endif
297
+ #elif defined(__loongarch_sx)
298
+ const __m128i values = __lsx_vld(
299
+ const_cast<void*>(static_cast<const void*>(control)), 0);
300
+ const __m128i matches = __lsx_vseq_b(
301
+ values, __lsx_vreplgr2vr_b(static_cast<int>(needle)));
302
+ const __m128i mask = __lsx_vmskltz_b(matches);
303
+ return static_cast<std::uint32_t>(__lsx_vpickve2gr_du(mask, 0));
304
+ #else
305
+ std::uint32_t mask = 0;
306
+ for (std::size_t lane = 0; lane < 16U; ++lane) {
307
+ mask |= static_cast<std::uint32_t>(control[lane] == needle) << lane;
308
+ }
309
+ return mask;
310
+ #endif
311
+ }
312
+
313
+ template <typename Token>
314
+ class SwissMap {
315
+ public:
316
+ explicit SwissMap(std::size_t expected) {
317
+ const std::size_t capacity = table_capacity(expected);
318
+ control_.assign(capacity, kEmpty);
319
+ keys_.resize(capacity);
320
+ values_.resize(capacity);
321
+ mask_ = capacity - 1U;
322
+ }
323
+
324
+ Token intern(std::uint32_t key) {
325
+ const std::uint64_t hash = mix(key);
326
+ std::size_t group =
327
+ (static_cast<std::size_t>(hash) & mask_) & ~(kGroupWidth - 1U);
328
+ const std::uint8_t fingerprint = static_cast<std::uint8_t>((hash >> 57U) & 0x7fU);
329
+ for (;;) {
330
+ std::uint32_t matches = control_mask(control_.data() + group, fingerprint);
331
+ while (matches != 0U) {
332
+ const auto lane = static_cast<std::size_t>(std::countr_zero(matches));
333
+ const std::size_t slot = group + lane;
334
+ if (keys_[slot] == key) return values_[slot];
335
+ matches &= matches - 1U;
336
+ }
337
+
338
+ const std::uint32_t empty = control_mask(control_.data() + group, kEmpty);
339
+ if (empty != 0U) {
340
+ if (size_ >= static_cast<std::size_t>(std::numeric_limits<Token>::max())) {
341
+ throw std::length_error("combined alphabet does not fit requested token width");
342
+ }
343
+ const auto lane = static_cast<std::size_t>(std::countr_zero(empty));
344
+ const std::size_t slot = group + lane;
345
+ const Token value = static_cast<Token>(++size_);
346
+ control_[slot] = fingerprint;
347
+ keys_[slot] = key;
348
+ values_[slot] = value;
349
+ return value;
350
+ }
351
+ group = (group + kGroupWidth) & mask_;
352
+ }
353
+ }
354
+
355
+ std::size_t size() const noexcept { return size_; }
356
+
357
+ private:
358
+ static constexpr std::uint8_t kEmpty = 0x80U;
359
+ static constexpr std::size_t kGroupWidth = 16U;
360
+
361
+ static std::uint64_t mix(std::uint32_t value) noexcept {
362
+ std::uint64_t result = value;
363
+ result ^= result >> 16U;
364
+ result *= 0x7feb352dULL;
365
+ result ^= result >> 15U;
366
+ result *= 0x846ca68bULL;
367
+ result ^= result >> 16U;
368
+ return result;
369
+ }
370
+
371
+ std::vector<std::uint8_t> control_;
372
+ std::vector<std::uint32_t> keys_;
373
+ std::vector<Token> values_;
374
+ std::size_t mask_ = 0;
375
+ std::size_t size_ = 0;
376
+ };
377
+
378
+ struct PackedPair32 {
379
+ std::vector<std::uint32_t> query;
380
+ std::vector<std::uint32_t> target;
381
+ std::size_t distinct = 0;
382
+ };
383
+
384
+ inline PackedPair32 pack_pair(
385
+ std::span<const std::uint32_t> query,
386
+ std::span<const std::uint32_t> target) {
387
+ if (query.size() > std::numeric_limits<std::size_t>::max() - target.size()) {
388
+ throw std::length_error("combined string length overflow");
389
+ }
390
+ SwissMap<std::uint32_t> map(query.size() + target.size());
391
+ PackedPair32 packed;
392
+ packed.query.reserve(query.size());
393
+ packed.target.reserve(target.size());
394
+ for (const std::uint32_t codepoint : query) {
395
+ packed.query.push_back(map.intern(codepoint));
396
+ }
397
+ for (const std::uint32_t codepoint : target) {
398
+ packed.target.push_back(map.intern(codepoint));
399
+ }
400
+ packed.distinct = map.size();
401
+ return packed;
402
+ }
403
+
404
+ template <typename Token>
405
+ inline std::vector<Token> narrow(std::span<const std::uint32_t> input) {
406
+ std::vector<Token> output;
407
+ output.reserve(input.size());
408
+ for (const std::uint32_t codepoint : input) {
409
+ output.push_back(static_cast<Token>(codepoint));
410
+ }
411
+ return output;
412
+ }
413
+
414
+ } // namespace detail
415
+
416
+ inline PreparedPair prepare_pair(
417
+ std::string_view query,
418
+ std::string_view target,
419
+ PreparationOptions options = {}) {
420
+ const bool query_ascii = is_ascii(query);
421
+ const bool target_ascii = is_ascii(target);
422
+
423
+ if (options.mode == PreparationMode::streaming) {
424
+ const auto promote = [](std::string_view input, bool ascii) {
425
+ if (!ascii) return detail::decode(input);
426
+ std::vector<std::uint32_t> output;
427
+ output.reserve(input.size());
428
+ for (const unsigned char value : input) output.push_back(value);
429
+ return output;
430
+ };
431
+ return PreparedPair{
432
+ TokenWidth::u32,
433
+ false,
434
+ false,
435
+ promote(query, query_ascii),
436
+ promote(target, target_ascii)};
437
+ }
438
+
439
+ if (query_ascii && target_ascii) {
440
+ auto view = [](std::string_view input) {
441
+ return std::span<const std::uint8_t>(
442
+ reinterpret_cast<const std::uint8_t*>(input.data()), input.size());
443
+ };
444
+ return PreparedPair{
445
+ TokenWidth::u8, true, false, view(query), view(target)};
446
+ }
447
+
448
+ auto query_codepoints = detail::decode(query);
449
+ auto target_codepoints = detail::decode(target);
450
+
451
+ if (query_codepoints.size() >
452
+ std::numeric_limits<std::size_t>::max() - target_codepoints.size()) {
453
+ throw std::length_error("combined string length overflow");
454
+ }
455
+ const std::size_t total_length =
456
+ query_codepoints.size() + target_codepoints.size();
457
+ if (options.non_ascii == NonAsciiPolicy::pack_long &&
458
+ total_length >= options.pack_threshold) {
459
+ auto packed = detail::pack_pair(query_codepoints, target_codepoints);
460
+ if (packed.distinct <=
461
+ static_cast<std::size_t>(std::numeric_limits<std::uint8_t>::max())) {
462
+ return PreparedPair{
463
+ TokenWidth::u8, false, true,
464
+ detail::narrow<std::uint8_t>(packed.query),
465
+ detail::narrow<std::uint8_t>(packed.target)};
466
+ }
467
+ if (packed.distinct <=
468
+ static_cast<std::size_t>(std::numeric_limits<std::uint16_t>::max())) {
469
+ return PreparedPair{
470
+ TokenWidth::u16, false, true,
471
+ detail::narrow<std::uint16_t>(packed.query),
472
+ detail::narrow<std::uint16_t>(packed.target)};
473
+ }
474
+ return PreparedPair{
475
+ TokenWidth::u32, false, true,
476
+ std::move(packed.query), std::move(packed.target)};
477
+ }
478
+
479
+ const auto max_codepoint = [](const std::vector<std::uint32_t>& input) {
480
+ std::uint32_t result = 0;
481
+ for (const std::uint32_t value : input) result = std::max(result, value);
482
+ return result;
483
+ };
484
+ if (std::max(max_codepoint(query_codepoints), max_codepoint(target_codepoints)) <= 0xffffU) {
485
+ return PreparedPair{
486
+ TokenWidth::u16,
487
+ false,
488
+ false,
489
+ detail::narrow<std::uint16_t>(query_codepoints),
490
+ detail::narrow<std::uint16_t>(target_codepoints)};
491
+ }
492
+ return PreparedPair{
493
+ TokenWidth::u32,
494
+ false,
495
+ false,
496
+ std::move(query_codepoints),
497
+ std::move(target_codepoints)};
498
+ }
499
+
500
+ inline std::vector<std::uint32_t> prepare_streaming(std::string_view input) {
501
+ // Streaming callers deliberately pay a uniform UCS-4 representation cost:
502
+ // no later string can force already-prepared state to change token width.
503
+ if (is_ascii(input)) {
504
+ std::vector<std::uint32_t> output;
505
+ output.reserve(input.size());
506
+ for (const unsigned char value : input) output.push_back(value);
507
+ return output;
508
+ }
509
+ return detail::decode(input);
510
+ }
511
+
512
+ } // namespace stride_align::utf8