Skip to main content

xpict_core/
align.rs

1//! Rigid 2D alignment (Kabsch) — Indigo / no-RDKit fallback.
2//!
3//! **Template depiction** (fix matched atoms and redraw the rest) stays at
4//! the language edge via RDKit. This module only rotates/translates point
5//! sets so a caller-supplied correspondence lands on a **template** frame.
6//!
7//! Keep in sync with `python/xpict/align.py` (`_kabsch_2d`, `_apply_transform`).
8
9/// Rigid transform: `x' = cos·x ∓ sin·y + tx`, `y' = sin·x ± cos·y + ty`
10/// (sign of the sin terms flips when [`RigidTransform::det`] is negative).
11#[derive(Debug, Clone, Copy, PartialEq)]
12pub struct RigidTransform {
13    pub cos: f64,
14    pub sin: f64,
15    pub tx: f64,
16    pub ty: f64,
17    /// `+1` proper rotation, `-1` improper (reflection).
18    pub det: f64,
19}
20
21impl Default for RigidTransform {
22    fn default() -> Self {
23        Self {
24            cos: 1.0,
25            sin: 0.0,
26            tx: 0.0,
27            ty: 0.0,
28            det: 1.0,
29        }
30    }
31}
32
33impl RigidTransform {
34    pub fn apply(&self, x: f64, y: f64) -> (f64, f64) {
35        if self.det >= 0.0 {
36            (
37                self.cos * x - self.sin * y + self.tx,
38                self.sin * x + self.cos * y + self.ty,
39            )
40        } else {
41            (
42                self.cos * x + self.sin * y + self.tx,
43                self.sin * x - self.cos * y + self.ty,
44            )
45        }
46    }
47
48    pub fn apply_all(&self, pts: &[(f64, f64)]) -> Vec<(f64, f64)> {
49        pts.iter().map(|&(x, y)| self.apply(x, y)).collect()
50    }
51}
52
53/// Kabsch / Procrustes 2D: map `src` → `dst` (paired points).
54///
55/// Returns `(cos, sin, tx, ty, det_sign)`. Empty input → identity.
56pub fn kabsch_2d(src: &[(f64, f64)], dst: &[(f64, f64)], allow_reflect: bool) -> RigidTransform {
57    assert_eq!(src.len(), dst.len(), "src/dst length mismatch");
58    let n = src.len();
59    if n == 0 {
60        return RigidTransform::default();
61    }
62    let sx: f64 = src.iter().map(|p| p.0).sum::<f64>() / n as f64;
63    let sy: f64 = src.iter().map(|p| p.1).sum::<f64>() / n as f64;
64    let dx: f64 = dst.iter().map(|p| p.0).sum::<f64>() / n as f64;
65    let dy: f64 = dst.iter().map(|p| p.1).sum::<f64>() / n as f64;
66    let mut sxx = 0.0;
67    let mut syy = 0.0;
68    let mut sxy = 0.0;
69    let mut syx = 0.0;
70    for (&(x, y), &(u, v)) in src.iter().zip(dst.iter()) {
71        let x0 = x - sx;
72        let y0 = y - sy;
73        let u0 = u - dx;
74        let v0 = v - dy;
75        sxx += x0 * u0;
76        sxy += x0 * v0;
77        syx += y0 * u0;
78        syy += y0 * v0;
79    }
80
81    let rot_score = |c: f64, s: f64| c * (sxx + syy) + s * (sxy - syx);
82
83    let ang = (sxy - syx).atan2(sxx + syy);
84    let c1 = ang.cos();
85    let s1 = ang.sin();
86    let mut best_c = c1;
87    let mut best_s = s1;
88    let mut best_det = 1.0;
89    let mut best_sc = rot_score(c1, s1);
90    if allow_reflect {
91        let ang2 = (sxy + syx).atan2(sxx - syy);
92        let c2 = ang2.cos();
93        let s2 = ang2.sin();
94        let sc2 = c2 * (sxx - syy) + s2 * (sxy + syx);
95        if sc2 > best_sc {
96            best_c = c2;
97            best_s = s2;
98            best_det = -1.0;
99            best_sc = sc2;
100        }
101    }
102    let _ = best_sc;
103
104    let (tx, ty) = if best_det >= 0.0 {
105        (
106            dx - (best_c * sx - best_s * sy),
107            dy - (best_s * sx + best_c * sy),
108        )
109    } else {
110        (
111            dx - (best_c * sx + best_s * sy),
112            dy - (best_s * sx - best_c * sy),
113        )
114    };
115    RigidTransform {
116        cos: best_c,
117        sin: best_s,
118        tx,
119        ty,
120        det: best_det,
121    }
122}
123
124/// Align `other` atom coords onto a **template** using `mapping` (other→template).
125///
126/// `template` / `other` are `(index, x, y)`. Returns transformed other coords
127/// as `(index, x, y)` in the same order as `other`. Mapping pairs missing from
128/// either set are skipped; fewer than one pair → identity.
129pub fn rigid_align_coords(
130    template: &[(i32, f64, f64)],
131    other: &[(i32, f64, f64)],
132    mapping: &[(i32, i32)], // (other_index, template_index)
133) -> (Vec<(i32, f64, f64)>, RigidTransform) {
134    let tmpl: std::collections::HashMap<i32, (f64, f64)> =
135        template.iter().map(|&(i, x, y)| (i, (x, y))).collect();
136    let oth: std::collections::HashMap<i32, (f64, f64)> =
137        other.iter().map(|&(i, x, y)| (i, (x, y))).collect();
138    let mut src = Vec::new();
139    let mut dst = Vec::new();
140    for &(oi, ti) in mapping {
141        let Some(&s) = oth.get(&oi) else {
142            continue;
143        };
144        let Some(&d) = tmpl.get(&ti) else {
145            continue;
146        };
147        src.push(s);
148        dst.push(d);
149    }
150    if src.is_empty() {
151        return (other.to_vec(), RigidTransform::default());
152    }
153    let xf = kabsch_2d(&src, &dst, true);
154    let out = other
155        .iter()
156        .map(|&(i, x, y)| {
157            let (nx, ny) = xf.apply(x, y);
158            (i, nx, ny)
159        })
160        .collect();
161    (out, xf)
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    #[test]
169    fn identity_on_same_points() {
170        let pts = vec![(0.0, 0.0), (1.0, 0.0), (0.0, 1.0)];
171        let xf = kabsch_2d(&pts, &pts, true);
172        assert!((xf.cos - 1.0).abs() < 1e-9);
173        assert!(xf.sin.abs() < 1e-9);
174        assert!(xf.tx.abs() < 1e-9 && xf.ty.abs() < 1e-9);
175    }
176
177    #[test]
178    fn rotates_and_translates() {
179        // 90° CCW + translate (2, 3): (1,0)→(0,1)+offset, etc.
180        let src = vec![(1.0, 0.0), (0.0, 1.0), (-1.0, 0.0)];
181        let dst: Vec<(f64, f64)> = src
182            .iter()
183            .map(|&(x, y)| (-y + 2.0, x + 3.0))
184            .collect();
185        let xf = kabsch_2d(&src, &dst, false);
186        for (&(x, y), &(u, v)) in src.iter().zip(dst.iter()) {
187            let (nx, ny) = xf.apply(x, y);
188            assert!((nx - u).abs() < 1e-9, "{nx} vs {u}");
189            assert!((ny - v).abs() < 1e-9, "{ny} vs {v}");
190        }
191    }
192
193    #[test]
194    fn flip_allowed_when_reflected() {
195        let src = vec![(0.0, 0.0), (2.0, 0.0), (0.0, 1.0)];
196        let dst = vec![(0.0, 0.0), (2.0, 0.0), (0.0, -1.0)]; // reflect Y
197        let xf = kabsch_2d(&src, &dst, true);
198        assert!(xf.det < 0.0);
199        for (&(x, y), &(u, v)) in src.iter().zip(dst.iter()) {
200            let (nx, ny) = xf.apply(x, y);
201            assert!((nx - u).abs() < 1e-9);
202            assert!((ny - v).abs() < 1e-9);
203        }
204    }
205
206    #[test]
207    fn rigid_align_coords_onto_template() {
208        // Template: horizontal. Other: same shape, rotated 180° + shift.
209        let template = vec![(0, 0.0, 0.0), (1, 20.0, 0.0), (2, 20.0, 10.0)];
210        let other = vec![(0, 5.0, 5.0), (1, -15.0, 5.0), (2, -15.0, -5.0)];
211        // Map all atoms 1:1 by index.
212        let mapping = vec![(0, 0), (1, 1), (2, 2)];
213        let (aligned, _xf) = rigid_align_coords(&template, &other, &mapping);
214        for ((_, x, y), (_, u, v)) in aligned.iter().zip(template.iter()) {
215            assert!((x - u).abs() < 1e-6, "{x},{y} vs {u},{v}");
216            assert!((y - v).abs() < 1e-6);
217        }
218    }
219
220    #[test]
221    fn default_transform_is_identity() {
222        let xf = RigidTransform::default();
223        assert!((xf.cos - 1.0).abs() < 1e-12);
224        assert!(xf.sin.abs() < 1e-12);
225        assert!(xf.tx.abs() < 1e-12 && xf.ty.abs() < 1e-12);
226        assert!((xf.det - 1.0).abs() < 1e-12);
227        let pts = [(1.0, 2.0), (3.0, -4.0)];
228        assert_eq!(xf.apply_all(&pts), pts.to_vec());
229    }
230
231    #[test]
232    fn kabsch_empty_is_identity() {
233        let xf = kabsch_2d(&[], &[], true);
234        assert!((xf.cos - 1.0).abs() < 1e-12);
235        assert!(xf.det > 0.0);
236    }
237
238    #[test]
239    fn rigid_align_skips_missing_pairs_and_empty_mapping() {
240        let template = vec![(0, 0.0, 0.0), (1, 10.0, 0.0)];
241        let other = vec![(0, 1.0, 1.0), (1, 11.0, 1.0)];
242        // Missing other / template indices are skipped → no pairs → identity.
243        let (out, xf) = rigid_align_coords(&template, &other, &[(9, 0), (0, 9)]);
244        assert_eq!(out, other);
245        assert!((xf.cos - 1.0).abs() < 1e-12);
246        let (out2, _) = rigid_align_coords(&template, &other, &[]);
247        assert_eq!(out2, other);
248    }
249}