Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 27 additions & 0 deletions src/embkit/align.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,3 +110,30 @@ def procrustes_scale(X, Y):
denominators = np.sum(A * A, axis=0)
k = np.divide(numerators, denominators, out=np.zeros_like(numerators), where=denominators!=0)
return R, k, S


def procrustes_scale_centered(X, Y):
"""
Same as procrustes_scale, but mean-centers X and Y before fitting the rotation and scale.

To apply transformation:
(src - Xmean).dot(R) * k + Ymean # element-wise multiplication w per-dim scaling factors

Args:
X: The first matrix (N_points, N_dims).
Y: The second matrix (N_points, N_dims).

Returns:
R: The optimal rotation matrix (guaranteed det(R) = +1).
k: Per-dimension Scaling factors (shape: (N_dims,)), one scaling value per dim
Xmean: Per-dimension mean of X (shape: (N_dims,)), subtract before applying R/k
Ymean: Per-dimension mean of Y (shape: (N_dims,)), add back after applying R/k
S: Singular values of (X - Xmean).T @ (Y - Ymean)

"""
Xmean = np.array(X).mean(axis=0)
Ymean = np.array(Y).mean(axis=0)
Xc = np.array(X) - Xmean
Yc = np.array(Y) - Ymean
R, k, S = procrustes_scale(Xc, Yc)
return R, k, Xmean, Ymean, S
62 changes: 62 additions & 0 deletions tests/preprocessing/test_align.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,9 @@
from embkit.align import (
calc_rmsd,
procrustes,
procrustes_scale,
matrix_spearman_alignment_linear,
procrustes_scale_centered,
)


Expand Down Expand Up @@ -139,6 +141,66 @@ def test_procrustes_singular_values_are_sorted_nonnegative(self):
self.assertTrue(np.all(S >= 0))
np.testing.assert_array_equal(S, np.sort(S)[::-1])

# -------- procrustes_scale --------
def test_procrustes_scale_identity(self):
X = np.eye(3)
Y = np.eye(3)
R, k, S = procrustes_scale(X, Y)
np.testing.assert_array_almost_equal(R, np.eye(3))
np.testing.assert_array_almost_equal(k, np.ones(3))

def test_procrustes_scale_recovers_rotation_and_scale(self):
theta = np.pi / 5
R_true = np.array(
[[np.cos(theta), -np.sin(theta)],
[np.sin(theta), np.cos(theta)]]
)
k_true = np.array([3.0, 3.0])

rng = np.random.default_rng(5)
X = rng.standard_normal((80, 2))
Y = (X @ R_true) * k_true

R, k, S = procrustes_scale(X, Y)

np.testing.assert_array_almost_equal(R, R_true, decimal=6)
np.testing.assert_array_almost_equal(k, k_true, decimal=6)

Y_pred = (X @ R) * k
np.testing.assert_array_almost_equal(Y_pred, Y, decimal=6)

# -------- procrustes_scale_centered --------
def test_procrustes_scale_centered_means_are_correct(self):
rng = np.random.default_rng(3)
X = rng.standard_normal((40, 2)) + np.array([5.0, -3.0])
Y = rng.standard_normal((40, 2)) + np.array([-2.0, 7.0])

R, k, Xmean, Ymean, S = procrustes_scale_centered(X, Y)

np.testing.assert_array_almost_equal(Xmean, X.mean(axis=0))
np.testing.assert_array_almost_equal(Ymean, Y.mean(axis=0))

def test_procrustes_scale_centered_recovers_rotation_scale_and_translation(self):
theta = np.pi / 3
R_true = np.array(
[[np.cos(theta), -np.sin(theta)],
[np.sin(theta), np.cos(theta)]]
)
k_true = np.array([2.0, 2.0])
translation = np.array([10.0, -4.0])

rng = np.random.default_rng(4)
X = rng.standard_normal((100, 2)) + np.array([3.0, 3.0])
Y = ((X - X.mean(axis=0)) @ R_true) * k_true + translation

R, k, Xmean, Ymean, S = procrustes_scale_centered(X, Y)

np.testing.assert_array_almost_equal(R, R_true, decimal=6)
np.testing.assert_array_almost_equal(k, k_true, decimal=6)

Y_pred = ((X - Xmean) @ R) * k + Ymean
np.testing.assert_array_almost_equal(Y_pred, Y, decimal=6)


if __name__ == '__main__':
unittest.main()
Loading