Skip to main content

debruijn/
filter.rs

1// Copyright 2017 10x Genomics
2
3//! Methods for converting sequences into kmers, filtering observed kmers before De Bruijn graph construction, and summarizing 'color' annotations.
4
5use std::collections::HashSet;
6use std::fmt::Debug;
7use std::hash::Hash;
8use std::mem;
9use std::ops::Range;
10use std::sync::Arc;
11use std::sync::Mutex;
12use std::time::Instant;
13
14use boomphf::hashmap::BoomHashMap2;
15use indicatif::MultiProgress;
16use indicatif::ProgressBar;
17use indicatif::ProgressIterator;
18use indicatif::ProgressStyle;
19use itertools::Itertools;
20use log::debug;
21use log::warn;
22use rayon::current_num_threads;
23use rayon::prelude::*;
24
25use crate::KmerDataItem;
26use crate::reads::Read;
27use crate::reads::ReadsPaired;
28use crate::reads::Strandedness;
29use crate::summarizer::SummaryConfig;
30use crate::summarizer::SummaryData;
31use crate::Dir;
32use crate::Exts;
33use crate::Kmer;
34use crate::Vmer;
35use crate::BUCKETS;
36use crate::PROGRESS_STYLE;
37
38const HEAP_APPROX: f32 = 1.5;
39
40/// check which bucket a k-mer has to be sorted into according to first four bases
41pub fn bucket<K: Kmer>(kmer: K) -> usize {
42    if K::k() > 3 {
43        (kmer.get(0) as usize) << 6
44            | (kmer.get(1) as usize) << 4
45            | (kmer.get(2) as usize) << 2
46            | (kmer.get(3) as usize)
47    } else {
48        kmer.to_u64() as usize
49    }
50}
51
52fn bucket_flip<K: Kmer>(kmer: K, stranded: Strandedness) -> (usize, K) {
53    // if not stranded choose lexiographically lesser of kmer and rc of kmer
54    // if forward, use original kmer
55    // if reverse, use rc of kmer
56    let min_kmer = match stranded {
57        Strandedness::Unstranded => {
58            let (min_kmer, _) = kmer.min_rc_flip();
59            min_kmer
60        },
61        Strandedness::Forward => kmer,
62        Strandedness::Reverse => kmer.rc(),
63    };
64
65    // calculate which bucket this kmer belongs to
66    (bucket(min_kmer), min_kmer)
67}
68
69fn bucket_ext_flip<K: Kmer>(kmer: K, exts: Exts, stranded: Strandedness, bucket_range: Range<usize>) ->Option<(K, Exts, usize)> {
70    // if not stranded choose lexiographically lesser of kmer and rc of kmer
71    // if forward, use original kmer
72    // if reverse, use rc of kmer
73    let (min_kmer, flip_exts) = match stranded {
74        Strandedness::Unstranded => {
75            let (min_kmer, flip) = kmer.min_rc_flip();
76            let flip_exts = if flip { exts.rc() } else { exts };
77            (min_kmer, flip_exts)
78        },
79        Strandedness::Forward => (kmer, exts),
80        Strandedness::Reverse => (kmer.rc(), exts.rc()),
81    };
82
83    // calculate which bucket this kmer belongs to
84    let bucket = if K::k() > 3 { bucket(min_kmer) } else { min_kmer.to_u64() as usize };
85
86    // check if bucket is in current range
87    let in_range = bucket >= bucket_range.start && bucket < bucket_range.end;
88
89    if in_range {
90        Some((min_kmer, flip_exts, bucket))
91    } else {
92        None
93    }
94}
95
96/// increase the capacities for each bucket for the k-mers in one read
97fn add_seq_bucket_capacities<K: Kmer, DI: Clone + Copy>(read: &Read<DI>, capacities: &mut [usize; BUCKETS], unique_kmers: &mut HashSet<K>) {
98    // iterate through all kmers in seq
99    for kmer in read.seq().iter_kmers::<K>() {
100        // calculate which bucket this kmer belongs to
101        let (bucket, min_kmer) = bucket_flip(kmer, read.stranded());
102        capacities[bucket] += 1;
103
104        // count k-mer coverages
105        unique_kmers.insert(min_kmer);
106    }
107}
108
109/// caclulate the ranges (slices) in which the buckets are split
110fn bucket_ranges<K, SD, DI, I>(n_nodes: usize, input_kmers: usize, memory_size: f32, iter_capacities: I) -> Vec<Range<usize>> 
111where I: Iterator<Item = (usize, usize)>
112{
113    // calculate expected size per node
114    // some SDs will have additional content in heap, we approximate this by adding 50%
115    let exp_node_mem = mem::size_of::<K>() + mem::size_of::<Exts>() + (mem::size_of::<SD>() as f32 * HEAP_APPROX) as usize;
116
117    let graph_mem = n_nodes * exp_node_mem;
118    debug!("n final nodes: {n_nodes}");
119    debug!("expected graph memory: {graph_mem}");
120
121    // calculate numnber of slices for memory limit
122    let mem_per_kmer = mem::size_of::<(K, Exts, DI)>();
123    debug!("size used for calculation: {} bytes", mem_per_kmer);
124    debug!("size of kmer, E, D: {} bytes", mem::size_of::<(K, Exts, DI)>());
125    debug!("size of K: {} bytes, size of Exts: {} bytes, size of DI: {} bytes", mem::size_of::<K>(), mem::size_of::<Exts>(), mem::size_of::<DI>());
126    debug!("type DI: {}", std::any::type_name::<DI>());
127
128    let memory_limit = (memory_size * 10f32.powf(9.)) as usize;
129    let max_mem = memory_limit.saturating_sub(graph_mem);
130    let required_slices = if max_mem == 0 {
131        BUCKETS
132    } else {
133        mem_per_kmer * input_kmers / max_mem
134    }; 
135
136    let bucket_ranges = if required_slices >= BUCKETS {
137        warn!("supplied memory limit might not be sufficient, will construct graph with lowest possible memory usage");
138        // run each bucket in a separate slice
139        (0..BUCKETS).map(|i| i..(i+1)).collect::<Vec<_>>()
140    } else {
141        // if max_mem > 0, split k-mers into however many slices are needed
142        debug!("splitting k-mers into {required_slices} slices");
143
144        let mut start_bucket = 0;
145        let mut size = 0;
146
147        // maximum number of k-mers in a slice
148        let max_size = max_mem / mem_per_kmer;
149
150        let mut bucket_ranges = Vec::with_capacity(required_slices);
151
152        for (i, capacity) in iter_capacities {
153            size += capacity;
154            if size > max_size {
155                bucket_ranges.push(start_bucket..i);
156                start_bucket = i;
157                size = capacity;
158            }
159        }
160        bucket_ranges.push(start_bucket..BUCKETS);
161
162        bucket_ranges
163    };
164
165    debug!("bucket ranges: {:?}", bucket_ranges);
166    debug!("kmer_mem: {} bytes, max_mem: {} bytes, slices: {}", mem_per_kmer * input_kmers, max_mem, bucket_ranges.len());
167    debug!("bucket_ranges: {:?}", bucket_ranges);
168    assert!(bucket_ranges[bucket_ranges.len() - 1].end >= BUCKETS);
169
170    bucket_ranges
171}
172
173/// Process DNA sequences into kmers and determine the set of valid kmers,
174/// their extensions, and summarize associated label/'color' data. The input
175/// sequences are converted to kmers of type `K`, and like kmers are grouped together.
176/// All instances of each kmer, along with their label data are then proccessed with
177/// [`SummaryData::summarize`], which generates an implementation of [`SummaryData`],
178/// which is specified with the generic `SD`, decides if the k-mer is 'valid' 
179/// based on the parameters given in `summary_config`, and
180/// summarizes the the individual label into a single label data structure
181/// for the kmer. Care is taken to keep the memory consumption small.
182/// 
183/// Be aware that the configuration in `summary_config` only applies if the required
184/// informaiton can be supplied by the chosen implementation of `SummaryData`.
185/// E.g., the k-mers will not be filtered according to p-value when the `SummaryData`
186/// only contains the number of observations.
187///
188/// # Arguments
189///
190/// * `seqs` are the reads wrapped in a `Reads<u8>`. See [`Reads<D>`]
191/// * `summary_config` is a [`SummaryConfig`], which contains prameters and 
192///   information necessary for the filtering
193/// * `stranded`: if true, preserve the strandedness of the input sequences, effectively
194///   assuming they are all in the positive strand. If false, the kmers will be canonicalized
195///   to the lexicographic minimum of the kmer and it's reverse complement.
196/// * `report_all_kmers`: if true returns the vector of all the observed kmers and performs the
197///   kmer based filtering
198/// * `memory_size`: gives the size bound on the memory in GB to use and automatically determines
199///   the number of passes needed
200/// * `time`: print information about the time needed for each step
201/// # Returns
202/// BoomHashMap2 Object, check rust-boomphf for details
203/// 
204/// /// # Returns
205/// BoomHashMap2 Object, check rust-boomphf for details
206/// 
207/// # Examples:
208/// 
209/// ```
210/// use debruijn::summarizer::{SampleInfo, SummaryConfig, TagsCountsData, StatTest, GroupFrac};
211/// use debruijn::reads::{Reads, ReadsPaired, Strandedness};
212/// use debruijn::filter::filter_kmers_parallel;
213/// use debruijn::kmer::Kmer16;
214/// use debruijn::Exts;
215/// 
216/// let mut seqs = Reads::new(Strandedness::Unstranded);
217/// seqs.add_from_bytes("ACCGATCATATATTTTCGGGGCTAGGCGAAGCGATCTTATCGAGC".as_bytes(), None, 1u8);
218/// seqs.add_from_bytes("GCGATCGAGCATGCTCAGCTGACGTGACTGACGTAGCTATCTTTTCGTAGCTAC".as_bytes(), None, 1u8);
219/// seqs.add_from_bytes("GCGAGTTTGCGACTCGAGGCTATCTAGCTAGCTASGCTCTCGACTAGCTGACTTACGACGACTACG".as_bytes(), None, 2u8);
220/// seqs.add_from_bytes("CGATTAGCTACGTAGCTAGCTGACGTACTGGGGGGTATTTCGGATCTGCGGAGCGATCT".as_bytes(), None, 2u8);
221///       
222/// let sample_info = SampleInfo::new(
223///     0b000011,
224///     0b111100,
225///     vec![23423, 3463454, 2242234, 2233243, 234322434, 2323234],
226/// );
227///     
228/// let summary_config = SummaryConfig::new(sample_info)
229///     .with_min_kmer_obs(3)
230///     .with_group_frac(GroupFrac::One, 0.333)
231///     .with_stat_test(StatTest::StudentsTTest);
232///    
233/// let (hashed_kmers, _) = filter_kmers_parallel::<Kmer16, TagsCountsData, _>(
234///     &ReadsPaired::Unpaired { reads: seqs },
235///     &summary_config,
236///     false,
237///     10.,
238///     false,
239/// );
240/// ```
241#[inline(never)]
242//pub fn filter_kmers_parallel<K: Kmer + Sync + Send, V: Vmer + Sync, D1: Clone + Debug + Sync, DS: Clone + Sync + Send, S: KmerSummarizer<D1, DS, (usize, usize)> +  Send>(
243pub fn filter_kmers_parallel<K, SD, DI>(
244    seqs: &ReadsPaired<DI>,
245    summariy_config: &SummaryConfig,
246    report_all_kmers: bool,
247    memory_size: f32,
248    time: bool,
249) -> (BoomHashMap2<K, Exts, SD>, Vec<K>)
250where 
251K: Kmer + Sync + Send,
252SD: Clone + std::fmt::Debug + Send + SummaryData<DI>,
253DI: Clone + Copy + Send + Sync
254{
255    // take timestamp before all processes
256    let before_all = Instant::now();
257
258    // progress bars
259    let multi_pb = MultiProgress::new();
260    let style = ProgressStyle::with_template(PROGRESS_STYLE).unwrap().progress_chars("#/-");
261
262    // split all reads into ranges to be processed in parallel for counting capacities and picking kmers
263    let n_threads = current_num_threads();
264    let n_reads = seqs.n_reads();
265    let sz = n_reads / n_threads + 1;
266
267    debug!("n_reads: {}", n_reads);
268    debug!("sz: {}", sz);
269
270    let mut parallel_ranges = Vec::with_capacity(n_threads);
271    let mut start = 0;
272    while start < n_reads {
273        parallel_ranges.push(start..start + sz);
274        start += sz;
275    }
276
277    let last_start = parallel_ranges.pop().expect("no kmers in parallel ranges").start;
278    parallel_ranges.push(last_start..n_reads);
279    debug!("parallel ranges: {:?}", parallel_ranges);
280
281    let capacities = Arc::new(Mutex::new(vec![[0; BUCKETS]; n_threads]));
282    let unique_kmers = Arc::new(Mutex::new(vec![HashSet::<K>::new(); n_threads]));
283
284    let pb_size_buckets = multi_pb.add(ProgressBar::new(seqs.n_reads() as u64));
285    pb_size_buckets.set_style(style.clone());
286    pb_size_buckets.set_message(format!("{:<32}", "finding bucket sizes"));
287
288    parallel_ranges.clone().into_par_iter().enumerate().for_each(|(i, range)| {
289
290        // first go trough all kmers to find the length of all buckets (to reserve capacity)
291        let mut thread_capacities = [0usize; BUCKETS];
292        let mut thread_kmers = HashSet::new();
293        for ref read in seqs.iter_partial(range.clone())
294        { 
295            // add the required capacities to the respective buckets
296            add_seq_bucket_capacities(read, &mut thread_capacities, &mut thread_kmers);
297
298            pb_size_buckets.inc(1);
299        }
300
301        let mut cap = capacities.lock().expect("error locking capacity mutex");
302        cap[i] = thread_capacities;
303
304        let mut u_kmers = unique_kmers.lock().expect("error locking coverages mutex");
305        u_kmers[i] = thread_kmers;
306    });
307
308    let capacities = capacities.lock().expect("error in final lock capacites");
309    let input_kmers = capacities.iter().flatten().sum::<usize>();
310
311    let mut unique_kmers = unique_kmers.lock().expect("error final lock coverages");
312
313    // combine unique k-mers found by separate threads
314    let mut combined_kmers = unique_kmers.pop().expect("no k-mers, empty graph");
315
316    while let Some(k) = unique_kmers.pop() {
317        combined_kmers.extend(k);
318    }
319
320    let iter_capacities = (0..BUCKETS).map(|i| (i, capacities.iter().map(|c_bucket| c_bucket[i]).sum::<usize>()));
321    let bucket_ranges = bucket_ranges::<K, SD, DI, _>(combined_kmers.len(), input_kmers, memory_size, iter_capacities);
322
323    let n_slices = bucket_ranges.len();
324
325    debug!("n of seqs: {}", seqs.n_reads());
326
327    let mut time_picking_par = 0.;
328    let mut time_picking = 0.;
329    let mut time_summarizing = 0.;
330
331    let shared_target_vecs = Arc::new(Mutex::new((Vec::new(), Vec::new(), Vec::new(), Vec::new())));
332
333    if time { println!("time all prepariations before sliced in filter_kmers (s): {}", before_all.elapsed().as_secs_f32()) }
334
335    let pb_bucket_ranges = multi_pb.add(ProgressBar::new(bucket_ranges.len() as u64));
336    pb_bucket_ranges.set_style(style.clone());
337    pb_bucket_ranges.set_message(format!("{:<32}", "filtering kmers"));
338
339    for (i, bucket_range) in bucket_ranges.into_iter().enumerate() {
340
341        debug!("Processing slice {} of {}", i+1, n_slices);
342
343        let before_kmer_picking = Instant::now();
344        // first step: picking kmers with their exts & data from the reads
345        // go through all kmers and sort into bucket according to first four bases
346        // all kmers starting with "AAAA" go in kmer_buckets[0], all starting with AAAC go in kmer_buckets[1] and so on
347        // when using the first four bases, this needs 256 buckets
348        // the buckets are split in to the bucket_ranges to save memory
349
350
351        let kmer_buckets = Arc::new(Mutex::new(vec![Vec::new(); n_threads]));
352
353        let before_picking_parallel = Instant::now();
354
355        let pb_fill_buckets = multi_pb.add(ProgressBar::new(seqs.n_reads() as u64));
356        pb_fill_buckets.set_style(style.clone());
357        pb_fill_buckets.set_message(format!("{:<32}", "filling buckets with k-mers"));
358
359        let pb_sum_buckets = multi_pb.add(ProgressBar::new(BUCKETS as u64));
360        pb_sum_buckets.set_style(style.clone());
361        pb_sum_buckets.set_message(format!("{:<32}", "summarizing k-mers in buckets"));
362
363        parallel_ranges.clone().into_par_iter().enumerate().for_each(|(i, range)| {
364            let mut kmer_buckets1d = Vec::with_capacity(BUCKETS); 
365            
366            // reserve capacities needed for current range in each bucket
367            for (i, capacity) in capacities[i].into_iter().enumerate() {
368                if bucket_range.contains(&i) {
369                    // capacity is in bucket range, allocate bucket with capacity
370                    kmer_buckets1d.push(Vec::with_capacity(capacity));
371                } else {
372                    // not in current range, add empty placeholder vector
373                    kmer_buckets1d.push(Vec::new());
374                }
375                
376            }
377
378            // fill buckets with kmers
379            for ref read in seqs.iter_partial(range.clone())
380            {
381                for (kmer, exts, quality) in read.iter_kmer_exts_quality::<K>() {
382                    // if needed, flip kmer and exts
383                    // check if bucket is in current range and if so, push kmer to bucket
384                    if let Some((min_kmer, flip_exts, bucket)) = bucket_ext_flip(kmer, exts, read.stranded(), bucket_range.clone()) {
385                        kmer_buckets1d[bucket].push(KmerDataItem::new(min_kmer, flip_exts, read.data(), quality));
386                    }
387                }
388
389                pb_fill_buckets.inc(1);
390            }
391
392            // clone and lock kmer_buckets to safely share across threads
393            let _kb_clone = Arc::clone(&kmer_buckets);
394            let mut kb2d = kmer_buckets.lock().expect("lock kmer buckets 2d");
395            // replace empty vec too keep order
396            kb2d[i] = kmer_buckets1d;
397
398        });
399
400        time_picking_par += before_picking_parallel.elapsed().as_secs_f32();
401
402        // unlock kmer buckets and move out of guard so they can be turned into iterator
403        let mut kmer_buckets = kmer_buckets.lock().expect("unlock kmer_buckets final");
404        let kmer_buckets = mem::take(&mut *kmer_buckets);
405
406        // all combined buckets go into this vector
407        let mut new_buckets = vec![Vec::new(); BUCKETS];
408        // flatten kmer buckets
409        for thread_vec in kmer_buckets.into_iter() {
410            for (i, mut bucket) in thread_vec.into_iter().enumerate() {
411                new_buckets[i].reserve_exact(bucket.len());
412                new_buckets[i].append(&mut bucket);
413            }
414        }
415
416        time_picking += before_kmer_picking.elapsed().as_secs_f32();
417        
418        let before_parallel = Instant::now();
419        
420        // parallel start
421        // summarize kmers in buckets      
422        new_buckets.into_par_iter().for_each(|mut kmer_vec| {
423            kmer_vec.sort_by_key(|elt| elt.kmer);
424
425            let size = kmer_vec.iter().chunk_by(|elt| elt.kmer).into_iter().count();
426
427            let mut all_kmers = Vec::with_capacity(size);
428            let mut valid_kmers = Vec::with_capacity(size);
429            let mut valid_exts = Vec::with_capacity(size);
430            let mut valid_data = Vec::with_capacity(size);
431
432
433            for (kmer, kmer_obs_iter) in kmer_vec.into_iter().chunk_by(|elt| elt.kmer).into_iter() {
434                let (is_valid, exts, summary_data) = SD::summarize(kmer_obs_iter, summariy_config);
435                if report_all_kmers {
436                    all_kmers.push(kmer);
437                }
438                if is_valid {
439                    valid_kmers.push(kmer);
440                    valid_exts.push(exts);
441                    valid_data.push(summary_data);
442                }
443            }
444
445            // if there are valid k-mers in this bucket, append them to the shared target vectors
446            // important that this is done in one step so each kmer has the same index with its exts and data
447            if !valid_kmers.is_empty() {
448
449                let _stv_clone = Arc::clone(&shared_target_vecs);
450                let mut stv = shared_target_vecs.lock().expect("lock target vectors");
451                // valid kmers
452                stv.0.reserve_exact(valid_kmers.len());
453                stv.0.append(&mut valid_kmers); 
454                // valid exts
455                stv.1.reserve_exact(valid_exts.len());
456                stv.1.append(&mut valid_exts); 
457                // valid data
458                stv.2.reserve_exact(valid_data.len());
459                stv.2.append(&mut valid_data); 
460            }
461
462            // if kmers were collected into all_kmers, append them to shared target vector
463            if !all_kmers.is_empty() {
464                let _stv_clone = Arc::clone(&shared_target_vecs);
465                let mut stv = shared_target_vecs.lock().expect("lock target vectors");
466                // all kmers
467                stv.3.reserve_exact(all_kmers.len());
468                stv.3.append(&mut all_kmers);
469            }
470
471            pb_sum_buckets.inc(1);
472        });
473        // parallel end
474
475        pb_bucket_ranges.inc(1);
476
477        time_summarizing += before_parallel.elapsed().as_secs_f32();
478
479        debug!("processed slice {}", i+1);
480    }
481    pb_bucket_ranges.finish_and_clear();
482
483    if time { 
484        println!("time collecting par (s): {}", time_picking_par);
485        println!("time collecting (s): {}", time_picking);
486        println!("time summarizing (s): {}", time_summarizing);
487    }
488
489    let stv = shared_target_vecs.lock().expect("final lock target vectors");
490
491    debug!("valid kmers - capacity: {}, size: {}, mem: {}", stv.0.capacity(), stv.0.len(), mem::size_of_val(&*stv.0));
492    debug!("valid exts - capacity: {}, size: {}, mem: {}", stv.1.capacity(), stv.1.len(), mem::size_of_val(&*stv.1));
493    debug!("valid data - capacity: {}, size: {}, struct mem: {}, real mem: {}", stv.2.capacity(), stv.2.len(), mem::size_of_val(&*stv.2), {
494        let mut data_size = 0;
495        for data in &stv.2 {
496            data_size += data.mem();
497        }
498        data_size
499    });
500    debug!("all kmers - capacity: {}, size: {}, mem: {}", stv.3.capacity(), stv.3.len(), mem::size_of_val(&*stv.3));
501    
502
503    let before_hash = Instant::now();
504    let hm = BoomHashMap2::new_parallel(stv.0.to_vec(), stv.1.to_vec(), stv.2.to_vec());
505    let after_hash = before_hash.elapsed().as_secs_f32();
506    let all_kmers = stv.3.to_vec();
507
508    let filter_kmers_inner = before_all.elapsed().as_secs_f32();
509    if time { 
510        println!("time filter_kmers inner (s): {}", filter_kmers_inner);
511        println!("time build filtered hash map (s): {}", after_hash);
512    }
513
514    (
515        hm,
516        all_kmers,
517    )
518}
519
520// TODO add conditional filters to all SummaryDatas
521
522/// Process DNA sequences into kmers and determine the set of valid kmers,
523/// their extensions, and summarize associated label/'color' data. The input
524/// sequences are converted to kmers of type `K`, and like kmers are grouped together.
525/// All instances of each kmer, along with their label data are then proccessed with
526/// [`SummaryData::summarize`], which generates an implementation of [`SummaryData`],
527/// which is specified with the generic `SD`, decides if the k-mer is 'valid' 
528/// based on the parameters given in `summary_config`, and
529/// summarizes the the individual label into a single label data structure
530/// for the kmer. Care is taken to keep the memory consumption small.
531/// 
532/// Be aware that the configuration in `summary_config` only applies if the required
533/// informaiton can be supplied by the chosen implementation of `SummaryData`.
534/// E.g., the k-mers will not be filtered according to p-value when the `SummaryData`
535/// only contains the number of observations.
536///
537/// # Arguments
538///
539/// * `seqs` are the reads wrapped in a `Reads<u8>`. See [`Reads<DI>`]
540/// * `summary_config` is a [`SummaryConfig`], which contains prameters and 
541///   information necessary for the filtering
542/// * `stranded`: if true, preserve the strandedness of the input sequences, effectively
543///   assuming they are all in the positive strand. If false, the kmers will be canonicalized
544///   to the lexicographic minimum of the kmer and it's reverse complement.
545/// * `report_all_kmers`: if true returns the vector of all the observed kmers and performs the
546///   kmer based filtering
547/// * `memory_size`: gives the size bound on the memory in GB to use and automatically determines
548///   the number of passes needed.
549/// 
550/// # Returns
551/// BoomHashMap2 Object, check rust-boomphf for details
552/// 
553/// # Examples:
554/// 
555/// ```
556/// use debruijn::summarizer::{SampleInfo, SummaryConfig, TagsCountsData, StatTest, GroupFrac};
557/// use debruijn::reads::{Reads, ReadsPaired, Strandedness};
558/// use debruijn::filter::filter_kmers;
559/// use debruijn::kmer::Kmer16;
560/// use debruijn::Exts;
561/// 
562/// let mut seqs = Reads::new(Strandedness::Unstranded);
563/// seqs.add_from_bytes("ACCGATCATATATTTTCGGGGCTAGGCGAAGCGATCTTATCGAGC".as_bytes(), None, 1u8);
564/// seqs.add_from_bytes("GCGATCGAGCATGCTCAGCTGACGTGACTGACGTAGCTATCTTTTCGTAGCTAC".as_bytes(), None, 1u8);
565/// seqs.add_from_bytes("GCGAGTTTGCGACTCGAGGCTATCTAGCTAGCTASGCTCTCGACTAGCTGACTTACGACGACTACG".as_bytes(), None, 2u8);
566/// seqs.add_from_bytes("CGATTAGCTACGTAGCTAGCTGACGTACTGGGGGGTATTTCGGATCTGCGGAGCGATCT".as_bytes(), None, 2u8);
567///       
568/// let sample_info = SampleInfo::new(
569///     0b000011,
570///     0b111100,
571///     vec![23423, 3463454, 2242234, 2233243, 234322434, 2323234],
572/// );
573///     
574/// let summary_config = SummaryConfig::new(sample_info)
575///     .with_min_kmer_obs(3)
576///     .with_group_frac(GroupFrac::One, 0.333)
577///     .with_stat_test(StatTest::StudentsTTest);
578///
579///    
580/// let (hashed_kmers, _) = filter_kmers::<TagsCountsData, Kmer16, _>(
581///     &ReadsPaired::Unpaired { reads: seqs },
582///     &summary_config,
583///     false,
584///     10.,
585///    false,
586/// );
587/// ```
588#[inline(never)]
589pub fn filter_kmers<SD, K, DI>(
590    seqs: &ReadsPaired<DI>,
591    summary_config: &SummaryConfig,
592    report_all_kmers: bool,
593    memory_size: f32,
594    time: bool,
595) -> (BoomHashMap2<K, Exts, SD>, Vec<K>)
596where
597    SD: Debug + SummaryData<DI>, 
598    K: Kmer, 
599    DI: Copy + Clone + Debug + Hash + Eq
600{
601    let before_all = Instant::now();
602
603    // progress bars
604    let multi_pb = MultiProgress::new();
605    let style = ProgressStyle::with_template(PROGRESS_STYLE).unwrap().progress_chars("#/-");
606
607    let pb = multi_pb.add(ProgressBar::new(seqs.n_reads() as u64));
608    pb.set_style(style.clone());
609    pb.set_message(format!("{:<32}", "finding bucket lengths"));
610
611    // go trough all kmers to find the length of all buckets (to reserve capacity)
612    let mut capacities = [0; BUCKETS];
613
614    // also track coverage to predict final graph size
615    let mut unique_kmers = HashSet::<K>::new();
616
617    for ref read in seqs.iter().progress_with(pb)         
618    {
619        // add the required capacities to the respective buckets
620            add_seq_bucket_capacities(read, &mut capacities, &mut unique_kmers);
621    }
622    
623    debug!("kmer capacities: {:?}, times {}", capacities, mem::size_of::<(K, Exts, DI)>());
624    let input_kmers = capacities.iter().sum::<usize>();
625
626    if time { println!("time counting kmers (s): {}", before_all.elapsed().as_secs_f32()) }
627
628    let iter_capacities = capacities.iter().enumerate().map(|(i, &c)| (i, c));
629    let bucket_ranges = bucket_ranges::<K, SD, DI, _>(unique_kmers.len(), input_kmers, memory_size, iter_capacities);
630
631    let n_slices = bucket_ranges.len();
632
633    debug!("n of seqs: {}", seqs.n_reads());
634
635    let mut all_kmers = Vec::new();
636    let mut valid_kmers = Vec::new();
637    let mut valid_exts = Vec::new();
638    let mut valid_data = Vec::new();
639
640    let mut time_collecting = 0.;
641    let mut time_summarizing = 0.;
642
643    if time { println!("time all prepariations before sliced in filter_kmers (s): {}", before_all.elapsed().as_secs_f32()) }
644
645    let pb_bucket_ranges = multi_pb.add(ProgressBar::new(bucket_ranges.len() as u64));
646    pb_bucket_ranges.set_style(style.clone());
647    pb_bucket_ranges.set_message(format!("{:<32}", "filtering k-mers"));
648
649
650    // iterate over the bucket ranges
651    for (i, bucket_range) in bucket_ranges.into_iter().enumerate() {
652        debug!("Processing slice {} of {}", i+1, n_slices);
653
654        let before_kmer_picking = Instant::now();
655        // first step: picking kmers with their exts & data from the reads
656        // go through all kmers and sort into bucket according to first four bases
657        // all kmers starting with "AAAA" go in kmer_buckets[0], all starting with AAAC go in kmer_buckets[1] and so on
658        // when using the first four bases, this needs 256 buckets
659        // the buckets are split in to the bucket_ranges to save memory
660        
661        let mut kmer_buckets = Vec::with_capacity(BUCKETS);
662        // reserve needed capacity in each bucket
663        for (i, capacity) in capacities.iter().enumerate() {
664            if bucket_range.contains(&i) {
665                // bucket is in range, add with prepared capacity
666                kmer_buckets.push(Vec::with_capacity(*capacity));
667            } else {
668                // bucket is not in range, add empty placeholder
669                kmer_buckets.push(Vec::new());
670            }
671            
672        }
673
674        // then go through all kmers and add to bucket according to first four bases and current bucket_range
675        let pb = multi_pb.add(ProgressBar::new(seqs.n_reads() as u64));
676        pb.set_style(style.clone());
677        pb.set_message(format!("{:<32}", "filling buckets with kmers"));
678
679        for ref read in seqs.iter().progress_with(pb)             
680        {
681            // iterate trough all kmers in seq
682            for (kmer, exts, quality) in read.iter_kmer_exts_quality::<K>() {
683                // if needed, flip kmer and exts
684                // check if bucket is in current range and if so, push kmer to bucket
685                if let Some((min_kmer, flip_exts, bucket)) = bucket_ext_flip(kmer, exts, read.stranded(), bucket_range.clone()) {
686                    kmer_buckets[bucket].push(KmerDataItem::new(min_kmer, flip_exts, read.data(), quality));
687                }
688            }
689        }
690
691        time_collecting += before_kmer_picking.elapsed().as_secs_f32();
692
693        let before_summarizing = Instant::now();
694
695        // go trough all buckets and summarize the contents
696        let pb = multi_pb.add(ProgressBar::new(kmer_buckets.len() as u64));
697        pb.set_style(style.clone());
698        pb.set_message(format!("{:<32}", "summarizing k-mers in buckets"));
699
700        for mut kmer_vec in kmer_buckets.into_iter().progress_with(pb) {
701            //debug!("kmers in this bucket: {}", kmer_vec.len());
702            kmer_vec.sort_by_key(|elt| elt.kmer);
703
704            
705            // predict amount of unique k-mers found in this bucket
706            let size = kmer_vec.iter().chunk_by(|elt| elt.kmer).into_iter().count();
707
708            // only works perfectly if min k-mer count is 1, else this might reserve too much 
709            // still better than doubling the vector
710            // also reserve_exact considers pre-existing free capacity -> this might mostly add to runtime
711            valid_kmers.reserve_exact(size);
712            valid_exts.reserve_exact(size);
713            valid_data.reserve_exact(size);
714
715            // if all k-mers should be reported, also grow all_kmers by exact amount -> this should always be the perfect capacity
716            if report_all_kmers {
717                all_kmers.reserve_exact(size);
718            }
719
720            // group the tuples by the k-mers and iterate over the groups
721            for (kmer, kmer_obs_iter) in kmer_vec.into_iter().chunk_by(|elt| elt.kmer).into_iter() {
722                // summarize group with chosen summarizer and add result to vectors
723                let (is_valid, exts, summary_data) = SD::summarize(kmer_obs_iter, summary_config);
724                if report_all_kmers {
725                    all_kmers.push(kmer);
726                }
727                if is_valid {
728                    valid_kmers.push(kmer);
729                    valid_exts.push(exts);
730                    valid_data.push(summary_data); 
731                }
732            }
733        }
734
735        pb_bucket_ranges.inc(1);
736
737        time_summarizing += before_summarizing.elapsed().as_secs_f32();
738
739        debug!("valid kmers - capacity: {}, size: {}, mem: {} Bytes", valid_kmers.capacity(), valid_kmers.len(), mem::size_of_val(&*valid_kmers));
740        debug!("valid exts - capacity: {}, size: {}, mem: {} Bytes", valid_exts.capacity(), valid_exts.len(), mem::size_of_val(&*valid_exts));
741        debug!("valid data - capacity: {}, size: {}, mem: {} Bytes", valid_data.capacity(), valid_data.len(), mem::size_of_val(&*valid_data));
742    }
743
744    pb_bucket_ranges.finish_and_clear();
745
746    if time { 
747        println!("time collecting (s): {}", time_collecting);
748        println!("time summarizing (s): {}", time_summarizing);
749    }
750
751
752    debug!(
753        "Unique kmers: {}, All kmers (if returned): {}",
754        valid_kmers.len(),
755        all_kmers.len(),
756    );
757
758    debug!("size of valid kmers: {} Bytes
759        size of valid exts: {} Bytes
760        size of valid data: {} Bytes", mem::size_of_val(&*valid_kmers), mem::size_of_val(&*valid_exts), mem::size_of_val(&*valid_data));
761
762    let before_hash = Instant::now();
763    let hm = BoomHashMap2::new(valid_kmers, valid_exts, valid_data);
764    let after_hash = before_hash.elapsed().as_secs_f32();
765
766    let filter_kmers_inner = before_all.elapsed().as_secs_f32();
767    if time { 
768        println!("time filter_kmers inner (s): {}", filter_kmers_inner);
769        println!("time build filter hash map (s): {}", after_hash);
770    }
771    (
772        hm,
773        all_kmers,
774    )
775}
776
777/// Remove extensions in valid_kmers that point to censored kmers. A censored kmer
778/// exists in `all_kmers` but not `valid_kmers`. Since the kmer exists in this partition,
779/// but was censored, we know that we can delete extensions to it.
780/// In sharded kmer processing, we will have extensions to kmers in other shards. We don't
781/// know whether these are censored until later, so we retain these extension.
782pub fn remove_censored_exts_sharded<K: Kmer, D>(
783    stranded: bool,
784    valid_kmers: &mut [(K, (Exts, D))],
785    all_kmers: &[K],
786) {
787    for idx in 0..valid_kmers.len() {
788        let mut new_exts = Exts::empty();
789        let kmer = valid_kmers[idx].0;
790        let exts = (valid_kmers[idx].1).0;
791
792        for dir in [Dir::Left, Dir::Right].iter() {
793            for i in 0..4 {
794                if exts.has_ext(*dir, i) {
795                    let _ext_kmer = kmer.extend(i, *dir);
796
797                    let ext_kmer = if stranded {
798                        _ext_kmer
799                    } else {
800                        _ext_kmer.min_rc()
801                    };
802
803                    let censored = if valid_kmers.binary_search_by_key(&ext_kmer, |d| d.0).is_ok() {
804                        // ext_kmer is valid. not censored
805                        false
806                    } else {
807                        // ext_kmer is not valid. if it was in this shard, then we censor it
808                        all_kmers.binary_search(&ext_kmer).is_ok()
809                    };
810
811                    if !censored {
812                        new_exts = new_exts.set(*dir, i);
813                    }
814                }
815            }
816        }
817
818        (valid_kmers[idx].1).0 = new_exts;
819    }
820}
821
822/// Remove extensions in valid_kmers that point to censored kmers. Use this method in a non-partitioned
823/// context when valid_kmers includes _all_ kmers that will ultimately be included in the graph.
824pub fn remove_censored_exts<K: Kmer, D>(stranded: bool, valid_kmers: &mut [(K, (Exts, D))]) {
825    for idx in 0..valid_kmers.len() {
826        let mut new_exts = Exts::empty();
827        let kmer = valid_kmers[idx].0;
828        let exts = (valid_kmers[idx].1).0;
829
830        for dir in [Dir::Left, Dir::Right].iter() {
831            for i in 0..4 {
832                if exts.has_ext(*dir, i) {
833                    let ext_kmer = if stranded {
834                        kmer.extend(i, *dir)
835                    } else {
836                        kmer.extend(i, *dir).min_rc()
837                    };
838
839                    let kmer_valid = valid_kmers.binary_search_by_key(&ext_kmer, |d| d.0).is_ok();
840
841                    if kmer_valid {
842                        new_exts = new_exts.set(*dir, i);
843                    }
844                }
845            }
846        }
847
848        (valid_kmers[idx].1).0 = new_exts;
849    }
850}
851
852#[cfg(test)]
853mod tests {
854    use boomphf::hashmap::BoomHashMap2;
855    use crate::{dna_string::DnaString, filter::*, kmer::{Kmer2, Kmer6}, reads::Reads, summarizer::{SampleInfo, TagsSumData}, test::{random_dna, random_kmer}, Exts};
856
857    #[test]
858    fn test_filter_kmers() {
859        let fastq = [
860            (DnaString::from_dna_string("AAAAATTT"), Exts::empty(), 6u8),
861            (DnaString::from_dna_string("TTTTTTTTTTAAAAAA"), Exts::empty(), 6u8),
862            (DnaString::from_dna_string("AAAAAAAAAAAAA"), Exts::empty(), 7u8),
863        ];
864
865        let reads = Reads::from_vmer_vec(fastq, crate::reads::Strandedness::Unstranded);
866
867        let sample_info = SampleInfo::new(0, 0, Vec::new());
868
869        let config = SummaryConfig::new(sample_info).with_stat_test(crate::summarizer::StatTest::StudentsTTest);
870
871
872        let (hm, _): (BoomHashMap2<Kmer6, Exts, TagsSumData>, Vec<_>) = filter_kmers(
873            &ReadsPaired::Unpaired { reads }, 
874            &config,
875            false, 
876            1.,
877            false,
878         );
879
880         println!("{:?}", hm);
881
882    }
883
884    #[test]
885    fn test_filter_kmers_parallel() {
886        /* let fastq = [
887            (DnaString::from_dna_string("AAAAATTT"), Exts::empty(), 6u8),
888            (DnaString::from_dna_string("TTTTTTTTTTAAAAAA"), Exts::empty(), 6u8),
889            (DnaString::from_dna_string("AAAAAAAAAAAAA"), Exts::empty(), 7u8),
890        ];
891
892        let mut reads = Reads::new();
893
894        for (read, exts, data) in fastq {
895            reads.add_read(read, exts, data);
896
897        } */
898
899        let mut reads = Reads::new(crate::reads::Strandedness::Unstranded);
900
901        for _i in 0..10000 {
902            let dna = random_dna(150);
903            reads.add_from_bytes(&dna, None, 0u8);
904        }
905
906        let sample_info = SampleInfo::new(0, 0, Vec::new());
907        let config = SummaryConfig::new(sample_info.clone()).with_stat_test(crate::summarizer::StatTest::StudentsTTest);
908
909
910        let (hm, _): (BoomHashMap2<Kmer6, Exts, TagsSumData>, Vec<_>) = filter_kmers_parallel(
911            &ReadsPaired::Unpaired { reads }, 
912            &config,
913            false, 
914            1.,
915            false,         
916        );
917
918        println!("{:?}", hm);
919
920    }
921
922
923    #[test]
924    fn test_par_hmap_combination() {
925        let n_threads = rayon::current_num_threads();
926        let kmers = Arc::new(Mutex::new(vec![HashSet::new(); n_threads]));
927        
928        (0..n_threads).into_par_iter().for_each(|i| {
929            let mut ks = HashSet::new();
930            for _j in 0..100 {
931                let kmer = random_kmer::<Kmer2>();
932                ks.insert(kmer);
933            }
934            let mut c = kmers.lock().expect("cov lock");
935            c[i] = ks;
936        });
937
938        println!("k-mers uncombinded: {:?}", kmers);
939
940        let mut kmers = kmers.lock().expect("final lock coverages");
941
942        let mut combined_kmers = kmers.pop().expect("no kmers");
943        while  let Some(kmap) = kmers.pop() {
944                combined_kmers.extend(kmap);
945        }
946
947        assert_eq!(combined_kmers.len(), 16);
948
949        println!("combined k-mers: {:?}", combined_kmers)
950    }
951}
952