ftui_simd/lib.rs
1#![forbid(unsafe_code)]
2#![feature(portable_simd)]
3
4//! Portable-SIMD kernels for FrankenTUI's two lane-parallel hot loops.
5//!
6//! # Role in FrankenTUI
7//! Two places in the render path do the same narrow thing over a long run of
8//! bytes: the diff compares one row of 16-byte cells against the previous
9//! frame's row, and the width fast path asks whether a string is entirely
10//! ASCII. Both are a wide equality test followed by "where was the first
11//! difference", which is what SIMD lanes are for.
12//!
13//! # How it fits in the system
14//! Every kernel here is safe `std::simd` and carries a `_scalar` twin with
15//! identical semantics. The twins are not dead weight: they are what the
16//! parity tests compare against, what the benches measure against, and what
17//! callers use when the `simd` feature is off. Nothing in the workspace
18//! depends on this crate unless that feature is enabled, so the scalar paths
19//! stay the default until benches justify otherwise.
20//!
21//! # Why `u128` and not `Cell`
22//! `ftui_render::Cell` is `#[repr(C, align(16))]` over four `u32` fields, so a
23//! cell is bit-for-bit a `u128`. Keeping the kernels on `u128` leaves this
24//! crate free of a dependency on the render crate, and leaves the conversion
25//! (which must stay safe, so it composes the fields rather than transmuting)
26//! on the caller's side.
27//!
28//! # Lane widths
29//! The kernels use 512-bit vectors (`u64x8`, `u8x64`). That is wider than most
30//! targets execute natively; `std::simd` splits them into whatever the target
31//! has, which keeps one code path across x86-64 and aarch64 and gives the
32//! compiler a full unrolled chunk to work with.
33//!
34//! # Which of these are worth calling
35//! Measured 2026-09-18, full numbers in
36//! `docs/perf/simd_kernels_2026-09-18.md`:
37//!
38//! - [`all_ascii`] and [`ascii_width`] beat their twins by 5x at 64 bytes and
39//! by 20-44x from a kilobyte up. Call these.
40//! - [`first_mismatch_u128`] and [`rows_equal_u128`] are **3-4x slower** than
41//! their twins and are deliberately not wired into the diff. `#![forbid(unsafe_code)]`
42//! leaves no way to view `&[u128]` as lanes, so the kernel must build each
43//! vector with shifts and masks, while the scalar `a[i] != b[i]` over `u128`
44//! is already one 128-bit compare after LLVM is done with it. They are kept
45//! because the measurement is worth keeping, and because the parity tests
46//! over them are what prove the lane indexing is right.
47//!
48//! Prefer the `_scalar` twin for cell comparison. That is not a placeholder.
49
50use std::simd::cmp::{SimdPartialEq, SimdPartialOrd};
51use std::simd::{Mask, Simd, u8x64, u64x8};
52
53/// Cells compared per vector chunk: eight `u64` lanes is four `u128` cells.
54const CELLS_PER_CHUNK: usize = 4;
55
56/// Bytes compared per vector chunk.
57const BYTES_PER_CHUNK: usize = 64;
58
59/// Index of the first element where `a` and `b` differ, or `None` when the
60/// common prefix runs to the end of the shorter slice.
61///
62/// Only the first `a.len().min(b.len())` elements are examined; a length
63/// difference past that point is not a mismatch this reports, because the
64/// diff's callers size both rows to the same width and a shorter slice means
65/// a clipped row rather than a changed cell.
66///
67/// # Examples
68/// ```
69/// # use ftui_simd::first_mismatch_u128;
70/// assert_eq!(first_mismatch_u128(&[1, 2, 3], &[1, 9, 3]), Some(1));
71/// assert_eq!(first_mismatch_u128(&[1, 2, 3], &[1, 2, 3]), None);
72/// ```
73#[must_use]
74pub fn first_mismatch_u128(a: &[u128], b: &[u128]) -> Option<usize> {
75 let len = a.len().min(b.len());
76 let mut offset = 0;
77
78 while offset + CELLS_PER_CHUNK <= len {
79 let lhs = load_cells(&a[offset..offset + CELLS_PER_CHUNK]);
80 let rhs = load_cells(&b[offset..offset + CELLS_PER_CHUNK]);
81 let differing: Mask<i64, 8> = lhs.simd_ne(rhs);
82 if differing.any() {
83 // Each cell is two consecutive u64 lanes, so the first differing
84 // lane identifies the cell by halving its index. Both halves of a
85 // cell can differ; the earlier lane is the one `first_set` gives.
86 let lane = differing.first_set().unwrap_or(0);
87 return Some(offset + lane / 2);
88 }
89 offset += CELLS_PER_CHUNK;
90 }
91
92 // Tail shorter than one chunk.
93 (offset..len).find(|&i| a[i] != b[i])
94}
95
96/// Scalar twin of [`first_mismatch_u128`], and the reference its parity tests
97/// compare against.
98#[must_use]
99pub fn first_mismatch_u128_scalar(a: &[u128], b: &[u128]) -> Option<usize> {
100 let len = a.len().min(b.len());
101 (0..len).find(|&i| a[i] != b[i])
102}
103
104/// Whether two equal-length runs of cells are bitwise identical.
105///
106/// Returns `false` for slices of differing length: callers use this to decide
107/// whether a row can be skipped entirely, and a row whose width changed cannot.
108///
109/// # Examples
110/// ```
111/// # use ftui_simd::rows_equal_u128;
112/// assert!(rows_equal_u128(&[7, 7], &[7, 7]));
113/// assert!(!rows_equal_u128(&[7, 7], &[7, 7, 7]));
114/// ```
115#[must_use]
116pub fn rows_equal_u128(a: &[u128], b: &[u128]) -> bool {
117 a.len() == b.len() && first_mismatch_u128(a, b).is_none()
118}
119
120/// Scalar twin of [`rows_equal_u128`].
121#[must_use]
122pub fn rows_equal_u128_scalar(a: &[u128], b: &[u128]) -> bool {
123 a.len() == b.len() && first_mismatch_u128_scalar(a, b).is_none()
124}
125
126/// Whether every byte is ASCII, i.e. has its high bit clear.
127///
128/// This is the question the width fast path actually asks: if no byte is
129/// continuation or lead, the run needs no Unicode width lookup and its display
130/// width is its length.
131///
132/// # Examples
133/// ```
134/// # use ftui_simd::all_ascii;
135/// assert!(all_ascii(b"plain text"));
136/// assert!(!all_ascii("caf\u{e9}".as_bytes()));
137/// assert!(all_ascii(b""));
138/// ```
139#[must_use]
140pub fn all_ascii(bytes: &[u8]) -> bool {
141 let mut offset = 0;
142
143 while offset + BYTES_PER_CHUNK <= bytes.len() {
144 let chunk = u8x64::from_slice(&bytes[offset..offset + BYTES_PER_CHUNK]);
145 if (chunk & u8x64::splat(0x80)).simd_ne(u8x64::splat(0)).any() {
146 return false;
147 }
148 offset += BYTES_PER_CHUNK;
149 }
150
151 bytes[offset..].iter().all(u8::is_ascii)
152}
153
154/// Scalar twin of [`all_ascii`].
155#[must_use]
156pub fn all_ascii_scalar(bytes: &[u8]) -> bool {
157 bytes.iter().all(u8::is_ascii)
158}
159
160/// Display width of a run that is entirely printable ASCII, or `None`.
161///
162/// `None` means "this run needs the real width tables", and covers both
163/// non-ASCII bytes and ASCII control characters: C0 (`0x00..=0x1f`) and DEL
164/// (`0x7f`) have no single agreed display width, so they are handed back to
165/// the caller rather than counted as one column each. Space through `~` are
166/// one column apiece, so the width of an accepted run is its length.
167///
168/// # Examples
169/// ```
170/// # use ftui_simd::ascii_width;
171/// assert_eq!(ascii_width(b"hello"), Some(5));
172/// assert_eq!(ascii_width(b"tab\there"), None);
173/// assert_eq!(ascii_width("\u{e9}".as_bytes()), None);
174/// ```
175#[must_use]
176pub fn ascii_width(bytes: &[u8]) -> Option<usize> {
177 let mut offset = 0;
178
179 while offset + BYTES_PER_CHUNK <= bytes.len() {
180 let chunk = u8x64::from_slice(&bytes[offset..offset + BYTES_PER_CHUNK]);
181 // Printable ASCII is 0x20..=0x7e, so one unsigned-wrapping subtract
182 // puts every acceptable byte in 0..=0x5e and everything else above it:
183 // 0x7f wraps to 0x5f, and any high-bit byte lands at 0x60 or more.
184 let shifted = chunk - u8x64::splat(0x20);
185 if shifted.simd_gt(u8x64::splat(0x5e)).any() {
186 return None;
187 }
188 offset += BYTES_PER_CHUNK;
189 }
190
191 if bytes[offset..].iter().all(|b| (0x20..=0x7e).contains(b)) {
192 Some(bytes.len())
193 } else {
194 None
195 }
196}
197
198/// Scalar twin of [`ascii_width`].
199#[must_use]
200pub fn ascii_width_scalar(bytes: &[u8]) -> Option<usize> {
201 if bytes.iter().all(|b| (0x20..=0x7e).contains(b)) {
202 Some(bytes.len())
203 } else {
204 None
205 }
206}
207
208/// Reinterpret four `u128` cells as the eight `u64` lanes of one vector.
209///
210/// Splitting each cell into two `u64` halves keeps the whole thing in safe
211/// code: there is no 128-bit lane type to compare with, and `u64` is the
212/// widest lane every supported target handles well.
213#[inline]
214fn load_cells(cells: &[u128]) -> u64x8 {
215 debug_assert_eq!(cells.len(), CELLS_PER_CHUNK);
216 let mut lanes = [0_u64; 8];
217 for (i, cell) in cells.iter().enumerate() {
218 lanes[i * 2] = (*cell & u128::from(u64::MAX)) as u64;
219 lanes[i * 2 + 1] = (*cell >> 64) as u64;
220 }
221 Simd::from_array(lanes)
222}
223
224#[cfg(test)]
225mod tests {
226 use super::*;
227
228 #[test]
229 fn first_mismatch_reports_none_for_identical_rows() {
230 let row: Vec<u128> = (0..200).collect();
231 assert_eq!(first_mismatch_u128(&row, &row), None);
232 assert_eq!(first_mismatch_u128_scalar(&row, &row), None);
233 }
234
235 #[test]
236 fn first_mismatch_finds_the_earliest_difference_at_every_position() {
237 // Every index matters: inside the first chunk, straddling a chunk
238 // boundary, and in the scalar tail.
239 for len in [1_usize, 3, 4, 5, 8, 63, 64, 65, 200] {
240 let base: Vec<u128> = (0..len as u128).collect();
241 for idx in 0..len {
242 let mut changed = base.clone();
243 changed[idx] = u128::MAX;
244 assert_eq!(
245 first_mismatch_u128(&base, &changed),
246 Some(idx),
247 "len {len}, idx {idx}"
248 );
249 assert_eq!(
250 first_mismatch_u128_scalar(&base, &changed),
251 Some(idx),
252 "scalar len {len}, idx {idx}"
253 );
254 }
255 }
256 }
257
258 #[test]
259 fn first_mismatch_detects_a_difference_in_either_half_of_a_cell() {
260 // The low half alone, the high half alone, and both: all three must
261 // report the same cell, since a cell is two lanes.
262 for delta in [1_u128, 1_u128 << 64, (1_u128 << 64) | 1] {
263 let base = vec![0_u128; 8];
264 let mut changed = base.clone();
265 changed[5] = delta;
266 assert_eq!(
267 first_mismatch_u128(&base, &changed),
268 Some(5),
269 "delta {delta}"
270 );
271 }
272 }
273
274 #[test]
275 fn first_mismatch_stops_at_the_shorter_slice() {
276 let long: Vec<u128> = (0..16).collect();
277 let short = &long[..6];
278 assert_eq!(first_mismatch_u128(&long, short), None);
279 assert_eq!(first_mismatch_u128(short, &long), None);
280 }
281
282 #[test]
283 fn first_mismatch_handles_empty_input() {
284 assert_eq!(first_mismatch_u128(&[], &[]), None);
285 assert_eq!(first_mismatch_u128(&[], &[1, 2]), None);
286 }
287
288 #[test]
289 fn rows_equal_requires_matching_length() {
290 assert!(rows_equal_u128(&[1, 2, 3], &[1, 2, 3]));
291 assert!(!rows_equal_u128(&[1, 2, 3], &[1, 2]));
292 assert!(!rows_equal_u128(&[1, 2, 3], &[1, 2, 4]));
293 assert!(rows_equal_u128(&[], &[]));
294 }
295
296 #[test]
297 fn all_ascii_agrees_with_the_scalar_twin_across_lengths() {
298 for len in [0_usize, 1, 63, 64, 65, 127, 128, 1000] {
299 let ascii = vec![b'a'; len];
300 assert!(all_ascii(&ascii), "len {len}");
301 assert_eq!(all_ascii(&ascii), all_ascii_scalar(&ascii));
302
303 if len > 0 {
304 // A single non-ASCII byte must be found wherever it sits.
305 for idx in [0, len / 2, len - 1] {
306 let mut probe = ascii.clone();
307 probe[idx] = 0xC3;
308 assert!(!all_ascii(&probe), "len {len}, idx {idx}");
309 assert_eq!(all_ascii(&probe), all_ascii_scalar(&probe));
310 }
311 }
312 }
313 }
314
315 #[test]
316 fn all_ascii_accepts_control_bytes() {
317 // Control characters are ASCII; only ascii_width is fussy about them.
318 assert!(all_ascii(b"\x00\x09\x1b\x7f"));
319 }
320
321 #[test]
322 fn ascii_width_counts_printable_runs_and_rejects_the_rest() {
323 assert_eq!(ascii_width(b""), Some(0));
324 assert_eq!(ascii_width(b" "), Some(1));
325 assert_eq!(ascii_width(b"~"), Some(1));
326 assert_eq!(ascii_width(b"hello world"), Some(11));
327
328 let long = vec![b'x'; 300];
329 assert_eq!(ascii_width(&long), Some(300));
330
331 // The boundaries either side of printable ASCII, and a high byte.
332 assert_eq!(ascii_width(b"\x1f"), None);
333 assert_eq!(ascii_width(b"\x7f"), None);
334 assert_eq!(ascii_width(b"\xc3\xa9"), None);
335 }
336
337 #[test]
338 fn ascii_width_rejects_a_bad_byte_anywhere_including_past_a_full_chunk() {
339 for len in [64_usize, 65, 129, 300] {
340 for idx in [0, len / 2, len - 1] {
341 for bad in [0x00_u8, 0x1f, 0x7f, 0x80, 0xff] {
342 let mut probe = vec![b'a'; len];
343 probe[idx] = bad;
344 assert_eq!(ascii_width(&probe), None, "len {len}, idx {idx}, bad {bad}");
345 assert_eq!(ascii_width(&probe), ascii_width_scalar(&probe));
346 }
347 }
348 }
349 }
350
351 #[test]
352 fn kernels_agree_with_their_twins_over_a_deterministic_sweep() {
353 // A cheap xorshift keeps this reproducible without a dev-dependency.
354 let mut state = 0x2545_F491_4F6C_DD1D_u64;
355 let mut next = move || {
356 state ^= state << 13;
357 state ^= state >> 7;
358 state ^= state << 17;
359 state
360 };
361
362 for len in 0..200_usize {
363 // Fill both halves of each cell. Widening a single u64 left every
364 // base value with a zero high half, so a kernel that only ever
365 // compared the low 64 bits of equal cells would have passed —
366 // the perturbation below can flip a high bit, but two *equal*
367 // cells were never wide.
368 let a: Vec<u128> = (0..len)
369 .map(|_| (u128::from(next()) << 64) | u128::from(next()))
370 .collect();
371 let mut b = a.clone();
372 if len > 0 {
373 let idx = (next() as usize) % len;
374 // Half the time perturb a cell, half the time leave them equal.
375 if next() % 2 == 0 {
376 b[idx] ^= 1 << (next() % 128);
377 }
378 }
379 assert_eq!(
380 first_mismatch_u128(&a, &b),
381 first_mismatch_u128_scalar(&a, &b),
382 "len {len}"
383 );
384 assert_eq!(rows_equal_u128(&a, &b), rows_equal_u128_scalar(&a, &b));
385
386 let bytes: Vec<u8> = (0..len).map(|_| (next() % 256) as u8).collect();
387 assert_eq!(all_ascii(&bytes), all_ascii_scalar(&bytes), "len {len}");
388 assert_eq!(ascii_width(&bytes), ascii_width_scalar(&bytes), "len {len}");
389 }
390 }
391
392 /// The sweep above draws bytes uniformly from `0..=255`, so a run of
393 /// `BYTES_PER_CHUNK` bytes is entirely ASCII with probability 2^-64. Every
394 /// chunk it ever evaluates therefore takes the rejecting branch on its
395 /// first iteration: the accepting path through the chunk loop - advancing
396 /// `offset`, and the tail that is measured from wherever it stopped - is
397 /// never reached for any input long enough to have a chunk at all.
398 ///
399 /// These runs are printable ASCII by construction, then poisoned one byte
400 /// at a time so the rejecting branch is exercised at every position rather
401 /// than only near the front.
402 #[test]
403 fn ascii_kernels_agree_with_their_twins_on_runs_that_reach_the_chunk_loop() {
404 // Past two full chunks, so a bad byte can land in the first chunk, a
405 // later chunk, or the tail.
406 for len in 0..=(2 * BYTES_PER_CHUNK + 5) {
407 // 0x20..=0x7e, cycled: printable, and not a single repeated byte.
408 let mut run: Vec<u8> = (0..len).map(|i| 0x20 + (i % 0x5f) as u8).collect();
409
410 assert!(all_ascii(&run), "len {len}");
411 assert_eq!(all_ascii(&run), all_ascii_scalar(&run), "len {len}");
412 assert_eq!(ascii_width(&run), Some(len), "len {len}");
413 assert_eq!(ascii_width(&run), ascii_width_scalar(&run), "len {len}");
414
415 for pos in 0..len {
416 // NUL and 0x1f are ASCII but not printable, so the two kernels
417 // disagree with each other by design - each is compared only
418 // against its own twin. 0x7f is the byte the wrapping subtract
419 // in `ascii_width` folds to the top of the accepted range.
420 for bad in [0x00_u8, 0x1f, 0x7f, 0x80, 0xff] {
421 let saved = run[pos];
422 run[pos] = bad;
423 assert_eq!(
424 all_ascii(&run),
425 all_ascii_scalar(&run),
426 "all_ascii len {len} pos {pos} byte {bad:#04x}"
427 );
428 assert_eq!(
429 ascii_width(&run),
430 ascii_width_scalar(&run),
431 "ascii_width len {len} pos {pos} byte {bad:#04x}"
432 );
433 run[pos] = saved;
434 }
435 }
436 }
437 }
438}