Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
193 changes: 185 additions & 8 deletions include/stringzilla/find/serial.h
Original file line number Diff line number Diff line change
Expand Up @@ -431,6 +431,186 @@ SZ_HELPER_NOINLINE sz_cptr_t sz_find_3byte_serial_(sz_cptr_t haystack, sz_size_t
return SZ_NULL_CHAR;
}

/**
* @brief Find the last occurrence of a @b two-character needle in an arbitrary length haystack.
* This implementation uses hardware-agnostic SWAR technique, to process 8 possible offsets at a time.
*/
SZ_HELPER_NOINLINE sz_cptr_t sz_rfind_2byte_serial_(sz_cptr_t haystack, sz_size_t haystack_length, sz_cptr_t needle,
sz_size_t needle_length) {

// This is an internal method, and the haystack is guaranteed to be at least 2 bytes long.
sz_assert_(haystack_length >= 2 && "The haystack is too short.");
sz_unused_(needle_length); //? We keep this argument only for `sz_rfind_t` signature compatibility.
sz_cptr_t last = haystack + haystack_length - 2;

// On big-endian systems, skip SWAR and use simple serial search
#if SZ_IS_BIG_ENDIAN_
for (; last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) == 2) return last;
return SZ_NULL_CHAR;
#endif

// Process the misaligned tail, to void UB on unaligned 64-bit loads.
#if !SZ_USE_MISALIGNED_LOADS
for (; ((sz_size_t)(last + 1) & 7ull) && last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) == 2) return last;
#endif

sz_u64_vec_t haystack_even_vec, haystack_odd_vec, needle_vec, matches_even_vec, matches_odd_vec;
needle_vec.u64 = 0;
needle_vec.u8s[0] = needle[0], needle_vec.u8s[1] = needle[1];
needle_vec.u64 *= 0x0001000100010001ull; // broadcast

// This code simulates hyper-scalar execution, analyzing 8 offsets `[last - 7, last]` at a time.
for (; last >= haystack + 7; last -= 8) {
sz_cptr_t const window = last - 7;
haystack_even_vec.u64 = *(sz_u64_t *)window;
haystack_odd_vec.u64 = (haystack_even_vec.u64 >> 8) | ((sz_u64_t)(sz_u8_t)window[8] << 56);
matches_even_vec = sz_u64_each_2byte_equal_(haystack_even_vec, needle_vec);
matches_odd_vec = sz_u64_each_2byte_equal_(haystack_odd_vec, needle_vec);
matches_even_vec.u64 >>= 8;
if (matches_even_vec.u64 + matches_odd_vec.u64) {
sz_u64_t match_indicators = matches_even_vec.u64 | matches_odd_vec.u64;
return window + (63 - sz_u64_clz(match_indicators)) / 8;
}
}

for (; last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) == 2) return last;
return SZ_NULL_CHAR;
}

/**
* @brief Find the last occurrence of a @b three-character needle in an arbitrary length haystack.
* This implementation uses hardware-agnostic SWAR technique, to process 8 possible offsets at a time.
*/
SZ_HELPER_NOINLINE sz_cptr_t sz_rfind_3byte_serial_(sz_cptr_t haystack, sz_size_t haystack_length, sz_cptr_t needle,
sz_size_t needle_length) {

// This is an internal method, and the haystack is guaranteed to be at least 3 bytes long.
sz_assert_(haystack_length >= 3 && "The haystack is too short.");
sz_unused_(needle_length); //? We keep this argument only for `sz_rfind_t` signature compatibility.
sz_cptr_t last = haystack + haystack_length - 3;

// On big-endian systems, skip SWAR and use simple serial search
#if SZ_IS_BIG_ENDIAN_
for (; last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) + (last[2] == needle[2]) == 3) return last;
return SZ_NULL_CHAR;
#endif

// Process the misaligned tail, to void UB on unaligned 64-bit loads.
#if !SZ_USE_MISALIGNED_LOADS
for (; ((sz_size_t)(last + 1) & 7ull) && last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) + (last[2] == needle[2]) == 3) return last;
#endif

sz_u64_vec_t haystack0_vec, haystack1_vec, haystack2_vec, haystack3_vec, haystack4_vec;
sz_u64_vec_t matches0_vec, matches1_vec, matches2_vec, matches3_vec, matches4_vec;
sz_u64_vec_t needle_vec;
needle_vec.u64 = 0;
needle_vec.u8s[0] = needle[0], needle_vec.u8s[1] = needle[1], needle_vec.u8s[2] = needle[2];
needle_vec.u64 *= 0x0000000001000001ull; // broadcast

// This code simulates hyper-scalar execution, analyzing 8 offsets `[last - 7, last]` at a time.
// Those need exactly the 8-byte word at `last - 7` and the subsequent two-byte word.
sz_u64_t haystack_page_current, haystack_page_next;
for (; last >= haystack + 7; last -= 8) {
sz_cptr_t const window = last - 7;
haystack_page_current = *(sz_u64_t *)window;
haystack_page_next = *(sz_u16_t *)(window + 8);
haystack0_vec.u64 = (haystack_page_current);
haystack1_vec.u64 = (haystack_page_current >> 8) | (haystack_page_next << 56);
haystack2_vec.u64 = (haystack_page_current >> 16) | (haystack_page_next << 48);
haystack3_vec.u64 = (haystack_page_current >> 24) | (haystack_page_next << 40);
haystack4_vec.u64 = (haystack_page_current >> 32) | (haystack_page_next << 32);
matches0_vec = sz_u64_each_3byte_equal_(haystack0_vec, needle_vec);
matches1_vec = sz_u64_each_3byte_equal_(haystack1_vec, needle_vec);
matches2_vec = sz_u64_each_3byte_equal_(haystack2_vec, needle_vec);
matches3_vec = sz_u64_each_3byte_equal_(haystack3_vec, needle_vec);
matches4_vec = sz_u64_each_3byte_equal_(haystack4_vec, needle_vec);

if (matches0_vec.u64 | matches1_vec.u64 | matches2_vec.u64 | matches3_vec.u64 | matches4_vec.u64) {
matches0_vec.u64 >>= 16;
matches1_vec.u64 >>= 8;
matches3_vec.u64 <<= 8;
matches4_vec.u64 <<= 16;
sz_u64_t match_indicators = matches0_vec.u64 | matches1_vec.u64 | matches2_vec.u64 | matches3_vec.u64 |
matches4_vec.u64;
return window + (63 - sz_u64_clz(match_indicators)) / 8;
}
}

for (; last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) + (last[2] == needle[2]) == 3) return last;
return SZ_NULL_CHAR;
}

/**
* @brief Find the last occurrence of a @b four-character needle in an arbitrary length haystack.
* This implementation uses hardware-agnostic SWAR technique, to process 8 possible offsets at a time.
*/
SZ_HELPER_NOINLINE sz_cptr_t sz_rfind_4byte_serial_(sz_cptr_t haystack, sz_size_t haystack_length, sz_cptr_t needle,
sz_size_t needle_length) {

// This is an internal method, and the haystack is guaranteed to be at least 4 bytes long.
sz_assert_(haystack_length >= 4 && "The haystack is too short.");
sz_unused_(needle_length); //? We keep this argument only for `sz_rfind_t` signature compatibility.
sz_cptr_t last = haystack + haystack_length - 4;

// On big-endian systems, skip SWAR and use simple serial search
#if SZ_IS_BIG_ENDIAN_
for (; last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) + (last[2] == needle[2]) + (last[3] == needle[3]) == 4)
return last;
return SZ_NULL_CHAR;
#endif

// Process the misaligned tail, to void UB on unaligned 64-bit loads.
#if !SZ_USE_MISALIGNED_LOADS
for (; ((sz_size_t)(last + 1) & 7ull) && last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) + (last[2] == needle[2]) + (last[3] == needle[3]) == 4)
return last;
#endif

sz_u64_vec_t haystack0_vec, haystack1_vec, haystack2_vec, haystack3_vec, needle_vec;
sz_u64_vec_t matches0_vec, matches1_vec, matches2_vec, matches3_vec;
needle_vec.u64 = 0;
needle_vec.u8s[0] = needle[0], needle_vec.u8s[1] = needle[1], needle_vec.u8s[2] = needle[2],
needle_vec.u8s[3] = needle[3];
needle_vec.u64 *= 0x0000000100000001ull; // broadcast

// This code simulates hyper-scalar execution, analyzing 8 offsets `[last - 7, last]` at a time.
// Those need exactly the 8-byte word at `last - 7` and the subsequent three bytes.
sz_u64_t haystack_page_current, haystack_page_next;
for (; last >= haystack + 7; last -= 8) {
sz_cptr_t const window = last - 7;
haystack_page_current = *(sz_u64_t *)window;
haystack_page_next = (sz_u64_t)(*(sz_u16_t *)(window + 8)) | ((sz_u64_t)(sz_u8_t)window[10] << 16);
haystack0_vec.u64 = (haystack_page_current);
haystack1_vec.u64 = (haystack_page_current >> 8) | (haystack_page_next << 56);
haystack2_vec.u64 = (haystack_page_current >> 16) | (haystack_page_next << 48);
haystack3_vec.u64 = (haystack_page_current >> 24) | (haystack_page_next << 40);
matches0_vec = sz_u64_each_4byte_equal_(haystack0_vec, needle_vec);
matches1_vec = sz_u64_each_4byte_equal_(haystack1_vec, needle_vec);
matches2_vec = sz_u64_each_4byte_equal_(haystack2_vec, needle_vec);
matches3_vec = sz_u64_each_4byte_equal_(haystack3_vec, needle_vec);

if (matches0_vec.u64 | matches1_vec.u64 | matches2_vec.u64 | matches3_vec.u64) {
matches0_vec.u64 >>= 24;
matches1_vec.u64 >>= 16;
matches2_vec.u64 >>= 8;
sz_u64_t match_indicators = matches0_vec.u64 | matches1_vec.u64 | matches2_vec.u64 | matches3_vec.u64;
return window + (63 - sz_u64_clz(match_indicators)) / 8;
}
}

for (; last >= haystack; --last)
if ((last[0] == needle[0]) + (last[1] == needle[1]) + (last[2] == needle[2]) + (last[3] == needle[3]) == 4)
return last;
return SZ_NULL_CHAR;
}

/**
* @brief Boyer-Moore-Horspool algorithm for exact matching of patterns up to @b 256-bytes long.
* Uses the Raita heuristic to match the first two, the last, and the middle character of the pattern.
Expand Down Expand Up @@ -682,10 +862,9 @@ SZ_API_COMPTIME sz_cptr_t sz_rfind_serial(sz_cptr_t haystack, sz_size_t haystack
sz_find_t backends[] = {
// For very short strings brute-force SWAR makes sense.
sz_rfind_1byte_serial_,
// TODO: implement reverse-order SWAR for 2/3/4 byte variants.
// TODO: sz_rfind_2byte_serial_,
// TODO: sz_rfind_3byte_serial_,
// TODO: sz_rfind_4byte_serial_,
sz_rfind_2byte_serial_,
sz_rfind_3byte_serial_,
sz_rfind_4byte_serial_,
// To avoid constructing the skip-table, let's use the prefixed approach.
// sz_rfind_over_4bytes_serial_,
// For longer needles - use skip tables.
Expand All @@ -695,11 +874,9 @@ SZ_API_COMPTIME sz_cptr_t sz_rfind_serial(sz_cptr_t haystack, sz_size_t haystack

return backends[
// For very short strings brute-force SWAR makes sense.
0 +
// To avoid constructing the skip-table, let's use the prefixed approach.
(needle_length > 1) +
(needle_length > 1) + (needle_length > 2) + (needle_length > 3) +
// For longer needles - use skip tables.
(needle_length > 256)](haystack, haystack_length, needle, needle_length);
(needle_length > 4) + (needle_length > 256)](haystack, haystack_length, needle, needle_length);
}

#ifdef __cplusplus
Expand Down