Skip to main content

smriti/ml/
alignment.rs

1//! Face alignment via 5-point similarity transform.
2//!
3//! Maps detected 5-point landmarks (left eye, right eye, nose, left mouth,
4//! right mouth) onto the canonical InsightFace 112x112 template via a least-squares
5//! similarity transform (uniform scale + rotation + translation), then bilinearly
6//! samples the source image to produce a 112x112 RGB crop ready for the embedder.
7//!
8//! Replaces the naive eye-center-crop approach which did not correct for head
9//! tilt and produced different embeddings for the same face at small rotations.
10
11use image::{DynamicImage, Rgb, RgbImage};
12
13/// InsightFace canonical 5-point template (ArcFace / GLinTR) at 112x112.
14/// Order: left eye, right eye, nose, left mouth, right mouth.
15pub const CANONICAL_TEMPLATE_112: [(f32, f32); 5] = [
16    (38.2946, 51.6963),
17    (73.5318, 51.5014),
18    (56.0252, 71.7366),
19    (41.5493, 92.3655),
20    (70.7299, 92.2041),
21];
22
23/// 2x3 affine transform: q = M * [p_x, p_y, 1]^T.
24///
25/// Represents a similarity transform (uniform scale + rotation + translation).
26#[derive(Debug, Clone, Copy)]
27pub struct SimilarityTransform {
28    /// a = s*cos(theta)
29    pub a: f32,
30    /// b = s*sin(theta)
31    pub b: f32,
32    pub tx: f32,
33    pub ty: f32,
34}
35
36impl SimilarityTransform {
37    /// Apply inverse: target point -> source point.
38    #[inline]
39    pub fn apply_inverse(&self, q: (f32, f32)) -> Option<(f32, f32)> {
40        let det = self.a * self.a + self.b * self.b;
41        if det <= f32::EPSILON {
42            return None;
43        }
44        let qx = q.0 - self.tx;
45        let qy = q.1 - self.ty;
46        Some((
47            (self.a * qx + self.b * qy) / det,
48            (-self.b * qx + self.a * qy) / det,
49        ))
50    }
51
52    #[cfg(test)]
53    #[inline]
54    fn apply(&self, p: (f32, f32)) -> (f32, f32) {
55        (
56            self.a * p.0 - self.b * p.1 + self.tx,
57            self.b * p.0 + self.a * p.1 + self.ty,
58        )
59    }
60}
61
62/// Least-squares similarity transform from source -> target (both 5 points).
63///
64/// Solves for (a, b, tx, ty) such that:
65///   target_x = a*src_x - b*src_y + tx
66///   target_y = b*src_x + a*src_y + ty
67///
68/// Closed-form via centering and normal equations. Returns None if the source
69/// points are degenerate (zero spread).
70pub fn estimate_similarity(
71    src: &[(f32, f32); 5],
72    dst: &[(f32, f32); 5],
73) -> Option<SimilarityTransform> {
74    let n = src.len() as f32;
75
76    let (mut mpx, mut mpy, mut mqx, mut mqy) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
77    for i in 0..src.len() {
78        mpx += src[i].0;
79        mpy += src[i].1;
80        mqx += dst[i].0;
81        mqy += dst[i].1;
82    }
83    mpx /= n;
84    mpy /= n;
85    mqx /= n;
86    mqy /= n;
87
88    let mut num_a = 0.0f32;
89    let mut num_b = 0.0f32;
90    let mut den = 0.0f32;
91    for i in 0..src.len() {
92        let px = src[i].0 - mpx;
93        let py = src[i].1 - mpy;
94        let qx = dst[i].0 - mqx;
95        let qy = dst[i].1 - mqy;
96
97        num_a += px * qx + py * qy;
98        num_b += px * qy - py * qx;
99        den += px * px + py * py;
100    }
101
102    if den <= f32::EPSILON {
103        return None;
104    }
105
106    let a = num_a / den;
107    let b = num_b / den;
108    let tx = mqx - a * mpx + b * mpy;
109    let ty = mqy - b * mpx - a * mpy;
110
111    Some(SimilarityTransform { a, b, tx, ty })
112}
113
114/// Warp a source image onto the 112x112 canonical face template using the
115/// given 5-point landmarks. Bilinear sampling, black fill for out-of-bounds.
116///
117/// Returns None only if the landmarks are degenerate (e.g. all collinear or
118/// coincident), which signals a bad detection that should be rejected.
119pub fn align_face_112(image: &DynamicImage, landmarks: &[(f32, f32); 5]) -> Option<RgbImage> {
120    let transform = estimate_similarity(landmarks, &CANONICAL_TEMPLATE_112)?;
121    let rgb = image.to_rgb8();
122    let (w, h) = (rgb.width() as i32, rgb.height() as i32);
123
124    let mut out = RgbImage::new(112, 112);
125
126    for v in 0..112u32 {
127        for u in 0..112u32 {
128            let (sx, sy) = match transform.apply_inverse((u as f32, v as f32)) {
129                Some(p) => p,
130                None => continue,
131            };
132
133            let pixel = bilinear_sample(&rgb, sx, sy, w, h);
134            out.put_pixel(u, v, pixel);
135        }
136    }
137
138    Some(out)
139}
140
141#[inline]
142fn bilinear_sample(img: &RgbImage, x: f32, y: f32, w: i32, h: i32) -> Rgb<u8> {
143    if x < 0.0 || y < 0.0 || x > (w - 1) as f32 || y > (h - 1) as f32 {
144        return Rgb([0, 0, 0]);
145    }
146
147    let x0 = x.floor() as i32;
148    let y0 = y.floor() as i32;
149    let x1 = (x0 + 1).min(w - 1);
150    let y1 = (y0 + 1).min(h - 1);
151
152    let dx = x - x0 as f32;
153    let dy = y - y0 as f32;
154
155    let p00 = img.get_pixel(x0 as u32, y0 as u32).0;
156    let p10 = img.get_pixel(x1 as u32, y0 as u32).0;
157    let p01 = img.get_pixel(x0 as u32, y1 as u32).0;
158    let p11 = img.get_pixel(x1 as u32, y1 as u32).0;
159
160    let w00 = (1.0 - dx) * (1.0 - dy);
161    let w10 = dx * (1.0 - dy);
162    let w01 = (1.0 - dx) * dy;
163    let w11 = dx * dy;
164
165    let blend = |i: usize| -> u8 {
166        let v =
167            p00[i] as f32 * w00 + p10[i] as f32 * w10 + p01[i] as f32 * w01 + p11[i] as f32 * w11;
168        v.round().clamp(0.0, 255.0) as u8
169    };
170
171    Rgb([blend(0), blend(1), blend(2)])
172}
173
174#[cfg(test)]
175mod tests {
176    use super::*;
177
178    #[test]
179    fn identity_transform_on_canonical() {
180        let t = estimate_similarity(&CANONICAL_TEMPLATE_112, &CANONICAL_TEMPLATE_112).unwrap();
181        assert!((t.a - 1.0).abs() < 1e-4);
182        assert!(t.b.abs() < 1e-4);
183        assert!(t.tx.abs() < 1e-3);
184        assert!(t.ty.abs() < 1e-3);
185    }
186
187    #[test]
188    fn scaled_rotated_translated_recovers_parameters() {
189        let theta = 0.3f32;
190        let s = 1.7f32;
191        let (tx, ty) = (15.0f32, -8.0f32);
192        let (c, sn) = (theta.cos(), theta.sin());
193
194        let dst: [(f32, f32); 5] = core::array::from_fn(|i| {
195            let (px, py) = CANONICAL_TEMPLATE_112[i];
196            (s * (c * px - sn * py) + tx, s * (sn * px + c * py) + ty)
197        });
198
199        let t = estimate_similarity(&CANONICAL_TEMPLATE_112, &dst).unwrap();
200        assert!((t.a - s * c).abs() < 1e-3);
201        assert!((t.b - s * sn).abs() < 1e-3);
202        assert!((t.tx - tx).abs() < 1e-2);
203        assert!((t.ty - ty).abs() < 1e-2);
204    }
205
206    #[test]
207    fn inverse_roundtrip() {
208        let t = SimilarityTransform {
209            a: 0.8,
210            b: 0.4,
211            tx: 10.0,
212            ty: -5.0,
213        };
214        let p = (42.0f32, 17.0f32);
215        let q = t.apply(p);
216        let p_back = t.apply_inverse(q).unwrap();
217        assert!((p_back.0 - p.0).abs() < 1e-3);
218        assert!((p_back.1 - p.1).abs() < 1e-3);
219    }
220}