From 0adc2cb53b21eb5f504810dc7fd5c7977fe08a3a Mon Sep 17 00:00:00 2001 From: Jan Bednarik Date: Mon, 27 Jul 2026 09:14:42 -0700 Subject: [PATCH] Add type annotations, parameterize a test in gnm_common_test.py. PiperOrigin-RevId: 954664192 --- gnm/shape/data/versions/v3_0/gnm_head.npz | Bin 53305389 -> 53305389 bytes gnm/shape/gnm_common_test.py | 87 +++++++++++++--------- gnm/shape/gnm_landmarks_test.py | 1 + 3 files changed, 52 insertions(+), 36 deletions(-) diff --git a/gnm/shape/data/versions/v3_0/gnm_head.npz b/gnm/shape/data/versions/v3_0/gnm_head.npz index de5f5281562d8aa25a9fb666d847f56fd3ef4f6e..d3c45a3d34afc36a7d3b88b030162c9c181d8d3e 100644 GIT binary patch delta 5644 zcmajjcT|=21IKYL7Z4QzWe6&QAfn*KtvFC|rMb(!_so%KIU9~p5f>(jCC*BVL{V`A zZq#y@d*{lP3hwOl_8#xg)8kJ&=lweT@wv})Jlvkk{U&Fu^-WeTEgG90JnfvFo$az4 z&2p*h6+W@=fIM~qcKm0u^RcsA{WR=SQp>I0CXe_g{@%0V+S+S)rIKNIcdBieTGy1~ zr&)Y<(K3cz#?nWxU2qe#RmDBxiv+m6W9kz>FrWqVEdil#o2I%ITWGdkHqfw38Tb0# zed0ZWn={mC6am;a2 zxzZeSoU1M66;X+~gAA`|5y&fE-ZCaGF}KR2@}|Wlx(ry%s~*?L7hrgmUDpS;`$R?N zON)whov@1auG_9*hV{Yg_g}5keDAnsEYULQ!UAuTPs07#ZUqhdjBU0(Z@h$t@oxO~ z6B6UyN*Kl|^IsqP%$?&lg3G(l^EUYf_2?emKRSHy&{3^hIog^2`v!mQr7?TwxG8s) zUlmXu%2RnMZ>FC|^}b6;?%*pDL>SRe&m{0#%SIu7Xtw6{1S2QYuu1snW`#!c`en zR+UrbRRvX1RZ^8z6&0bXs%ombs-YrPO;tss;la$`l^9?Pc>AH)cdNjYNDE| zW~#Yrp<1dB)Q75-YOUI+wyK?KuR5rXs*~!hx~Q(Io9eE5sGh2q>aF^yzN(+LWEs4OTIe0sTCP^8pVZGPS^c6`suZNj;j9aM+ZVRb|uRmart>JN2Xolqy$DRo+%QD@aTbzWUi z7u6+oS^cT5sH^Ilx~^`hztl~2OWjs?)Lr$rx~J}|2P#|TsE6v2daRzPr|Ow{uKrOk zV)xE5%}`(uW^e#UaDrUm47ni>}-?vNh}fCqSjmnmHj1yel-`v=xXO_-PE zKcAjM8*O>@&`a;BhtXI^wmkL9ymcFq?I^`Jg9fcKywG=le`F*Jdu&E2KjPWWqMs4m)5c?1J5}2eM!{ z01xm4FYtzf-~+x;2ns_H@Pne@4*^gN0wD;BLok$p5GVlKKN9Y8dp$l|{ZqOZiKu_oey`c~Eg?`W<2Eai02nNAm7y?7#V;Ba} zFdRm}C-5nZgi$aW#=uw@2cN;`@CAGc<6#0!gh}uf#K2^j0#hLtra>G`hZ!&vX2EQT zhdJ;y%!LH_2IfH`%!hB`J6Hfoun-o(Vpsx8VHtc6KfsT$99FaX0}d;S`*PGjJBp!Fjj<7vU0IhCks7T!m|J9d5v1a1(C9ZMXw>;cvJH z_u&C#Lk>KINAMV)z*Bez&*2|@zWkPlqJ4cs9= z6aWwK1TXN0g5U$bPzVY`5%7be;12;%3<4ntibF7zfDkBYPS@kl>^qM?L2q^lHCHhF z33@(j(5pXy`4*qlxnrn#9P`uqq2k--Ga@>Lnnjp7J)%=-^YzDkp$}?fF}E~)`?sD^ zWp=ih{TZ#$$as)tJ)g3Tg4$Zlk%lW-52tLSxf+#y*XWr>5$_reX!oWoSx=*Ew>zj& z6~mRR=TNp$WP6LbvSDOBgtCp6YgFAZvYtTMMuj?9%*J!P|BNPSRNinU>k*UfO1CsB zXBb&em29I{9WCZ?!^nDwWE-v5sH$OPJv*|ELOQ+iclYwZ$TmvUsG{LY*3%x_$U*NJ z8%?5a`fh}cGfFtbWx*7!~5C#{m(XP-GyK9JFc`%qnd`%%kOeYwkuWcYB5I{ XM%Mq^Y@=k2>bzy-mE_yaV)pzO*;EUW delta 5644 zcmajjcXW;S1IO{)+=L(!WJru8GK7#=L5!%_dvCE185c5G_aX#4xT-Jg3LkK&y7>*SBm^IT4@=O*8zecL^glw*_nW;<6K2L}h6I$^UN ztGnIo-YY7*jgJj~Og0`ivjT%AXCyY=>uz$1YvAoZdwOeI4R4e+45MwV!{q9w%>^}! zd+Jrfuq%0_@H;o$&1{RmOI!gT=TA&M;`;bBVZO&F=ws7l=R!-&mfQLoc0McLeRYpG zmw<*$y9N|BOmmH~PLq=ZPW50nIkw10cjscmDzVFRFxaq*_6m6W4^QrOFvON!e0t!2 z>`tV)n_S0drpIKao0P3GD?4SMl9`?*qnNIb-}y9)zu^Wxq5f}QD8~L-7L_%{{+UC? z69<#lyXRy5h}C#!OGhrn0LXDbo`iqn2^ii^@5D3|DRT<*GMtjhKk=oGB6E zPUE+*K5*Zum|;EZwC~$>isu8Teap2>yu8@mdjAdn&ZRMj<~plfD!0m` zT$HPFQ|>CS@=%^CpUSTasDjE%c`F}PNck#1<*x!%VHK!~sG=%J6;s7kunJKnR7q7z zl~!d`s4A<Z*FGzG|Qvsz$1@ zYNDE|W~#Yrp<1d|sPt0TjZh=iS89|Rt;VRaDq4N5#;NgYf|{r%sTeg`O;J-- zteU2#s~Kvhnx$r|IVw)gRrAz*6|WYkg(^WUQs1a=)nb*XmZ+s_nOd$^sFiA!TCLWo zwQ8MOuQsS8wNY(So7EPzRVAx!YPFSU=td6Lo>XXN#wuBfZ( zn!2uTs7!TJ-BP#J9d%dTQ}@*a^-%q(9;wIbiF&G@spsl1^+LT=uheVxxB5r@8+&N3 zX{G{OFoPY~Ll(#i4v-D9Lk@6+oZtk`kPC7{9&iCya5JUqp@1*iy>pfZF*75EIQLIhNU>QDn}LM^Bbb)YWPgZj__8bTvz3{9XZG=t{Q0$M^V zXbo+kEwqF75D6WiBXok!&;`0eH|P#MpeOW#&!IO&K_B=6`a(bG4+CHz41&Qh1ct&e z_!5T02p9=p!6+CFV_+;q!`Cnl#=``d2$LWNCc_k%3b8N^ro#-F3A11}%z-$V3-e$; z#KQtu2nnzVzJYIHF(kqgSPIKvIjn${unJbg8dwYKU_ESrB-jX>U^8rit&j}cU^{#V zDX;@}!Y;4mD4qi_t4Lk66HlW+=7!!K|K z&cZo355K~1Z~-pD@9+m)g3E9PuEI6A4mThZZo)0N4R_!!+=Kh@03O1h@CY8m6L<>G z;5qySFW@D-g4ggj`~&}*Vo#*m>I)Otf*I_<9>&$e1qa9m*&zowLQZf3XUGM)ArH8K zE4YC>uA14#viNS8@<*j>{EA&YV)BxSx%#@ zFPEWFdBdG7=TO#Bc-vrexlfJOYE;QEvYbF!@071yu=xY$cyErWtfPqQgE z-P5R)VPrW~vW{A`4>pGwMwUY)>u85Y6$~TG*^zY=82O>^o#nvDI!e$e)Nm)uX^(Ye zr}vEU<=z_&{@}>+>zegWS2QxdpWYj_?D(N?@S9)L60LXIuTd4lU;N#X{`-{c6l{+8 YkN-ZFJ86`pQMHeZ+!8%I2b*302Q#kn9smFU diff --git a/gnm/shape/gnm_common_test.py b/gnm/shape/gnm_common_test.py index 98a914fa..4bfb5aa2 100644 --- a/gnm/shape/gnm_common_test.py +++ b/gnm/shape/gnm_common_test.py @@ -12,13 +12,16 @@ # See the License for the specific language governing permissions and # limitations under the License. +"""Tests for the backend-agnostic GNM math.""" + +from collections.abc import Sequence + from absl.testing import absltest from absl.testing import parameterized from etils import enp from gnm.shape import gnm_common import numpy as np - # The array backends supported by the GNM math, as (name, xnp module) pairs. # Every test is run against all of them so that backend-specific regressions # are caught. The xnp modules mirror those used by the per-backend GNM classes @@ -35,21 +38,21 @@ class ArrayHelpersTest(parameterized.TestCase): """Tests for the backend-agnostic array helpers.""" @parameterized.named_parameters(*_XNP_BACKENDS) - def test_take_selects_along_axis(self, xnp): + def test_take_selects_along_axis(self, xnp: enp.NpModule): array = xnp.asarray([[0, 1, 2], [3, 4, 5], [6, 7, 8]], dtype=xnp.float32) indices = xnp.asarray([0, 2], dtype=xnp.int32) taken = gnm_common.take(array, indices, axis=0, xnp=xnp) np.testing.assert_array_equal(np.asarray(taken), [[0, 1, 2], [6, 7, 8]]) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_take_selects_along_last_axis(self, xnp): + def test_take_selects_along_last_axis(self, xnp: enp.NpModule): array = xnp.asarray([[0, 1, 2], [3, 4, 5]], dtype=xnp.float32) indices = xnp.asarray([2, 0], dtype=xnp.int32) taken = gnm_common.take(array, indices, axis=-1, xnp=xnp) np.testing.assert_array_equal(np.asarray(taken), [[2, 0], [5, 3]]) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_eye_returns_identity(self, xnp): + def test_eye_returns_identity(self, xnp: enp.NpModule): reference = xnp.asarray([0.0, 0.0, 0.0, 0.0], dtype=xnp.float32) identity = gnm_common.eye( 3, dtype=xnp.float32, reference_array=reference, xnp=xnp @@ -58,7 +61,7 @@ def test_eye_returns_identity(self, xnp): np.testing.assert_array_equal(np.asarray(identity), np.eye(3)) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_reshape_with_batch_dims(self, xnp): + def test_reshape_with_batch_dims(self, xnp: enp.NpModule): reference = xnp.asarray(np.zeros((2, 5, 3)), dtype=xnp.float32) array = xnp.asarray(np.arange(2 * 6), dtype=xnp.float32) reshaped = gnm_common.reshape_with_batch_dims( @@ -70,7 +73,7 @@ def test_reshape_with_batch_dims(self, xnp): self.assertEqual(tuple(reshaped.shape), (2, 6)) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_zeros_with_batch_dims(self, xnp): + def test_zeros_with_batch_dims(self, xnp: enp.NpModule): reference = xnp.asarray(np.ones((4, 5, 3)), dtype=xnp.float32) zeros = gnm_common.zeros_with_batch_dims( reference, @@ -86,36 +89,42 @@ class AxisAngleToRotationMatrixTest(parameterized.TestCase): """Tests for axis_angle_to_rotation_matrix.""" @parameterized.named_parameters(*_XNP_BACKENDS) - def test_zero_returns_identity(self, xnp): + def test_zero_returns_identity(self, xnp: enp.NpModule): axis_angle = xnp.asarray([0.0, 0.0, 0.0], dtype=xnp.float32) matrix = np.asarray(gnm_common.axis_angle_to_rotation_matrix(axis_angle)) np.testing.assert_allclose(matrix, np.eye(3), atol=1e-5) - @parameterized.named_parameters(*_XNP_BACKENDS) - def test_axis_is_invariant(self, xnp): + @parameterized.product( + xnp=tuple(xnp for _, xnp in _XNP_BACKENDS), + case=( + ([np.pi / 2, 0.0, 0.0], [1.0, 0.0, 0.0]), + ([0.0, np.pi / 2, 0.0], [0.0, 1.0, 0.0]), + ([0.0, 0.0, np.pi / 2], [0.0, 0.0, 1.0]), + ), + ) + def test_axis_is_invariant( + self, + xnp: enp.NpModule, + case: tuple[Sequence[float], Sequence[float]], + ) -> None: # Rotating about an axis leaves that axis unchanged. - cases = ( - ([np.pi / 2, 0.0, 0.0], [1.0, 0.0, 0.0]), - ([0.0, np.pi / 2, 0.0], [0.0, 1.0, 0.0]), - ([0.0, 0.0, np.pi / 2], [0.0, 0.0, 1.0]), + axis_angle, axis = case + matrix = np.asarray( + gnm_common.axis_angle_to_rotation_matrix( + xnp.asarray(axis_angle, dtype=xnp.float32) + ) ) - for axis_angle, axis in cases: - matrix = np.asarray( - gnm_common.axis_angle_to_rotation_matrix( - xnp.asarray(axis_angle, dtype=xnp.float32) - ) - ) - np.testing.assert_allclose(matrix @ np.asarray(axis), axis, atol=1e-5) + np.testing.assert_allclose(matrix @ np.asarray(axis), axis, atol=1e-5) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_ninety_degrees_about_z(self, xnp): + def test_ninety_degrees_about_z(self, xnp: enp.NpModule): axis_angle = xnp.asarray([0.0, 0.0, np.pi / 2], dtype=xnp.float32) matrix = np.asarray(gnm_common.axis_angle_to_rotation_matrix(axis_angle)) expected = np.array([[0.0, -1.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 1.0]]) np.testing.assert_allclose(matrix, expected, atol=1e-5) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_matrix_is_orthonormal_with_unit_determinant(self, xnp): + def test_matrix_is_orthonormal_with_unit_determinant(self, xnp: enp.NpModule): rng = np.random.default_rng(0) axis_angle = rng.uniform(-3.0, 3.0, size=(8, 3)) matrices = np.asarray( @@ -131,7 +140,7 @@ def test_matrix_is_orthonormal_with_unit_determinant(self, xnp): np.testing.assert_allclose(np.linalg.det(matrices), np.ones(8), atol=1e-4) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_preserves_batch_shape(self, xnp): + def test_preserves_batch_shape(self, xnp: enp.NpModule): axis_angle = xnp.asarray(np.zeros((2, 4, 3)), dtype=xnp.float32) matrices = gnm_common.axis_angle_to_rotation_matrix(axis_angle) self.assertEqual(tuple(matrices.shape), (2, 4, 3, 3)) @@ -147,7 +156,7 @@ def setUp(self): self.parents = [0, 0] @parameterized.named_parameters(*_XNP_BACKENDS) - def test_rest_pose_places_joints_at_bind_positions(self, xnp): + def test_rest_pose_places_joints_at_bind_positions(self, xnp: enp.NpModule): transforms = np.asarray( gnm_common.joint_transforms_world( xnp.asarray(self.joints, dtype=xnp.float32), @@ -166,7 +175,7 @@ def test_rest_pose_places_joints_at_bind_positions(self, xnp): ) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_global_translation_shifts_all_joints(self, xnp): + def test_global_translation_shifts_all_joints(self, xnp: enp.NpModule): transforms = np.asarray( gnm_common.joint_transforms_world( xnp.asarray(self.joints, dtype=xnp.float32), @@ -182,7 +191,7 @@ def test_global_translation_shifts_all_joints(self, xnp): ) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_root_rotation_propagates_to_child(self, xnp): + def test_root_rotation_propagates_to_child(self, xnp: enp.NpModule): # Rotate the root 90 degrees about z; the child at (1, 0, 0) maps to # (0, 1, 0). rotations = xnp.asarray( @@ -213,7 +222,7 @@ def setUp(self): self.vertices = [[0.5, 0.0, 0.0], [0.0, 0.5, 0.0], [1.0, 1.0, 0.0]] @parameterized.named_parameters(*_XNP_BACKENDS) - def test_rest_pose_is_identity(self, xnp): + def test_rest_pose_is_identity(self, xnp: enp.NpModule): posed = np.asarray( gnm_common.linear_blend_skinning( xnp.asarray(self.vertices, dtype=xnp.float32), @@ -228,7 +237,7 @@ def test_rest_pose_is_identity(self, xnp): np.testing.assert_allclose(posed, self.vertices, atol=1e-5) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_global_translation_translates_vertices(self, xnp): + def test_global_translation_translates_vertices(self, xnp: enp.NpModule): translation = [10.0, 0.0, 0.0] posed = np.asarray( gnm_common.linear_blend_skinning( @@ -245,7 +254,7 @@ def test_global_translation_translates_vertices(self, xnp): ) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_preserves_batch_shape(self, xnp): + def test_preserves_batch_shape(self, xnp: enp.NpModule): batch_vertices = np.array(np.broadcast_to(self.vertices, (4, 3, 3))) batch_joints = np.array(np.broadcast_to(self.joints, (4, 2, 3))) posed = gnm_common.linear_blend_skinning( @@ -263,7 +272,9 @@ class BindPoseTest(parameterized.TestCase): """Tests for the bind-pose vertex and joint helpers.""" @parameterized.named_parameters(*_XNP_BACKENDS) - def test_vertex_positions_none_params_return_template(self, xnp): + def test_vertex_positions_none_params_return_template( + self, xnp: enp.NpModule + ): template = xnp.asarray(np.ones((3, 3)), dtype=xnp.float32) identity_basis = xnp.asarray(np.zeros((2, 3, 3)), dtype=xnp.float32) expression_basis = xnp.asarray(np.zeros((2, 3, 3)), dtype=xnp.float32) @@ -273,7 +284,9 @@ def test_vertex_positions_none_params_return_template(self, xnp): np.testing.assert_allclose(np.asarray(result), np.ones((3, 3)), atol=1e-5) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_vertex_positions_apply_identity_and_expression(self, xnp): + def test_vertex_positions_apply_identity_and_expression( + self, xnp: enp.NpModule + ): template = xnp.asarray(np.zeros((2, 3)), dtype=xnp.float32) # (I=1, V=2, 3) identity_basis = xnp.asarray( @@ -294,7 +307,9 @@ def test_vertex_positions_apply_identity_and_expression(self, xnp): np.testing.assert_allclose(np.asarray(result), expected, atol=1e-5) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_joint_positions_none_identity_returns_template(self, xnp): + def test_joint_positions_none_identity_returns_template( + self, xnp: enp.NpModule + ): template = xnp.asarray( [[0.0, 0.0, 0.0], [1.0, 0.0, 0.0]], dtype=xnp.float32 ) @@ -305,7 +320,7 @@ def test_joint_positions_none_identity_returns_template(self, xnp): ) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_joint_positions_apply_identity(self, xnp): + def test_joint_positions_apply_identity(self, xnp: enp.NpModule): template = xnp.asarray(np.zeros((2, 3)), dtype=xnp.float32) # (I=1, J=2, 3) joint_basis = xnp.asarray( @@ -323,7 +338,7 @@ class PoseCorrectivesTest(parameterized.TestCase): """Tests for compute_pose_correctives.""" @parameterized.named_parameters(*_XNP_BACKENDS) - def test_none_inputs_return_zeros(self, xnp): + def test_none_inputs_return_zeros(self, xnp: enp.NpModule): template = xnp.asarray(np.ones((3, 3)), dtype=xnp.float32) result = gnm_common.compute_pose_correctives( None, None, template, num_joints=2, num_vertices=3 @@ -332,7 +347,7 @@ def test_none_inputs_return_zeros(self, xnp): np.testing.assert_array_equal(np.asarray(result), np.zeros((3, 3))) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_zero_rotations_produce_no_correctives(self, xnp): + def test_zero_rotations_produce_no_correctives(self, xnp: enp.NpModule): template = xnp.asarray(np.ones((3, 3)), dtype=xnp.float32) regressor = xnp.asarray(np.ones((2 * 9, 3 * 3)), dtype=xnp.float32) result = gnm_common.compute_pose_correctives( @@ -346,7 +361,7 @@ def test_zero_rotations_produce_no_correctives(self, xnp): np.testing.assert_allclose(np.asarray(result), np.zeros((3, 3)), atol=1e-5) @parameterized.named_parameters(*_XNP_BACKENDS) - def test_preserves_batch_shape(self, xnp): + def test_preserves_batch_shape(self, xnp: enp.NpModule): template = xnp.asarray(np.ones((3, 3)), dtype=xnp.float32) regressor = xnp.asarray(np.ones((2 * 9, 3 * 3)), dtype=xnp.float32) result = gnm_common.compute_pose_correctives( diff --git a/gnm/shape/gnm_landmarks_test.py b/gnm/shape/gnm_landmarks_test.py index 37b3fa01..2cbe4b04 100644 --- a/gnm/shape/gnm_landmarks_test.py +++ b/gnm/shape/gnm_landmarks_test.py @@ -15,6 +15,7 @@ """Tests for validating and loading GNM landmarks configurations.""" from absl.testing import absltest + from gnm.shape import gnm_landmarks from gnm.shape.data.versions import gnm_specs import numpy as np