1use rayon::prelude::*;
9
10use crate::curve::iters::{CurveIter, SymCurveIterator};
11use crate::curve::matrix::RollType;
12use crate::fasta::RecordPiece;
13
14#[derive(Debug, Clone)]
21pub struct PieceCurves {
22 pub start: usize,
23 pub curve_start: usize,
24 pub curves: Vec<f64>,
25}
26
27#[derive(Debug, Clone, Copy)]
32pub struct CurveParams {
33 pub roll_type: RollType,
34 pub step_b: usize,
36 pub step_c: usize,
38 pub curve_scale: f64,
39}
40
41impl CurveParams {
42 pub fn lead_in(&self) -> usize {
50 self.step_b + self.step_c + 1
51 }
52}
53
54#[derive(Debug, Clone, Copy, PartialEq, Eq)]
56pub enum Stage {
57 Curvature,
59 Symmetry,
61}
62
63#[derive(Debug, Clone, Copy)]
65pub struct SymParams {
66 pub win: usize,
68 pub step: usize,
70}
71
72impl SymParams {
73 pub fn lead_in(&self, params: &CurveParams) -> usize {
78 params.lead_in() + self.win
79 }
80
81 pub fn score_count(&self, piece_len: usize, params: &CurveParams) -> usize {
83 let curve_len = piece_len.saturating_sub(2 * params.lead_in());
84 let dyads = curve_len.saturating_sub(2 * self.win);
85 dyads.div_ceil(self.step.max(1))
86 }
87}
88
89pub const DEFAULT_CHUNK_SCORES: usize = 1 << 20;
94
95#[derive(Debug, Clone, Copy, PartialEq, Eq)]
104pub struct Chunk {
105 pub in_start: usize,
107 pub in_end: usize,
109 pub out_start: usize,
111}
112
113pub fn chunks(piece_len: usize, params: &CurveParams, chunk_scores: usize) -> Vec<Chunk> {
117 let lead = params.lead_in();
118 let out_len = piece_len.saturating_sub(2 * lead);
119 let chunk_scores = chunk_scores.max(1);
120 (0..out_len)
121 .step_by(chunk_scores)
122 .map(|out_start| {
123 let out_end = (out_start + chunk_scores).min(out_len);
124 Chunk {
125 in_start: out_start,
127 in_end: out_end + 2 * lead,
128 out_start,
129 }
130 })
131 .collect()
132}
133
134pub fn sym_chunks(
141 piece_len: usize,
142 params: &CurveParams,
143 sym: &SymParams,
144 chunk_scores: usize,
145) -> Vec<Chunk> {
146 let step = sym.step.max(1);
147 let lead = params.lead_in();
148 let out_len = sym.score_count(piece_len, params);
149 let chunk_scores = chunk_scores.max(1);
150 (0..out_len)
151 .step_by(chunk_scores)
152 .map(|out_start| {
153 let out_end = (out_start + chunk_scores).min(out_len);
154 Chunk {
155 in_start: out_start * step,
156 in_end: (out_end - 1) * step + 2 * sym.win + 2 * lead + 1,
157 out_start,
158 }
159 })
160 .collect()
161}
162
163pub fn score_symmetry_bases(bases: &[u8], params: &CurveParams, sym: &SymParams) -> Vec<f64> {
165 CurveIter::new(
166 bases.iter().copied(),
167 params.roll_type,
168 params.step_b,
169 params.step_c,
170 params.curve_scale,
171 )
172 .sym_curve_iter(sym.win, sym.step)
173 .collect()
174}
175
176pub fn score_bases(bases: &[u8], params: &CurveParams) -> Vec<f64> {
178 CurveIter::new(
179 bases.iter().copied(),
180 params.roll_type,
181 params.step_b,
182 params.step_c,
183 params.curve_scale,
184 )
185 .collect()
186}
187
188pub fn score_piece(piece: &RecordPiece, params: &CurveParams) -> PieceCurves {
190 score_piece_chunked(piece, params, DEFAULT_CHUNK_SCORES)
191}
192
193pub fn score_piece_chunked(
199 piece: &RecordPiece,
200 params: &CurveParams,
201 chunk_scores: usize,
202) -> PieceCurves {
203 let bases = piece.bases();
205 let per_chunk: Vec<Vec<f64>> = chunks(bases.len(), params, chunk_scores)
208 .into_par_iter()
209 .map(|chunk| score_bases(&bases[chunk.in_start..chunk.in_end], params))
210 .collect();
211 let curves: Vec<f64> = per_chunk.concat();
212 let start = usize::from(piece.start);
213 PieceCurves {
214 start,
215 curve_start: start + params.lead_in(),
216 curves,
217 }
218}
219
220pub fn score_pieces(pieces: &[RecordPiece], params: &CurveParams) -> Vec<PieceCurves> {
226 pieces
227 .par_iter()
228 .map(|piece| score_piece(piece, params))
229 .collect()
230}
231
232#[cfg(test)]
233mod tests {
234 use super::*;
235 use crate::fasta::split_seq_by_gaps;
236 use approx::assert_relative_eq;
237
238 fn params() -> CurveParams {
239 CurveParams {
240 roll_type: RollType::Simple,
241 step_b: 5,
242 step_c: 15,
243 curve_scale: 0.33335,
244 }
245 }
246
247 fn record_of(seq: &[u8]) -> noodles_fasta::Record {
248 let mut src = b">chr42\n".to_vec();
249 src.extend_from_slice(seq);
250 src.push(b'\n');
251 let mut reader = noodles_fasta::io::Reader::new(&src[..]);
252 reader.records().next().unwrap().unwrap()
253 }
254
255 #[test]
259 fn test_record_piece_is_send_and_sync() {
260 fn assert_send_sync<T: Send + Sync>() {}
261 assert_send_sync::<RecordPiece>();
262 assert_send_sync::<PieceCurves>();
263 assert_send_sync::<CurveParams>();
264 }
265
266 #[test]
267 fn test_parallel_matches_serial() {
268 let unit = b"CCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATC";
270 let mut seq = Vec::new();
271 for reps in [3usize, 1, 5, 2] {
272 for _ in 0..reps {
273 seq.extend_from_slice(unit);
274 }
275 seq.extend_from_slice(b"NNNN");
276 }
277 let pieces = split_seq_by_gaps(record_of(&seq));
278 assert_eq!(pieces.len(), 4);
279
280 let parallel = score_pieces(&pieces, ¶ms());
281 let serial: Vec<PieceCurves> = pieces.iter().map(|p| score_piece(p, ¶ms())).collect();
282
283 assert_eq!(parallel.len(), serial.len());
284 for (par, ser) in parallel.iter().zip(&serial) {
285 assert_eq!(par.start, ser.start);
287 assert_eq!(par.curves.len(), ser.curves.len());
288 for (a, b) in par.curves.iter().zip(&ser.curves) {
289 assert_relative_eq!(a, b, epsilon = 1e-12);
290 }
291 }
292 }
293
294 #[test]
295 fn test_piece_start_positions_are_reported() {
296 let seq = b"CCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATCNNNNCCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATC";
297 let pieces = split_seq_by_gaps(record_of(seq));
298 let scored = score_pieces(&pieces, ¶ms());
299 assert_eq!(scored.len(), 2);
300 assert_eq!(scored[0].start, 1);
301 assert_eq!(scored[1].start, 55);
302 }
303
304 #[test]
305 fn test_lead_in_matches_the_reference_convention() {
306 let p = params();
311 assert_eq!(p.lead_in(), 21); let unit = b"CCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATC";
313 let mut seq = Vec::new();
314 for _ in 0..4 {
315 seq.extend_from_slice(unit);
316 }
317 let n = seq.len();
318 let pieces = split_seq_by_gaps(record_of(&seq));
319 let scored = score_pieces(&pieces, &p);
320 assert_eq!(scored.len(), 1);
321 assert_eq!(scored[0].curves.len(), n - 2 * p.lead_in());
322 assert_eq!(scored[0].start, 1);
323 assert_eq!(scored[0].curve_start, 1 + p.lead_in());
324 }
325
326 #[test]
327 fn test_curve_start_is_offset_from_the_piece_not_the_record() {
328 let seq = b"NNNNNCCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATCCCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATC";
329 let pieces = split_seq_by_gaps(record_of(seq));
330 assert_eq!(pieces.len(), 1);
331 let scored = score_pieces(&pieces, ¶ms());
332 assert_eq!(scored[0].start, 6); assert_eq!(scored[0].curve_start, 6 + params().lead_in());
334 }
335
336 #[test]
337 fn test_chunk_ranges_tile_the_output_exactly() {
338 let p = params();
339 let lead = p.lead_in();
340 let piece_len = 1000usize;
341 let out_len = piece_len - 2 * lead;
342 let cs = 100usize;
343 let cs_list = chunks(piece_len, &p, cs);
344 assert_eq!(cs_list.len(), out_len.div_ceil(cs));
345 let mut expected_out = 0usize;
347 for c in &cs_list {
348 assert_eq!(c.out_start, expected_out);
349 let produced = c.in_end - c.in_start - 2 * lead;
351 expected_out += produced;
352 assert!(c.in_end <= piece_len, "chunk reads past the piece");
353 }
354 assert_eq!(
355 expected_out, out_len,
356 "chunks do not cover the output exactly"
357 );
358 }
359
360 #[test]
361 fn test_short_piece_produces_no_chunks() {
362 let p = params();
363 for len in [0usize, 1, 2 * p.lead_in(), 2 * p.lead_in() + 1] {
364 let got = chunks(len, &p, 100);
365 let expect_scores = len.saturating_sub(2 * p.lead_in());
366 assert_eq!(got.is_empty(), expect_scores == 0, "len {len}");
367 }
368 }
369
370 #[test]
371 fn test_chunked_scoring_matches_unchunked() {
372 let unit = b"CCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATC";
376 let mut seq = Vec::new();
377 for _ in 0..40 {
378 seq.extend_from_slice(unit);
379 }
380 let pieces = split_seq_by_gaps(record_of(&seq));
381 let piece = &pieces[0];
382 let p = params();
383
384 let reference = score_bases(piece.bases(), &p);
385 assert!(reference.len() > 1000);
386
387 for chunk_scores in [
388 1usize,
389 7,
390 64,
391 500,
392 1000,
393 reference.len(),
394 reference.len() * 2,
395 ] {
396 let got = score_piece_chunked(piece, &p, chunk_scores);
397 assert_eq!(
398 got.curves.len(),
399 reference.len(),
400 "length differs at chunk_scores {chunk_scores}"
401 );
402 for (i, (a, b)) in reference.iter().zip(&got.curves).enumerate() {
403 assert_relative_eq!(a, b, epsilon = 1e-9, max_relative = 1e-9);
406 let _ = i;
407 }
408 }
409 }
410
411 fn sym() -> SymParams {
412 SymParams { win: 20, step: 1 }
413 }
414
415 #[test]
416 fn test_sym_chunks_tile_the_output_exactly() {
417 let p = params();
418 for step in [1usize, 3, 5] {
419 let sp = SymParams { win: 20, step };
420 let piece_len = 2000usize;
421 let out_len = sp.score_count(piece_len, &p);
422 assert!(out_len > 0);
423 for chunk_scores in [1usize, 7, 64, 500, out_len, out_len * 2] {
424 let plan = sym_chunks(piece_len, &p, &sp, chunk_scores);
425 assert_eq!(plan.len(), out_len.div_ceil(chunk_scores.max(1)));
426 let mut expected = 0usize;
427 for c in &plan {
428 assert_eq!(c.out_start, expected);
429 assert!(c.in_end <= piece_len, "chunk reads past the piece");
430 let produced = sp.score_count(c.in_end - c.in_start, &p);
432 expected += produced;
433 }
434 assert_eq!(expected, out_len, "step {step} chunk {chunk_scores}");
435 }
436 }
437 }
438
439 #[test]
440 fn test_chunked_symmetry_matches_unchunked() {
441 let unit = b"CCAACATTTTGACTTTTTGGGAGGGCACTAGCACCTATCTACCCTGAATC";
443 let mut seq = Vec::new();
444 for _ in 0..60 {
445 seq.extend_from_slice(unit);
446 }
447 let pieces = split_seq_by_gaps(record_of(&seq));
448 let piece = &pieces[0];
449 let p = params();
450
451 for step in [1usize, 3] {
452 let sp = SymParams { win: 20, step };
453 let reference = score_symmetry_bases(piece.bases(), &p, &sp);
454 assert!(reference.len() > 500, "input too small to be meaningful");
455
456 for chunk_scores in [1usize, 7, 64, 500, reference.len(), reference.len() * 2] {
457 let plan = sym_chunks(piece.bases().len(), &p, &sp, chunk_scores);
458 let rebuilt: Vec<f64> = plan
459 .iter()
460 .flat_map(|c| {
461 score_symmetry_bases(&piece.bases()[c.in_start..c.in_end], &p, &sp)
462 })
463 .collect();
464 assert_eq!(
465 rebuilt.len(),
466 reference.len(),
467 "length differs at step {step} chunk {chunk_scores}"
468 );
469 for (a, b) in reference.iter().zip(&rebuilt) {
470 assert_relative_eq!(a, b, epsilon = 1e-9, max_relative = 1e-6);
472 }
473 }
474 }
475 }
476
477 #[test]
478 fn test_symmetry_lead_in_and_count() {
479 let p = params();
480 let sp = sym();
481 assert_eq!(sp.lead_in(&p), p.lead_in() + sp.win);
483 for len in [0usize, 1, 2 * sp.lead_in(&p), 2 * sp.lead_in(&p) + 1, 5000] {
484 let expected = len.saturating_sub(2 * sp.lead_in(&p));
485 assert_eq!(sp.score_count(len, &p), expected, "len {len}");
486 }
487 }
488
489 #[test]
490 fn test_empty_input_is_not_an_error() {
491 let scored = score_pieces(&[], ¶ms());
492 assert!(scored.is_empty());
493 }
494
495 #[test]
496 fn test_piece_shorter_than_the_window_yields_nothing() {
497 let pieces = split_seq_by_gaps(record_of(b"ACGTACGT"));
499 assert_eq!(pieces.len(), 1);
500 let scored = score_pieces(&pieces, ¶ms());
501 assert_eq!(scored.len(), 1);
502 assert!(scored[0].curves.is_empty());
503 }
504}