From 3f40d6b0799b2f6f12abb10d79e4fee1710653e4 Mon Sep 17 00:00:00 2001 From: GNM Team Date: Thu, 6 Aug 2026 00:58:36 -0700 Subject: [PATCH] Fix the failing gnm test. PiperOrigin-RevId: 960128318 --- gnm/shape/data/versions/v3_0/gnm_head.npz | Bin 53305389 -> 53305389 bytes gnm/shape/gnm_common.py | 33 +++++++++- gnm/shape/gnm_numpy.py | 2 +- gnm/shape/gnm_xnp.py | 74 ++++++++++++++++------ 4 files changed, 88 insertions(+), 21 deletions(-) diff --git a/gnm/shape/data/versions/v3_0/gnm_head.npz b/gnm/shape/data/versions/v3_0/gnm_head.npz index cd3acf686de6cca2934d33933c131e4b6d5986b7..198711f1f7d7dab1b40b308171927e28c77062e0 100644 GIT binary patch delta 5652 zcmaLb1yEFb1BP*y6~sVP5L;0J5naXZ!tQRdySuwku~D(HyHT;R6}!8;TWtN_XV-gY z=N|c-d49v_d(JL`>+)Z>Z)?hR%C&kWa~ux`7Z(?Yvk9H2x(0g2E8nX_d1LSxcpH*5m3XSUtQKeSc)CDvJduMCEn zSKL3d0oKW0uPm`lB7S;pHW$*P1zDf#9}+lN9Qw5G(X)NmF4?tCm~lY{BMpHkwx zzW+QEcb_WOhiU2KXLuOfEMk$*qt@2VxRA_-jl&M>t#X>&Bd&$yFKpepAI)aiOGmyi+tRJ zonxF;+(_pbmx4>JR|L+9m(g&ALwT+U9xx}Kimwtx4w#d0@M7x~--@`VH(YV{e5QXM zBe1AzWMI)mBQ{yrM_(uMH>?*1S=aijiA?n-QIz#79(Hq@oAEniNV{|S7layDvR=0SyeV=QQ1`v6`*peTq?K9 zqw=bJD!(eA3aUb?uqvX8s$#0RDxpfMKvhbWR%KLKRZf*x6;wr4NmW)=R8>_?RaZ4s zO;tZ|&x{%U|4s0OLQYKR)DhNkrb5+pHABr*VQQ9|t>&n?YMz>}7N~HwP%To6)e^N-EmO(vIeQEgJ2RfO82wyJF^Qf*f|)K0Za?N)o#UKOSGsr~AJI;ak*!|I4Ss*b5> zbzGfLC)Fu+TAfj6)j4%uT~HU*C3RU{QCHP9bzR+1H`Oh5TisE2)jf4zJx~wTBlTE4 zQBTz~^<2GBFV!pcTD?(k)jRcGeNZ3OCl#YUt1s%S`li0CAL^(2rGAH;m~NV^z!A(4 z2b{ne;(`mrgZPjD5`rrv0yjtuNgye>g9jutMd?MsR4UHDFEz%Oxqx*)V}|%d)zItE zBq!_c{sK>v@kP#SjCJQ5?rD0F&bss1>uECnuk&YDQ-4^%Y~4-V?`g_mcr8ral>1FS zuRj#YAqAv_RFE3dKw3x#>A?#!Kt{*}-rxhikQw|S3;07;$Oab34mltIazZZ14S66h zOwuJ4-KFpG=jzu1WlkRG=t{Q0$M^VXbo+kEwqF7&;dF^C+G}apeuBP?$85zLNDkI zeV{M&gZ?l82Erg13`1Zj41?h?0!G3p7!6|}7{B%LN>5KcE|w%kP~u2ZpZ_9 zAs^(20#FbNL18EYMWGlJhZ0Z{0-+R?hB8nV%0YRk02QGURE8>06{E{JVHgaD5ik-)!Dtu*!7vua!FULP2`~{R!DN^MQ(+o}!gQDc zGa(FS!EBfVb73CLhXoK03t=06KsYE*aBN& z8$`l(*a16X7wm>Tuot3WAMA$%a1ai`VK@Ru;TS~2aX0}d;S`*PGjJBp!Fjj<7vU0I zhAVItuEBM<0XN|m+=e@F7w*A*cmNOK5j=(`@D!fGb9ezS;T61wH}DqT!F%`sAK??k zz-RaZU*Q{khad10e!*`?z5JQ=fe9SJ3~|5-oFOi_Ks<;K2_PZ3LLzX3#E=A%f;)IX zGVp}tkOERdDo71!AT6YW^xy>2zgc88M-^}IJOTFl0CJHu>@W;qRavmGwk z-R-qz0fw3FM9FSer;){MG0beoNOrRwniVk2Z0ARIGrz`v`dBRpj@i`nCUv1`JP delta 5652 zcmajj30%+j1IO`xeJhC$N^&Js5hap8~6CCKGo>!pch( z$xX>BI%%bO#luG4nJzSggddxOoTzQi(HrTfO4w_Kd=8)muXKf}lc8NTD$Wfxl%Xe%4qzhCmOro zR23Des;X+ruBxjVs-}ukwN!0YN7YsJRJ5wEVpIdwP&HDGRTI@z#j0j1PBmB0sTQiG zdS10stZ|&x{%U|4 zs0OLQYKR)DhN6>TBNenVzopqRm)Vi zTCP^8m1>n*t=?1bt2Jt^`arEy>(vIeQEgJ2)fTl?ZBsdFyV{{XRJrOSwNvdY;k19;+wnsrobRXohW$0#`7=4cx&43V|mS1}`WAMZp_%*8g?> z-gVSp);G-EtbKvD8kT3_mg(GY&MEz+C<#GO3Q9v6C=2Bv7|KHkRDg<52|^(Z!l5!m zKoy9Ds!$E=P#tPOO^AY8P#fw%U8o1qP#Fb|11~~bXb0^f9y&lrcnLZ|XXpZ5p&N9E9?%ndK?3xKKF}BXL4Ozk17Q#hh9NK% zhQV+c0V81)B*JJ&f-x`_l3^T-hY9d9yaE#;1zv^MU=mD**I^300aIZbyb04`2BgAU zFcaQ}GA zgM2swN8uQJ4qw1=H~}Z&6r6@L@Fjc&U&A+W7QTh=;2fNX@8Jiy06)S-xCB4JWw-)Y z;Tl|r8*meD!ELw$ci|rV48Opy@EhER0{9*NfCump9>HUH0#D&jTiU68u6n};u3&&0 zxPu220#7ImUQh&zf;aepFZe+*C=UJ*03{$0Ns6`&$if=~#9 zaHtFsPz55PDpUhIREHW+6QZCN)P_1x7wSPY)Q1>o01crLG=?V76k?$n#6fd-4q8A< zcph3oYj^?Lz>Clp+Ch7WhYrvYUV={08M;7M=my=P2lRwqkN~}*5A=n8&>sfCKo|sr zVF(O`VK5vuSFdOE;T$l&*AssT{U048_un-nO7A%G(uoRX-HY|q~uo70mYIqOc zhc&PkK7e(w9yY*6*aVwl3v7jLkOSLc2Yd**@Dc2UUGOn{0=r=kT_II0xt9d-wq^z>jbdF2PT5 z8Lq%pxCYnZ2Hb>Oa2xKxUAPB7!!PhF{08@-0Dgx*-~l{@NAMV)z*G3sRgXVIZ`i;U z3~&Q?@PI<#35CH6ia=5D1|RSRKPU#p!5;#k1O!4!2!c{j8p=RfCxYQ7f?3zYd3LPPu!|aX6Hn&@+A@2hS@i$E+X%gv*1I`Qo19+sq-OOjcXOOJInA1Mw;OdWGsl6G z(`=n)4JAn$@=4%`$UOc-D(Lj*6V_Htv}-$1i2A?|>_)8R8_n@O=``D IwHpC{0alF9XaE2J diff --git a/gnm/shape/gnm_common.py b/gnm/shape/gnm_common.py index dd5a3109..3a5f9d67 100644 --- a/gnm/shape/gnm_common.py +++ b/gnm/shape/gnm_common.py @@ -14,10 +14,41 @@ """Backend-agnostic core math functions of the GNM model using etils.enp.""" +import typing from typing import Any, Sequence from etils import enp -enpt = enp.typing +if typing.TYPE_CHECKING: + + class FloatArray: + + @classmethod + def __class_getitem__(cls, _): + return typing.Any + + class IntArray: + + @classmethod + def __class_getitem__(cls, _): + return typing.Any + + class BoolArray: + + @classmethod + def __class_getitem__(cls, _): + return typing.Any + + class _EnpTypingMock: + # pylint: disable=invalid-name + FloatArray = FloatArray + IntArray = IntArray + BoolArray = BoolArray + # pylint: enable=invalid-name + + enpt = _EnpTypingMock() +else: + enpt = enp.typing + _EPSILON = 1e-8 diff --git a/gnm/shape/gnm_numpy.py b/gnm/shape/gnm_numpy.py index b64563ce..3ebe5691 100644 --- a/gnm/shape/gnm_numpy.py +++ b/gnm/shape/gnm_numpy.py @@ -147,7 +147,7 @@ def compute_vertex_normals( vertex_normals = np.zeros_like(vertices_flat) np.add.at( vertex_normals, - (slice(None), self.triangles, slice(None)), + (slice(None), self.triangles, slice(None)), # pyrefly: ignore[bad-argument-type] face_normals_area[:, :, None, :], ) diff --git a/gnm/shape/gnm_xnp.py b/gnm/shape/gnm_xnp.py index 396384e9..5f524165 100644 --- a/gnm/shape/gnm_xnp.py +++ b/gnm/shape/gnm_xnp.py @@ -21,6 +21,7 @@ from collections.abc import Mapping, Sequence import dataclasses import functools +import typing from typing import Any, Self from absl import logging @@ -32,7 +33,41 @@ import numpy as np import numpy.typing as npt -enpt = enp.typing +if typing.TYPE_CHECKING: + + class FloatArray: + + @classmethod + def __class_getitem__(cls, _): + return typing.Any + + class IntArray: + + @classmethod + def __class_getitem__(cls, _): + return typing.Any + + class BoolArray: + + @classmethod + def __class_getitem__(cls, _): + return typing.Any + + class _EnpTypingMock: + # pylint: disable=invalid-name + FloatArray = FloatArray + IntArray = IntArray + BoolArray = BoolArray + # pylint: enable=invalid-name + + enpt = _EnpTypingMock() + + V = typing.Any + VPruned = typing.Any + L = typing.Any +else: + enpt = enp.typing + _NONZERO_THRESHOLD = 1e-4 _EPSILON = 1e-8 @@ -107,26 +142,26 @@ class GNM(gnm_base.GNMBase): version: gnm_specs.GNMVersion variant: gnm_specs.GNMVariant - template_vertex_positions: enpt.FloatArray - template_joint_positions: enpt.FloatArray - vertex_identity_basis: enpt.FloatArray - joint_identity_basis: enpt.FloatArray - expression_basis: enpt.FloatArray + template_vertex_positions: enpt.FloatArray['V 3'] + template_joint_positions: enpt.FloatArray['J 3'] + vertex_identity_basis: enpt.FloatArray['I V 3'] + joint_identity_basis: enpt.FloatArray['I J 3'] + expression_basis: enpt.FloatArray['E V 3'] identity_names: Sequence[str] joint_names: Sequence[str] expression_names: Sequence[str] joint_parent_indices: Sequence[int] - skinning_weights: enpt.FloatArray - quads: enpt.IntArray - triangles: enpt.IntArray - quad_uvs: enpt.FloatArray - triangle_uvs: enpt.FloatArray + skinning_weights: enpt.FloatArray['J V'] + quads: enpt.IntArray['Q 4'] + triangles: enpt.IntArray['T 3'] + quad_uvs: enpt.FloatArray['Q 4 2'] + triangle_uvs: enpt.FloatArray['T 3 2'] mesh_component_names: Sequence[str] - mirror_indices: enpt.IntArray - joint_regressor: enpt.FloatArray - pose_correctives_regressor: enpt.FloatArray - bone_aligned_template_joint_orientations: enpt.FloatArray - vertex_groups: enpt.FloatArray + mirror_indices: enpt.IntArray['V'] + joint_regressor: enpt.FloatArray['J V'] + pose_correctives_regressor: enpt.FloatArray['9*J 3*V'] + bone_aligned_template_joint_orientations: enpt.FloatArray['J 3 3'] + vertex_groups: enpt.FloatArray['G V'] vertex_group_names: Sequence[str] def __init__(self, *args: Any, **kwargs: Any) -> None: @@ -143,7 +178,8 @@ def __post_init__(self): } self._xnp = enp.get_np_module(self.template_vertex_positions) self._landmarks: dict[ - gnm_landmarks.GNMLandmarksType, tuple[enpt.IntArray, enpt.FloatArray] + gnm_landmarks.GNMLandmarksType, + tuple[enpt.IntArray['L'], enpt.FloatArray['L']], ] = {} @property @@ -571,7 +607,7 @@ def vertex_group(self, name: str) -> npt.NDArray[np.floating]: def vertex_group_mask( self, *names: str, threshold: float = _NONZERO_THRESHOLD - ) -> npt.NDArray[bool]: + ) -> npt.NDArray[np.bool_]: result_mask = np.zeros(self.num_vertices, dtype=bool) for name in names: operator, inverse = '|', False @@ -629,7 +665,7 @@ def compute_vertex_normals( 'Subclasses must implement compute_vertex_normals.' ) - def prune_vertices(self, keep_vertices: enpt.IntArray['V_pruned']) -> None: + def prune_vertices(self, keep_vertices: enpt.IntArray['VPruned']) -> None: """Prunes model vertices in-place.""" xnp = self.xnp num_vertices = self.num_vertices