1#[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 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
53pub 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
124pub fn rigid_align_coords(
130 template: &[(i32, f64, f64)],
131 other: &[(i32, f64, f64)],
132 mapping: &[(i32, i32)], ) -> (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 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)]; 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 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 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 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}