From 91cfd29bcf69ec3f94e6d816dd30023d5db1e667 Mon Sep 17 00:00:00 2001 From: Mirco Mazzoni Date: Sat, 5 Sep 2026 05:58:10 -0700 Subject: [PATCH] Explicitly set float32 as default precision on JAX. PiperOrigin-RevId: 976761424 --- gnm/shape/data/versions/v3_0/gnm_head.npz | Bin 53305389 -> 53305389 bytes gnm/shape/gnm_jax.py | 97 ++++++++++++++++++++++ 2 files changed, 97 insertions(+) diff --git a/gnm/shape/data/versions/v3_0/gnm_head.npz b/gnm/shape/data/versions/v3_0/gnm_head.npz index 0b3a3f337f58695a00e4132ca6f9a21301473a7f..b12cbf02a68fd822494be1b9d58b53a9f13f0488 100644 GIT binary patch delta 5658 zcmajj1ymG^0>*LJF|Z3!Y(Wu}wF4WwTd}*lTd&wHcI(>R-HYAb-QC^#zW=~G^4{q6 zp7Z~mJ!ih1-3!d#xw|rSqvuK$t6~`=lDmtWo14phd%AJ4vUyZkH@j7ImsBqNF%k@&{?(=o6-h-_QB+jrrlP6nDu#-wVyW0Fj*6?|srZWXxT}Q9LnTs)m8VLg zlB#4Xxk{l@s#Geq@>1R^jY_M0R66CW{8W0CLHR4Y%BV7_%qok@sFca;X57 zTjf!CRX&wp6;K6Lpem#at0JnXDyE975~`#srAn(Zs;nxf%Bu>hqN=1St17Chs-}We zbyY*vRJBxXRY%oT^;CV;Ks8j2RAbddHC4@2bJaq%RIOBN)kd{d?Noc!L3LD}RA<#i zbyeL|chy7nRJ~Mh)kpPJ|EPW{SoK!})Ic>z4OT zP$$(Xby}TKXVp1%UR_WZ)g^UVT~SxnHFaIxP&d^rbz9w0chx<0Up-I{)g$#-JyB29 zGxc1(P%qUh^;*4AZ`C{XUVTs>)hG2?eNkW4H}zfpP(Mdb|LPi|zy=0Hg2)gBqJkSl zgXjTSb(t;1917Gli z^pFAk!44TA6J&-gkQK5)cE|xaAr}NdZpZ_9As^(20#FbFp%4^?B2W~HL2)PnC7~3Q zhB8nV%0YRk02QGURE8>06{ALs|c&>sfCKo|srVF(O`VK5v< zz(^PcqhSn$z*rau<6#0!gh?@IU?XgT&9DWw!Zz3rJ76d50w;vRZrB5RVIS;=18@)y!C^Q8 zN8uP8hZArTPQht7183nJoQDf=5iY@HxB^$<8eE4Pa1(C9ZMXw>;U3(F2k;Oc!DDy= zPvIFnhZpb?Ucqa418?CSyoV3)5kA3Z_yS+y8+?Z!@Y7Wve>OdE1sfO;2_i!jhzf2H z4WdH~hzYSEHpGFr5D(%*0&urE_3>wmJdJxX$GnLiHebul*r@{Cj5%-VV1Qxnl1}xo zW&G2wz79T3VeV4w^sps!wA}1ri)~roYHFT!v<{AmY&_7bvrg|(@&E2G*L3B8~<^nt$c5A=gz=nn&6APj=RFa(CeFc=OaU?hx!(J%%=U@VM-@h|}X2L9(4Gx$Cb73CLhXt?@7Qtdz0!v{TEQb}a5>~-#SOaTe9ju29un{)F zX4nE-VH<3R9k3I2ffGVuH|&AEun+db0XPVU;4mD4qi_t4!wEPEr{FZ4fwOQ9&cg+` z2$$e8T!E`_4X(otxCyu5Hr#=`a1ZXo19%9J;4wUbr|=A(!wYx`ui!Pjfw%Au-opp@ z2%q3He1Wg<4Zgz<_-PwiZ>OOLu3!TLB0*$`0#U&YqCs?s0Wl#K#D+K!7ve#DNC57T z5Ii6eBnD4N0!bm6;nc^U5fsV%B;C9_xslfLNgA4J_J9BX+t-m~PjaKXIi9FT8UJnU z=WyMd+^ArVtLTySGx}CL&g`6^gJJ>b9kmQ zye#9unhxHs`5iz=uqGlkcp1)MPj92uJoCxVw|g7jmL~_l=gA>HJG>2%#&G)VNNbc` zU_PU(zE$wCJmXcAxjciH6Y0Pg{2)DKFq~@}_#1aKm}i+X%HQz+!&#h{buiFo?vjQ0 z8|f^&IUDBb)tohTuykH?x9Et!;cMBgxoYn2>EQjQ=I&T#yOH`2JM+60Uq^;4`i8|E z_thgm%Q)0FpBLYugX;`)*D=TMPI}6ZUHC~)X;Rr7r_E{q%^%GFjK^uR&>ZK^WjC^0 z-a&`di+KkFbx?nO{Kvnfz7;C`hn@Kj`kQwYX@cxVAjr^8b_?J(=Wc5lyAfcSh20K?pN${L1#|ZM zyP*iPVm0kXHp|Z~{8A75SKw@`HOu$+nU{Wo0hU?VOQiU~G$D;Bn5cXxLOc3^-=D*?qN&8&JWBmc+@`)z!%*QHDuzbGy|Y7t$t{lZO-koSl+8 zb*k?hy>g+!fx`oxHwRk2eKRySquO^%uD?u%hSjthRc1W2 z6e^`krBbUjDy{NR=~Q~Z$swfoiB4sUX!@HBn7fGu2$RP%Tv})mpVtZB;wfUUg6%RVUS1bx~ba zH`QJBP(4*I)m!yZeN{izUky+L)gU!k4N*hYFg09_P$ShSHCl~PW7Rk{UQJNJYNDE? zCaWoGs+y*zs~Kvhnx#V2Y&A#CRr6G+3RCmd0<};rQj66RwNx!r%hd|CQms;})f%-{ ztyAmO2DMRbQk&HlwN-6X+tm)WQ|(f_)gEP2d(}R*UmZ{f)gg6Q9Z^TsF?C#>P$$(X zby}TKXVp1%UR_WZ)g^UVT~SxnHFaIxP&d^rbz9w0chx<0Up-I{)g$#-JyB29Gxc1( zP%qUh^;*4AZ`C{XUVTs>)hG2?eNkW4H}zeGs~_s8`lWu44GDJ+R$u`GqJaxUhZx`r zF(DSjhBy!x;z4{!00|)xB!(oA6x<*gBnNj$0VyFBq=qz*7Caywqz6yP02#pxGC^j@ z0$Cv&c!L$PLk{qPoRAB0LmtQr`M?+aAU_mg|G+~!xC5u%V0UI zfR(TcR>K-t3+rG#Y=Dih2{ywP*b3WVJM4g+unTs>9eSMVC%z*~3+@8JV{gir7pzQ9-b2HznZe!x%o1;3qj`?KhUGg!cYXy5|TAqKcY zOo#=sAr8cacn}{FKtf0ai6M!_rrV#zWfsq3p$n$ESu+3SVa#?{2mK8DT*{elmh7gv zfOb5z*&6BKi!}DRG`rm_sX|+Cb+g3#+w^!HoZ#}`>8z%?xBc?Ps9Y$xp_UTKAUU{0 z3P=g5AT^|cwBP~hAU$|O2FM6rkO?wF7RUYnq@P}eh97;e*CQ+d zU+4$@VE_z-K`M7EQCd{7?!|NSO&{s1+0Wsuo~9DT383`VFPT0O|TiZz*g7>+hGUn zgk7*3_J9rc!amp!2jCzag2QkGj>0iG4kzFwoPyJE2F}7cI1d-#B3y#Ya0RZyHMkBp z;3nLH+i(Z&!acYT58xp@g2(U#p29PD4lm#(yn@&82HwIucn=@oBYc9-@CClYH~0?W z@B@CrFZgX4+iXrFs?zD*8u6;O5sPlo; zMiu*Tb-gm*qdnd;Y@4Ek=NH=NTD44Lq%~bZ&;(!Z<#Tn=Hn(M3W7>NAa9(#0BfV+( z?qNT!!2TIP@F0zdbdcV#4RZH1+AOfY`PEKO!_)NUfUEgAlh-a!Lu4>)Ub`|H6&Kmx z(OsV^c$wbeaj_up;B+DrWQHt|6|xz&^+DdogKYLyW{mSTyiKe4gtLn6nhyF~>~pDt zy^Tz!Ige6Bn5(0Mfod1_PLWjRwJEhJ^v}pMSO3w<EqF(tUeVgZkh{g#GL)h8f+iS5NI`u{q2(_=ekzI@!&B>WB&gh zOI*up6gHhS;=`vutBurHzJHH>)L0?YnB%D^^v`N7Yg-M!e~(>`a?J5G6lt|mb*x5i z(`u%%wHh;DTf`>+b8YGMKWM%-$M%kN(n%V#nojE2-I2yVY0Sqo7O}7YoV0#DtC7<* z=GeiJRy(LM^G}jv+eRA8QQvA5FsRa#Fm667JYb=US zF&&%o&sTa!W4@-7I<{k^)hacNYPE jt.Float[jnp.ndarray, "A1 ... An V 3"]: + """Evaluates the GNM mesh-generating function.""" + with jax.default_matmul_precision("float32"): + return super().__call__( + identity=identity, + expression=expression, + rotations=rotations, + translation=translation, + ) + + def vertex_positions_world( + self, + vertices: jt.Float[jnp.ndarray, "A1 ... An V 3"], + joints: jt.Float[jnp.ndarray, "A1 ... An J 3"], + rotations: jt.Float[jnp.ndarray, "A1 ... An J 3"], + translation: jt.Float[jnp.ndarray, "A1 ... An 3"], + ) -> jt.Float[jnp.ndarray, "A1 ... An V 3"]: + """Applies linear blend skinning to GNM vertices.""" + with jax.default_matmul_precision("float32"): + return super().vertex_positions_world( + vertices=vertices, + joints=joints, + rotations=rotations, + translation=translation, + ) + + apply_linear_blend_skinning = vertex_positions_world + + def vertex_positions_bind_pose( + self, + identity: jt.Float[jnp.ndarray, "A1 ... An I"] | None, + expression: jt.Float[jnp.ndarray, "A1 ... An E"] | None, + ) -> jt.Float[jnp.ndarray, "A1 ... An V 3"]: + """Computes vertices in the bind pose, with identity and expression applied.""" + with jax.default_matmul_precision("float32"): + return super().vertex_positions_bind_pose( + identity=identity, expression=expression + ) + + def joint_positions_bind_pose( + self, + identity: jt.Float[jnp.ndarray, "A1 ... An I"] | None, + ) -> jt.Float[jnp.ndarray, "A1 ... An J 3"]: + """Joint positions in the bind pose, with identity basis applied.""" + with jax.default_matmul_precision("float32"): + return super().joint_positions_bind_pose(identity=identity) + + def compute_pose_correctives( + self, + rotations: jt.Float[jnp.ndarray, "A1 ... An J 3"] | None, + ) -> jt.Float[jnp.ndarray, "A1 ... An V 3"]: + """Applies pose-dependent corrective shape offsets to vertices.""" + with jax.default_matmul_precision("float32"): + return super().compute_pose_correctives(rotations=rotations) + + def joint_transforms_world( + self, + joints: jt.Float[jnp.ndarray, "A1 ... An J 3"], + rotations: jt.Float[jnp.ndarray, "A1 ... An J 3"], + translation: jt.Float[jnp.ndarray, "A1 ... An 3"], + ) -> jt.Float[jnp.ndarray, "A1 ... An J 4 4"]: + """Computes the world-space transformation matrices for each joint.""" + with jax.default_matmul_precision("float32"): + return super().joint_transforms_world( + joints=joints, rotations=rotations, translation=translation + ) + + def get_posed_joint_transforms( + self, + identity: jt.Float[jnp.ndarray, "A1 ... An I"], + rotations: jt.Float[jnp.ndarray, "A1 ... An J 3"], + translation: jt.Float[jnp.ndarray, "A1 ... An 3"], + ) -> jt.Float[jnp.ndarray, "A1 ... An J 4 4"]: + """Computes the local-to-world transformation for every joint.""" + with jax.default_matmul_precision("float32"): + return super().get_posed_joint_transforms( + identity=identity, rotations=rotations, translation=translation + ) + def compute_vertex_normals( self, vertices: jt.Float[jnp.ndarray, "... V 3"], @@ -158,3 +245,13 @@ def compute_vertex_normals( ) vertex_normals = vertex_normals / jnp.maximum(normal_magnitudes, 1e-8) return vertex_normals.reshape(batch_dims + (num_vertices, 3)) + + +def axis_angle_to_rotation_matrix( + axis_angle: jt.Float[jnp.ndarray, "... 3"], + epsilon: float = 1e-8, +) -> jt.Float[jnp.ndarray, "... 3 3"]: + """Builds a 3x3 rotation matrix from an axis-angle vector.""" + with jax.default_matmul_precision("float32"): + return gnm_common.axis_angle_to_rotation_matrix(axis_angle, epsilon=epsilon) +