Skip to main content

detcore/edit_distance/
needleman_wunsch.rs

1/*
2 * Copyright (c) Meta Platforms, Inc. and affiliates.
3 * All rights reserved.
4 *
5 * This source code is licensed under the BSD-style license found in the
6 * LICENSE file in the root directory of this source tree.
7 */
8
9use std::cmp::min;
10
11type IndexAndAlignments = (usize, Vec<(Option<usize>, Option<usize>)>);
12
13#[derive(PartialEq, Eq, Clone, Copy, Debug)]
14pub enum Trace {
15    Stop,
16    Left,
17    Up,
18    MatchDiagonal,
19    MisMatchDiagonal,
20}
21
22#[derive(PartialEq, Eq, Debug)]
23pub enum NeedlemanWunschError {
24    MatrixTooLarge { cells: usize, max_cells: usize },
25}
26
27fn simple_score<T: PartialEq>(a: T, b: T) -> i32 {
28    if a == b { 1 } else { -5 }
29}
30
31pub struct NeedlemanWunsch<T: PartialEq + Clone> {
32    pub first_sequence: Vec<T>,
33    pub second_sequence: Vec<T>,
34    pub first_extra: Vec<T>,
35    pub second_extra: Vec<T>,
36    pub mismatches: Vec<(usize, usize)>,
37    pub num_mismatches: usize,
38}
39
40impl<T: PartialEq + Clone> Default for NeedlemanWunsch<T> {
41    fn default() -> NeedlemanWunsch<T> {
42        NeedlemanWunsch {
43            first_sequence: vec![],
44            second_sequence: vec![],
45            first_extra: vec![],
46            second_extra: vec![],
47            mismatches: vec![],
48            num_mismatches: 0,
49        }
50    }
51}
52
53impl<T: PartialEq + Clone> NeedlemanWunsch<T> {
54    pub fn match_sequences_base(&mut self) -> IndexAndAlignments {
55        self.match_sequences(None, None, None)
56    }
57
58    /// Take the first alignment in NeedlemanWunsch struct and changes
59    /// it using the mismatches to get closer to the second alignment
60    pub fn generate_midpoint_schedule(
61        &mut self,
62        start_index: usize,
63        alignment_difference: Vec<(Option<usize>, Option<usize>)>,
64    ) -> Vec<T> {
65        let mut midpoint_schedule: Vec<T> = vec![];
66        let total_swaps = self.num_mismatches / 2;
67        let available_full_swaps = self.mismatches.len();
68        let mut current_swaps = 0;
69        let mut full_swaps = 0;
70        {
71            let mut i = 0;
72            while i < start_index {
73                midpoint_schedule.push(self.first_sequence[i].clone());
74                i += 1;
75            }
76        }
77        {
78            let mut i = 0;
79            while i < alignment_difference.len() {
80                if alignment_difference[i].0.is_none() && alignment_difference[i].1.is_none() {
81                    if full_swaps < available_full_swaps && current_swaps < total_swaps {
82                        midpoint_schedule
83                            .push(self.first_sequence[self.mismatches[full_swaps].0].clone());
84                    } else if full_swaps < available_full_swaps {
85                        midpoint_schedule
86                            .push(self.second_sequence[self.mismatches[full_swaps].1].clone());
87                    }
88                    current_swaps += 1;
89                    full_swaps += 1;
90                } else if alignment_difference[i].0.is_some() && current_swaps < total_swaps {
91                    midpoint_schedule
92                        .push(self.first_sequence[alignment_difference[i].0.unwrap()].clone());
93                    current_swaps += 1;
94                } else if alignment_difference[i].1.is_some() && current_swaps >= total_swaps {
95                    midpoint_schedule
96                        .push(self.second_sequence[alignment_difference[i].1.unwrap()].clone());
97                    current_swaps += 1;
98                }
99                i += 1;
100            }
101        }
102        midpoint_schedule
103    }
104
105    /// Globally aligns the two sequences in NeedlemanWunsch struct.
106    /// Returns: (1) Index till which both sequences are complete same
107    ///          (2) Vector of (Matching index from sequence 1,
108    ///                         Matching index from sequence 2)
109    /// None, None represents no match - these are stored mismatch vector
110    pub fn match_sequences(
111        &mut self,
112        scoring_function: Option<fn(T, T) -> i32>,
113        gap_penalty: Option<i32>,
114        mismatch_penalty: Option<i32>,
115    ) -> IndexAndAlignments {
116        self.match_sequences_bounded(scoring_function, gap_penalty, mismatch_penalty, usize::MAX)
117            .expect("Needleman-Wunsch matrix dimensions overflowed")
118    }
119
120    /// Globally align the sequences, refusing work above `max_matrix_cells` before allocating the
121    /// traceback matrix. Scores use two rolling rows, so only traceback storage remains quadratic.
122    pub fn match_sequences_bounded(
123        &mut self,
124        scoring_function: Option<fn(T, T) -> i32>,
125        gap_penalty: Option<i32>,
126        mismatch_penalty: Option<i32>,
127        max_matrix_cells: usize,
128    ) -> Result<IndexAndAlignments, NeedlemanWunschError> {
129        let gap_penalty: i32 = gap_penalty.unwrap_or(-1);
130        let mismatch_penalty: i32 = mismatch_penalty.unwrap_or(-2);
131        let scoring_function = scoring_function.unwrap_or(simple_score);
132
133        self.mismatches.clear();
134        self.num_mismatches = 0;
135
136        let mut start_index = 0;
137        let max_length: usize = min(self.first_sequence.len(), self.second_sequence.len());
138        while start_index < max_length
139            && scoring_function(
140                self.first_sequence[start_index].clone(),
141                self.second_sequence[start_index].clone(),
142            ) > 0
143        {
144            start_index += 1;
145        }
146        let row = self.first_sequence.len() + 1 - start_index;
147        let col = self.second_sequence.len() + 1 - start_index;
148        let cells = row
149            .checked_mul(col)
150            .ok_or(NeedlemanWunschError::MatrixTooLarge {
151                cells: usize::MAX,
152                max_cells: max_matrix_cells,
153            })?;
154        if cells > max_matrix_cells {
155            return Err(NeedlemanWunschError::MatrixTooLarge {
156                cells,
157                max_cells: max_matrix_cells,
158            });
159        }
160
161        let mut previous_scores = (0..col).map(|j| gap_penalty * j as i32).collect::<Vec<_>>();
162        let mut current_scores = vec![0; col];
163        let mut tracing_matrix = vec![Trace::Stop; cells];
164        for trace in tracing_matrix.iter_mut().take(col).skip(1) {
165            *trace = Trace::Left;
166        }
167        for i in 1..row {
168            current_scores[0] = gap_penalty * i as i32;
169            tracing_matrix[i * col] = Trace::Up;
170            for j in 1..col {
171                let match_value: i32 = scoring_function(
172                    self.first_sequence[i - 1 + start_index].clone(),
173                    self.second_sequence[j - 1 + start_index].clone(),
174                );
175
176                let diagonal_score = previous_scores[j - 1]
177                    + if match_value > 0 {
178                        match_value
179                    } else {
180                        mismatch_penalty
181                    };
182                let horizontal_score = current_scores[j - 1] + gap_penalty;
183                let vertical_score = previous_scores[j] + gap_penalty;
184
185                let (score, trace) =
186                    if diagonal_score >= horizontal_score && diagonal_score >= vertical_score {
187                        (
188                            diagonal_score,
189                            if match_value > 0 {
190                                Trace::MatchDiagonal
191                            } else {
192                                Trace::MisMatchDiagonal
193                            },
194                        )
195                    } else if horizontal_score >= vertical_score {
196                        (horizontal_score, Trace::Left)
197                    } else {
198                        (vertical_score, Trace::Up)
199                    };
200                current_scores[j] = score;
201                tracing_matrix[i * col + j] = trace;
202            }
203            std::mem::swap(&mut previous_scores, &mut current_scores);
204        }
205
206        let mut aligned_seq: Vec<(Option<usize>, Option<usize>)> = vec![];
207        let (mut i, mut j) = (row - 1, col - 1);
208
209        while i > 0 || j > 0 {
210            match tracing_matrix[i * col + j] {
211                Trace::MatchDiagonal => {
212                    aligned_seq.push((Some(i - 1 + start_index), Some(j - 1 + start_index)));
213                    i -= 1;
214                    j -= 1;
215                }
216                Trace::Up => {
217                    aligned_seq.push((Some(i - 1 + start_index), None));
218                    i -= 1;
219                    self.num_mismatches += 1;
220                }
221                Trace::Left => {
222                    aligned_seq.push((None, Some(j - 1 + start_index)));
223                    j -= 1;
224                    self.num_mismatches += 1;
225                }
226                Trace::MisMatchDiagonal => {
227                    aligned_seq.push((None, None));
228                    self.num_mismatches += 2;
229                    self.mismatches
230                        .push((i - 1 + start_index, j - 1 + start_index));
231                    i -= 1;
232                    j -= 1;
233                }
234                Trace::Stop => unreachable!("traceback stopped before reaching matrix origin"),
235            }
236        }
237        self.mismatches.reverse();
238        aligned_seq.reverse();
239
240        Ok((start_index, aligned_seq))
241    }
242}
243
244#[cfg(test)]
245mod tests {
246    use super::*;
247
248    #[test]
249    fn test_single_complete_match() {
250        let mut sw_object = NeedlemanWunsch {
251            first_sequence: vec![10],
252            second_sequence: vec![10],
253            ..Default::default()
254        };
255        assert_eq!(sw_object.match_sequences_base(), (1, vec![]));
256    }
257
258    #[test]
259    fn test_single_complete_mismatch() {
260        let mut sw_object = NeedlemanWunsch {
261            first_sequence: vec![0],
262            second_sequence: vec![1],
263            ..Default::default()
264        };
265        assert_eq!(sw_object.match_sequences_base(), (0, vec![(None, None)],));
266    }
267
268    #[test]
269    fn test_partial_one_off_match() {
270        let mut sw_object = NeedlemanWunsch {
271            first_sequence: vec![3, 3, 4, 4, 3, 1, 2, 4, 1],
272            second_sequence: vec![4, 3, 4, 4, 1, 2, 3, 3],
273            ..Default::default()
274        };
275        assert_eq!(
276            sw_object.match_sequences_base(),
277            (
278                0,
279                vec![
280                    (None, None),
281                    (Some(1), Some(1)),
282                    (Some(2), Some(2)),
283                    (Some(3), Some(3)),
284                    (Some(4), None),
285                    (Some(5), Some(4)),
286                    (Some(6), Some(5)),
287                    (None, None),
288                    (None, None)
289                ]
290            )
291        );
292    }
293
294    #[test]
295    fn test_gap_match() {
296        let mut sw_object = NeedlemanWunsch {
297            first_sequence: vec![1, 2, 3],
298            second_sequence: vec![1, 4, 3],
299            ..Default::default()
300        };
301        assert_eq!(
302            sw_object.match_sequences_base(),
303            (1, vec![(None, None), (Some(2), Some(2))])
304        );
305    }
306
307    #[test]
308    fn test_gap_start_match() {
309        let mut sw_object = NeedlemanWunsch {
310            first_sequence: vec![2, 4, 3],
311            second_sequence: vec![1, 4, 3],
312            ..Default::default()
313        };
314        assert_eq!(
315            sw_object.match_sequences_base(),
316            (
317                0,
318                vec![(None, None), (Some(1), Some(1)), (Some(2), Some(2))],
319            )
320        );
321    }
322
323    #[test]
324    fn bounded_alignment_rejects_large_matrices() {
325        let mut sw_object = NeedlemanWunsch {
326            first_sequence: vec![1; 100],
327            second_sequence: vec![2; 100],
328            ..Default::default()
329        };
330        assert_eq!(
331            sw_object.match_sequences_bounded(None, None, None, 10_000),
332            Err(NeedlemanWunschError::MatrixTooLarge {
333                cells: 10_201,
334                max_cells: 10_000,
335            })
336        );
337    }
338
339    #[test]
340    fn alignment_includes_unmatched_suffix() {
341        let mut sw_object = NeedlemanWunsch {
342            first_sequence: vec![1, 2],
343            second_sequence: vec![1, 2, 3],
344            ..Default::default()
345        };
346        assert_eq!(sw_object.match_sequences_base(), (2, vec![(None, Some(2))]));
347    }
348}