Skip to main content

simd_minimizers/
syncmers.rs

1//! Collect (and dedup) SIMD-iterator values into a flat `Vec<u32>`.
2
3#![allow(clippy::uninit_vec)]
4
5use std::{
6    array::{self, from_fn},
7    cell::RefCell,
8};
9
10use crate::{S, minimizers::SKIPPED};
11use packed_seq::u32x8;
12use packed_seq::{ChunkIt, L, PaddedIt, intrinsics::transpose};
13use seq_hash::packed_seq;
14
15/// Collect positions of all syncmers.
16/// `OPEN`:
17/// - `false`: closed syncmers
18/// - `true`: open syncmers
19pub fn collect_syncmers_scalar<const OPEN: bool>(
20    w: usize,
21    it: impl Iterator<Item = u32>,
22    out_vec: &mut Vec<u32>,
23) {
24    if OPEN {
25        assert!(
26            w % 2 == 1,
27            "Open syncmers require odd window size, so that there is a unique middle element."
28        );
29    }
30    unsafe { out_vec.set_len(out_vec.capacity()) };
31    let mut idx = 0;
32    it.enumerate().for_each(|(i, min_pos)| {
33        let is_syncmer = if OPEN {
34            min_pos as usize == i + w / 2
35        } else {
36            min_pos as usize == i || min_pos as usize == i + w - 1
37        };
38        if is_syncmer {
39            if idx == out_vec.len() {
40                out_vec.reserve(1);
41                unsafe { out_vec.set_len(out_vec.capacity()) };
42            }
43            *unsafe { out_vec.get_unchecked_mut(idx) } = i as u32;
44            idx += 1;
45        }
46    });
47    out_vec.truncate(idx);
48}
49
50pub trait CollectSyncmers: Sized {
51    /// Collect all indices where syncmers start.
52    ///
53    /// Automatically skips `SIMD_SKIPPED` values for ambiguous windows for sequences shorter than 2^32-2 or so.
54    fn collect_syncmers<const OPEN: bool>(self, w: usize) -> Vec<u32> {
55        let mut v = vec![];
56        self.collect_syncmers_into::<OPEN>(w, &mut v);
57        v
58    }
59
60    /// Collect all indices where syncmers start into `out_vec`.
61    ///
62    /// Automatically skips `SIMD_SKIPPED` values for ambiguous windows for sequences shorter than 2^32-2 or so.
63    fn collect_syncmers_into<const OPEN: bool>(self, w: usize, out_vec: &mut Vec<u32>);
64}
65
66thread_local! {
67    static CACHE: RefCell<[Vec<u32>; 8]> = RefCell::new(array::from_fn(|_| Vec::new()));
68}
69
70impl<I: ChunkIt<u32x8>> CollectSyncmers for PaddedIt<I> {
71    // mostly copied from `Collect::collect_minimizers_into`
72    #[inline(always)]
73    fn collect_syncmers_into<const OPEN: bool>(self, w: usize, out_vec: &mut Vec<u32>) {
74        let Self { it, padding } = self;
75        CACHE.with(
76            #[inline(always)]
77            |v| {
78                let mut v = v.borrow_mut();
79
80                let mut write_idx = [0; 8];
81
82                let len = it.len();
83                let mut lane_offsets: u32x8 = u32x8::from(from_fn(|i| (i * len) as u32));
84
85                let mut mask = u32x8::ZERO;
86                let mut padding_i = 0;
87                let mut padding_idx = 0;
88                assert!(padding <= L * len, "padding {padding} <= L {L} * len {len}");
89                let mut remaining_padding = padding;
90                for i in (0..8).rev() {
91                    if remaining_padding >= len {
92                        mask.as_mut_array()[i] = u32::MAX;
93                        remaining_padding -= len;
94                        continue;
95                    }
96                    padding_i = len - remaining_padding;
97                    padding_idx = i;
98                    break;
99                }
100
101                // FIXME: Is this one slow?
102                let mut m = [u32x8::ZERO; 8];
103                let mut i = 0;
104                it.for_each(
105                    #[inline(always)]
106                    |x| {
107                        if i == padding_i {
108                            mask.as_mut_array()[padding_idx] = u32::MAX;
109                        }
110                        let x = x | mask;
111
112                        // Every non-syncmer minimizer pos is masked out.
113                        let is_syncmer = if OPEN {
114                            x.simd_eq(lane_offsets + S::splat((w / 2) as u32))
115                        } else {
116                            x.simd_eq(lane_offsets)
117                                | x.simd_eq(lane_offsets + S::splat(w as u32 - 1))
118                        };
119                        // current window position if syncmer, else u32::MAX
120                        let y = is_syncmer.blend(lane_offsets, u32x8::MAX);
121
122                        m[i % 8] = y;
123                        if i % 8 == 7 {
124                            let t = transpose(m);
125                            for j in 0..8 {
126                                let lane = t[j];
127                                if write_idx[j] + 8 > v[j].len() {
128                                    v[j].reserve(8);
129                                    unsafe {
130                                        let new_len = v[j].capacity();
131                                        v[j].set_len(new_len);
132                                    }
133                                }
134                                unsafe {
135                                    crate::intrinsics::append_filtered_vals(
136                                        lane,
137                                        // skip masked out values
138                                        lane.simd_eq(u32x8::MAX),
139                                        &mut v[j],
140                                        &mut write_idx[j],
141                                    );
142                                }
143                            }
144                        }
145                        i += 1;
146                        lane_offsets += S::ONE;
147                    },
148                );
149
150                for j in 0..8 {
151                    v[j].truncate(write_idx[j]);
152                }
153
154                // Manually write the unfinished parts of length k=i%8.
155                let t = transpose(m);
156                let k = i % 8;
157                for j in 0..8 {
158                    let lane = t[j].as_array();
159                    for &x in lane.iter().take(k) {
160                        if x < SKIPPED {
161                            v[j].push(x);
162                        }
163                    }
164                }
165
166                // Flatten v.
167                for lane in v.iter() {
168                    out_vec.extend_from_slice(lane.as_slice());
169                }
170            },
171        )
172    }
173}