diff --git a/bam_binary/Cargo.toml b/bam_binary/Cargo.toml index c9327f3..e9371cc 100644 --- a/bam_binary/Cargo.toml +++ b/bam_binary/Cargo.toml @@ -11,4 +11,5 @@ structopt = "0.3.21" bam_tools = {path = "../bam_tools"} md-5 = "0.9.1" byteorder = "1.2.3" -num_cpus = "0.2" \ No newline at end of file +num_cpus = "0.2" +tempdir = "0.3.7" \ No newline at end of file diff --git a/bam_binary/src/main.rs b/bam_binary/src/main.rs index 5adc94a..21deff0 100644 --- a/bam_binary/src/main.rs +++ b/bam_binary/src/main.rs @@ -4,6 +4,7 @@ use bam_tools::Reader; use bam_tools::MEGA_BYTE_SIZE; use md5::{Digest, Md5}; use std::env; +use tempdir::TempDir; use std::fs::File; use std::io::BufReader; @@ -61,12 +62,13 @@ fn main() { let out_file = File::create(opt.output.unwrap()).unwrap(); let mut writer = BufWriter::new(out_file); let tmp_dir_path = env::temp_dir(); + let dir = TempDir::new_in(tmp_dir_path, "BAM sort temporary directory.").unwrap(); sort_bam( 2000 * MEGA_BYTE_SIZE, reader, &mut writer, - tmp_dir_path, + &dir, 0, 5, bam_tools::sorting::sort::TempFilesMode::RegularFiles, @@ -77,7 +79,7 @@ fn main() { std::process::exit(0); } -fn generate_file_hash(reader: R) -> String { +fn generate_file_hash(reader: R) -> String { let mut bgzf_reader = Reader::new(reader, std::cmp::min(num_cpus::get(), 20)); let mut hasher = Md5::new(); diff --git a/bam_tools/src/reader.rs b/bam_tools/src/reader.rs index 211dfe5..e2a2e4d 100644 --- a/bam_tools/src/reader.rs +++ b/bam_tools/src/reader.rs @@ -6,6 +6,7 @@ use crate::MAGIC_NUMBER; use byteorder::{LittleEndian, ReadBytesExt, WriteBytesExt}; use std::ffi::CStr; use std::io; +use std::sync::Arc; use readahead::Readahead; use records::Records; @@ -19,7 +20,7 @@ pub struct Reader { } impl Reader { - pub fn new(inner: RSS, mut thread_num: usize) -> Self { + pub fn new(inner: RSS, mut thread_num: usize) -> Self { if thread_num > num_cpus::get() { thread_num = num_cpus::get(); } diff --git a/bam_tools/src/reader/readahead.rs b/bam_tools/src/reader/readahead.rs index 8d975ef..13f3291 100644 --- a/bam_tools/src/reader/readahead.rs +++ b/bam_tools/src/reader/readahead.rs @@ -3,141 +3,105 @@ use crate::util::{fetch_block, inflate_data}; // This module preparses BAM blocks to parallelize decompression use flume::{Receiver, Sender}; -use rayon::spawn; -use std::cmp::{Ord, Ordering, PartialEq, PartialOrd}; -use std::collections::BinaryHeap; use std::io::Read; +use std::sync::{Arc, Condvar, Mutex}; +use std::thread::{self, JoinHandle}; -#[allow(clippy::upper_case_acronyms)] -enum Status { - Success(Block), - EOF, -} - -struct Task(usize, Status); - -impl Ord for Task { - fn cmp(&self, other: &Self) -> Ordering { - // Smallest go first. - other.0.cmp(&self.0) - } -} - -impl PartialOrd for Task { - fn partial_cmp(&self, other: &Self) -> Option { - // Smallest go first. - Some(self.cmp(other)) - } -} - -impl Eq for Task {} - -impl PartialEq for Task { - fn eq(&self, other: &Self) -> bool { - // There shouldn't be two WorkUnits with the same number in the Heap - assert_ne!(self.0, other.0); - false - } -} - +type VectorOfSendersAndReceivers = Vec<(Sender, Receiver)>; /// Prefetches and decompresses GBAM blocks pub(crate) struct Readahead { - // Decompressing threadpool. - used_block_sender: Sender, - ready_to_processing_rx: Receiver, + circular_buf_channels: VectorOfSendersAndReceivers, + handles: Vec>, + current_task: usize, } impl Readahead { - pub fn new(mut thread_num: usize, mut reader: Box) -> Self { - // No less than 3 threads to avoid deadlock. - thread_num = std::cmp::max(thread_num, 3); - let (read_bufs_send, read_bufs_recv) = flume::unbounded(); - let (used_block_sender, used_block_receiver) = flume::unbounded(); - let (completed_task_tx, sorting_blocks_rx) = flume::unbounded(); - let (ready_tasks_tx, ready_to_processing_rx) = flume::unbounded(); - for _ in 0..thread_num { - read_bufs_send.send(Vec::new()).unwrap(); - used_block_sender.send(Block::default()).unwrap(); - } - let pool = rayon::ThreadPoolBuilder::new() - .num_threads(thread_num) - .build() - .unwrap(); - - // Ordering thread. - pool.spawn(move || { - // The heap is needed for cases when the blocks are not inflated in - // proper order (as coming from input stream). - let mut block_heap = BinaryHeap::::new(); - // Number of current block (ordered as read from input stream). - let mut cur_block_num = 0; - while let Ok(work_unit) = sorting_blocks_rx.recv() { - block_heap.push(work_unit); - // Fill queue with parsed blocks. - while !block_heap.is_empty() - && (block_heap.peek().unwrap().0 == cur_block_num) - && !ready_tasks_tx.is_disconnected() - { - ready_tasks_tx.send((block_heap.pop().unwrap()).1).unwrap(); - // The block is extracted. Wait for next one. - cur_block_num += 1; - } - } - }); - // Reading thread. - pool.spawn(move || { - let mut cur_task: usize = 0; - while let Ok(mut block) = used_block_receiver.recv() { - let mut read_buf = read_bufs_recv.recv().unwrap(); - let bytes_count = fetch_block(&mut reader, &mut read_buf, &mut block).unwrap(); - - let task_ready_to_sort_tx = completed_task_tx.clone(); - if bytes_count == 0 { - task_ready_to_sort_tx - .send(Task(cur_task, Status::EOF)) + pub fn new(mut thread_num: usize, reader: Box) -> Self { + thread_num = std::cmp::max(thread_num, 1); + + let mut circular_buf_channels = VectorOfSendersAndReceivers::new(); + + let mut handles: Vec> = Vec::new(); + + // Due to I/O unpredictable nature, it may happen that two or more threads would race to lock a mutex on a reader + // making uncompressed block stream unordered. Ensure order with condvar. + let cond_var_for_order: Arc<(Mutex, Condvar)> = + Arc::new((Mutex::new(0), Condvar::new())); + + let mutex_protected_reader = Arc::new(Mutex::new(reader)); + for i in 0..thread_num { + let (block_sender, block_receiver) = flume::bounded(1); + let (uncompressed_sender, uncompressed_receiver) = flume::bounded(1); + let clone_of_reader = mutex_protected_reader.clone(); + let cond_var_for_this_thread = cond_var_for_order.clone(); + + let thread = thread::spawn(move || { + let mut read_buf = Vec::new(); + + for mut block in block_receiver { + { + // Only one thread fetches data from file at a time. + let (lock, cvar) = &*cond_var_for_this_thread; + let mut my_turn = lock.lock().unwrap(); + + while *my_turn != i { + my_turn = cvar.wait(my_turn).unwrap(); + } + + let bytes_count = fetch_block( + clone_of_reader.lock().unwrap().as_mut(), + &mut read_buf, + &mut block, + ) .unwrap(); - // Reached EOF - return; - } - - let read_buf_sender = read_bufs_send.clone(); - spawn(move || { - decompress_block(&read_buf, &mut block); - task_ready_to_sort_tx - .send(Task(cur_task, Status::Success(block))) - .unwrap(); - if !read_buf_sender.is_disconnected() { - read_buf_sender.send(read_buf).unwrap(); + *my_turn += 1; + if *my_turn == thread_num { + *my_turn = 0; + } + cvar.notify_all(); + if bytes_count == 0 { + // Reached EOF. + return; + } + // After this line the mutex lock will be dropped, and the decompression will happen in parallel to other threads. } - }); - cur_task += 1; - } - }); + decompress_block(&read_buf, &mut block); + uncompressed_sender.send(block).unwrap(); + } + }); + block_sender.send(Block::default()).unwrap(); + handles.push(thread); + circular_buf_channels.push((block_sender, uncompressed_receiver)); + } Self { - used_block_sender, - ready_to_processing_rx, + circular_buf_channels, + handles, + current_task: 0, } } /// Receives prefetched block. This is a blocking function. In case there is - /// no uncompressed blocks in queue, the thread which called it will be + /// no uncompressed blocks in the queue, the thread which called it will be /// blocked until uncompressed buffer appears. pub fn get_block(&mut self, old_buf: Block) -> Option { - // eprintln!("3.6."); - if !self.used_block_sender.is_disconnected() { - // Ignore even if it errs. Even though the check has been passed at - // this point the threads might have been already terminated, so it - // will err on send attempt (no available receivers). - let _ = self.used_block_sender.send(old_buf); + if self.current_task == self.circular_buf_channels.len() { + self.current_task = 0; } - // eprintln!("3.7."); - match self.ready_to_processing_rx.recv().unwrap() { - Status::Success(block) => Some(block), - Status::EOF => None, + let cur_thread = &mut self.circular_buf_channels[self.current_task]; + let res = cur_thread.1.recv(); + // The thread reached EOF. + if res.is_err() { + // Join all of the other threads. The other threads should also reach the EOF by then. + for j in self.handles.drain(0..) { + j.join().unwrap(); + } + return None; } - // eprintln!("3.8."); + cur_thread.0.send(old_buf).unwrap(); + self.current_task += 1; + return Some(res.unwrap()); } } diff --git a/bam_tools/src/sorting/sort.rs b/bam_tools/src/sorting/sort.rs index 510b3f1..0a79da5 100644 --- a/bam_tools/src/sorting/sort.rs +++ b/bam_tools/src/sorting/sort.rs @@ -106,7 +106,7 @@ pub enum TempFilesMode { /// Memory limit won't be strictly obeyed, but it probably won't be overflowed significantly. #[allow(clippy::too_many_arguments)] -pub fn sort_bam( +pub fn sort_bam( mem_limit: usize, reader: R, sorted_sink: &mut W,